Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
132 changes: 124 additions & 8 deletions src/plugrl_server/algorithm/fpo/fpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,10 +88,12 @@ def __init__(self, config: FPOAlgoConfig, policy: BasePolicyGradientFlowPolicy):
+ list(self.policy.critic.parameters()),
device=config.master_weights_device,
)
self.optimizer = torch.optim.Adam(
self.master_weights.optimizer_params,
lr=config.learning_rate,
# MasterWeights keeps the order it was given and drops what is frozen,
# so the actor's optimizer parameters are the first this many.
self._n_actor_params = sum(
1 for p in self.policy.actor.parameters() if p.requires_grad
)
self.optimizer = self._build_optimizer()
self.global_step = 0
self.curr_train_itrs = 0
self.last_saved_itr = 0
Expand Down Expand Up @@ -170,6 +172,7 @@ def chunk_reduction(self) -> ChunkReduction:
steps=self.config.cfm_loss_steps,
dims=self.config.cfm_loss_dims,
sum_over_steps=self.config.cfm_loss_sum_over_steps,
huber_delta=self.config.cfm_loss_huber_delta,
)

def pre_learn(self) -> None:
Expand Down Expand Up @@ -373,7 +376,90 @@ def _scale_advantage(self, advantage: torch.Tensor) -> torch.Tensor:
# third iteration on and none of its metrics could show it, because the
# only advantage statistic logged was measured after normalising.
self._advantage_raw_std = float(advantage.std().detach().cpu())
if not self.config.normalize_advantage:
if (
not self.config.normalize_advantage
or self.config.normalize_advantage_per_minibatch
):
return advantage
return (advantage - advantage.mean()) / (advantage.std() + 1e-8)

@property
def _split_optimizer(self) -> bool:
"""Whether FPO++'s optimizer settings are in use, rather than FPO's one Adam."""
config = self.config
return (
config.critic_learning_rate is not None
or config.weight_decay != 0.0
or config.adam_eps != 1e-8
or config.actor_adam_beta2 != 0.999
)

@property
def _actor_optimizer_params(self) -> list[torch.nn.Parameter]:
return self.master_weights.optimizer_params[: self._n_actor_params]

@property
def _critic_optimizer_params(self) -> list[torch.nn.Parameter]:
return self.master_weights.optimizer_params[self._n_actor_params :]

def _build_optimizer(self) -> torch.optim.Optimizer:
"""FPO's one Adam over actor and critic, or FPO++'s two AdamW groups.

FPO++ keeps two AdamW optimizers; two parameter groups of one AdamW
are the same arithmetic and keep a single optimizer state to save.
"""
config = self.config
if not self._split_optimizer:
return torch.optim.Adam(
self.master_weights.optimizer_params, lr=config.learning_rate
)
critic_lr = (
config.learning_rate
if config.critic_learning_rate is None
else config.critic_learning_rate
)
return torch.optim.AdamW(
[
dict(
params=self._actor_optimizer_params,
lr=config.learning_rate,
betas=(0.9, config.actor_adam_beta2),
),
dict(
params=self._critic_optimizer_params,
lr=critic_lr,
betas=(0.9, 0.999),
),
],
eps=config.adam_eps,
weight_decay=config.weight_decay,
)

def _clip_gradients(self) -> tuple[float, float] | None:
"""Clip the actor's and the critic's gradients separately, as FPO++ does.

Returns their norms before clipping, or None when clipping is off.
"""
limit = self.config.max_grad_norm
if limit is None:
return None
norms = []
for params in (self._actor_optimizer_params, self._critic_optimizer_params):
with_grad = [p for p in params if p.grad is not None]
norms.append(
float(torch.nn.utils.clip_grad_norm_(with_grad, limit))
if with_grad
else 0.0
)
return norms[0], norms[1]

def _minibatch_advantage(self, advantage: torch.Tensor) -> torch.Tensor:
"""Normalise one minibatch's advantages, only when asked to, as FPO++ does.

Off by default: `_scale_advantage` explains why the buffer is the
default here.
"""
if not self.config.normalize_advantage_per_minibatch or advantage.numel() < 2:
return advantage
return (advantage - advantage.mean()) / (advantage.std() + 1e-8)

Expand Down Expand Up @@ -438,10 +524,12 @@ def _compute_loss(
initial_cfm_loss,
) -> tuple[torch.Tensor, dict[str, float], dict[str, float]]:
# `advantage` arrives already normalised over the whole buffer, in
# `_build_train_batch_cache`. It is deliberately not normalised again
# here: this method sees one minibatch, and at the batch sizes a VLA
# forces, a per-minibatch statistic manufactures signal rather than
# removing scale.
# `_build_train_batch_cache`. It is not normalised again here unless
# `normalize_advantage_per_minibatch` asks, as FPO++ does: this method
# sees one minibatch, and at the batch sizes a VLA forces, a
# per-minibatch statistic manufactures signal rather than removing
# scale.
advantage = self._minibatch_advantage(advantage)
obs_cache_started_at = time.perf_counter()
obs_cache = self.policy.build_obs_cache(obs)
_sync_cuda_if_needed(self.policy.device)
Expand Down Expand Up @@ -550,6 +638,11 @@ def learn_impl(self) -> tuple[int, dict]:
batch_to_device_total = 0.0
compute_loss_total = 0.0
stopped_early = False
grad_norms: dict[str, list[float]] = dict(actor=[], critic=[])
# How much of the return the critic that collected this buffer
# explains, read once the buffer's returns exist: before the first
# epoch, or after its refresh under `fpo_playground_trick`.
rollout_summary: dict | None = None
# Read once, here: `curr_train_itrs` is incremented at the end of this
# method, so asking again where the metrics are assembled would report
# the next iteration's answer.
Expand Down Expand Up @@ -588,13 +681,16 @@ def learn_impl(self) -> tuple[int, dict]:
ret_all = torch.from_numpy(self.rollout_buffer.returns[:num_items]).to(
self.policy.device
)
rollout_summary = self.rollout_buffer.description()

for _ in range(self.config.num_updates_per_batch):
if self.config.fpo_playground_trick:
dataloader_started_at = time.perf_counter()
value_all, advantage_all, ret_all = self._refresh_epoch_value_targets(
obs_all
)
if rollout_summary is None:
rollout_summary = self.rollout_buffer.description()
dataloader_total += time.perf_counter() - dataloader_started_at
dataloader_started_at = time.perf_counter()
indices = torch.randperm(num_items, device=self.policy.device)
Expand Down Expand Up @@ -634,6 +730,15 @@ def learn_impl(self) -> tuple[int, dict]:
if self.in_critic_warmup():
self._zero_actor_grads()
self.master_weights.grads_to_masters()
if in_warmup and self._split_optimizer:
# AdamW decays every parameter it steps, gradient or
# not, and skips one without a gradient altogether.
for param in self._actor_optimizer_params:
param.grad = None
norms = self._clip_gradients()
if norms is not None:
grad_norms["actor"].append(norms[0])
grad_norms["critic"].append(norms[1])
self.optimizer.step()
self.master_weights.masters_to_model()
_sync_cuda_if_needed(self.policy.device)
Expand All @@ -657,6 +762,7 @@ def learn_impl(self) -> tuple[int, dict]:
learn_loop_time = time.perf_counter() - learn_started_at
return self.global_step, dict(
train=dict(train_itrs=float(self.curr_train_itrs)),
rollout=rollout_summary or {},
losses=dict(
policy_loss=float(np.mean(metric_history["policy_loss"])),
value_loss=float(np.mean(metric_history["value_loss"])),
Expand All @@ -678,6 +784,16 @@ def learn_impl(self) -> tuple[int, dict]:
loss_delta_mean=float(np.mean(metric_history["loss_delta_mean"])),
loss_delta_min=float(np.min(metric_history["loss_delta_min"])),
loss_delta_max=float(np.max(metric_history["loss_delta_max"])),
# Before clipping, the largest of the iteration; only when
# `max_grad_norm` is set.
**(
dict(
actor_grad_norm_max=max(grad_norms["actor"]),
critic_grad_norm_max=max(grad_norms["critic"]),
)
if grad_norms["actor"]
else {}
),
),
learn_runtime=dict(
dataloader_time=dataloader_total,
Expand Down
30 changes: 30 additions & 0 deletions src/plugrl_server/algorithm/fpo/fpo_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,32 @@ class FPOAlgoConfig(BaseAlgoConfig):
# Off keeps one ratio per action, the samples averaged as
# `average_losses_before_exp` says and the difference clamped at +-3.
ratio_per_sample: bool = False
# FPO++'s square fine-tuning (amazon-far/fpo-control, `manipulation_
# experiments/finetune_online_rl.py`) has more than its policy loss, and
# E37, which ran only the loss, watched a behaviour-cloned policy fall
# while its randomly initialised critic never fit the returns. The rest,
# each off by default:
#
# Two AdamW optimizers. FPO++ gives the actor 1e-5 with betas (0.9, 0.99)
# and the critic 1e-4 with the default betas, both eps 1e-5 and weight
# decay 1e-6. Here they are two parameter groups of one AdamW, built when
# any of these differs from FPO's one Adam over both.
critic_learning_rate: float | None = None
adam_eps: float = 1e-8
weight_decay: float = 0.0
actor_adam_beta2: float = 0.999
# Clip the actor's and the critic's gradients separately to this norm
# (FPO++: 25 on square).
max_grad_norm: float | None = None
# Normalise the advantages within each minibatch rather than over the
# buffer, as FPO++ does with its 375-chunk minibatches. See
# `FPOAlgorithm._scale_advantage` for why the buffer is the default: at
# the minibatches of 8 a VLA forces, a per-minibatch statistic
# manufactures signal.
normalize_advantage_per_minibatch: bool = False
# FPO++'s Huber on the flow-matching error: d^2 within delta, 2 delta |d|
# - delta^2 beyond (FPO++: delta 1), before the chunk reduction.
cfm_loss_huber_delta: float | None = None
save_interval: int = 10
# A checkpoint to start this run from, and how much of it to take.
#
Expand Down Expand Up @@ -129,3 +155,7 @@ def __post_init__(self) -> None:
raise ValueError(
"restore only means something with policy_checkpoint_path set"
)
if self.normalize_advantage_per_minibatch and not self.normalize_advantage:
raise ValueError(
"normalize_advantage_per_minibatch needs normalize_advantage on"
)
26 changes: 23 additions & 3 deletions src/plugrl_server/algorithm/fpo/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,17 +18,35 @@ class ChunkReduction:
action dimensions - the ones the client executes and the environment
uses - and `sum_over_steps` sums the per-step means instead of averaging
them, as FPO++ does (Yi, Choi et al. 2026, amazon-far/fpo-control).
`huber_delta` replaces each element's squared error with FPO++'s Huber:
d^2 within delta, 2 delta |d| - delta^2 beyond, which meets d^2 at delta.
"""

steps: int | None = None
dims: int | None = None
sum_over_steps: bool = False
huber_delta: float | None = None


def _elementwise_error(diff: torch.Tensor, huber_delta: float | None) -> torch.Tensor:
if huber_delta is None:
return diff.pow(2)
magnitude = diff.abs()
return torch.where(
magnitude <= huber_delta,
diff.pow(2),
2.0 * huber_delta * magnitude - huber_delta**2,
)


def _reduce(
error: torch.Tensor, batch_size: int, sample_count: int, reduction: ChunkReduction
) -> torch.Tensor:
if reduction == ChunkReduction():
if (
reduction.steps is None
and reduction.dims is None
and not reduction.sum_over_steps
):
return error.reshape(batch_size, sample_count, -1).mean(dim=-1)
if error.dim() != 4:
raise ValueError(
Expand Down Expand Up @@ -75,8 +93,10 @@ def compute_cfm_loss(

if output_mode == "u":
target = loss_eps - action.unsqueeze(1)
return _reduce((v - target).pow(2), batch_size, sample_count, reduction)
error = _elementwise_error(v - target, reduction.huber_delta)
return _reduce(error, batch_size, sample_count, reduction)

x0_pred = x_t - loss_t_expand * v
x1_pred = x0_pred + v
return _reduce((loss_eps - x1_pred).pow(2), batch_size, sample_count, reduction)
error = _elementwise_error(loss_eps - x1_pred, reduction.huber_delta)
return _reduce(error, batch_size, sample_count, reduction)
5 changes: 5 additions & 0 deletions tests/test_fpo_advantage_normalisation.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,11 @@ def test_minibatch_loss_does_not_normalise_the_advantage():
running a policy forward. Same approach as tests/test_seeding.py. It has
to match the normalising expression rather than `advantage.std()`, which
also appears there legitimately as a logged metric.

FPO++ does normalise per minibatch, at 375 chunks a minibatch, and
`normalize_advantage_per_minibatch` opts into that through
`_minibatch_advantage`, off by default. tests/test_fpo_plus_plus_
optimizer.py checks both ways, on a running learn step.
"""
source = inspect.getsource(FPOAlgorithm._compute_loss)
assert NORMALISE not in source, (
Expand Down
Loading
Loading