From bce0625d7019b1002928a7e63a802659c7714047 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 21:51:30 -0400 Subject: [PATCH] FPO++'s optimizer, gradient clipping, advantage normalisation and Huber CFM loss, as options E37 ran FPO++'s policy loss on square from a behaviour-cloned start and the policy fell from 0.50 to 0.24-0.42 over fifty evaluation episodes, while its randomly initialised critic never fit the returns. FPO++'s square fine-tuning (amazon-far/fpo-control, manipulation_experiments/finetune_online_rl.py) has more than the loss, and this adds the rest, each off by default: critic_learning_rate, adam_eps, weight_decay, actor_adam_beta2 two AdamW groups (FPO++: actor 1e-5 betas (0.9, 0.99), critic 1e-4, eps 1e-5, weight decay 1e-6); left at their defaults, FPO's one Adam. During a critic warmup the actor's group gets no gradient at all, so weight decay cannot move it. max_grad_norm the actor's and the critic's gradients clipped separately (FPO++: 25); their pre-clip maxima are logged. normalize_advantage_per_minibatch as FPO++ does, instead of over the buffer. cfm_loss_huber_delta FPO++'s Huber on the flow-matching error, d^2 within delta and 2 delta |d| - delta^2 beyond (FPO++: 1). FPO's learn step now also reports the critic's explained variance, which DPPO always has and E37 could not read. --- src/plugrl_server/algorithm/fpo/fpo.py | 132 ++++++++- src/plugrl_server/algorithm/fpo/fpo_config.py | 30 +++ src/plugrl_server/algorithm/fpo/utils.py | 26 +- tests/test_fpo_advantage_normalisation.py | 5 + tests/test_fpo_plus_plus_optimizer.py | 255 ++++++++++++++++++ 5 files changed, 437 insertions(+), 11 deletions(-) create mode 100644 tests/test_fpo_plus_plus_optimizer.py diff --git a/src/plugrl_server/algorithm/fpo/fpo.py b/src/plugrl_server/algorithm/fpo/fpo.py index 786ee2b..2e99cf6 100644 --- a/src/plugrl_server/algorithm/fpo/fpo.py +++ b/src/plugrl_server/algorithm/fpo/fpo.py @@ -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 @@ -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: @@ -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) @@ -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) @@ -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. @@ -588,6 +681,7 @@ 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: @@ -595,6 +689,8 @@ def learn_impl(self) -> tuple[int, dict]: 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) @@ -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) @@ -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"])), @@ -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, diff --git a/src/plugrl_server/algorithm/fpo/fpo_config.py b/src/plugrl_server/algorithm/fpo/fpo_config.py index d56b4cc..c671931 100644 --- a/src/plugrl_server/algorithm/fpo/fpo_config.py +++ b/src/plugrl_server/algorithm/fpo/fpo_config.py @@ -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. # @@ -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" + ) diff --git a/src/plugrl_server/algorithm/fpo/utils.py b/src/plugrl_server/algorithm/fpo/utils.py index af0433b..299a1b0 100644 --- a/src/plugrl_server/algorithm/fpo/utils.py +++ b/src/plugrl_server/algorithm/fpo/utils.py @@ -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( @@ -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) diff --git a/tests/test_fpo_advantage_normalisation.py b/tests/test_fpo_advantage_normalisation.py index d3d0acb..5be4be2 100644 --- a/tests/test_fpo_advantage_normalisation.py +++ b/tests/test_fpo_advantage_normalisation.py @@ -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, ( diff --git a/tests/test_fpo_plus_plus_optimizer.py b/tests/test_fpo_plus_plus_optimizer.py new file mode 100644 index 0000000..78da0c4 --- /dev/null +++ b/tests/test_fpo_plus_plus_optimizer.py @@ -0,0 +1,255 @@ +"""FPO++'s optimizer, gradient clipping, advantage normalisation and CFM loss, as options. + +E37 fine-tuned a behaviour-cloned `fpo-policy` on robomimic square with +FPO++'s policy loss and watched it fall from 0.52 to 0.26-0.39, while DPPO +lifted the same clone to 0.78-0.80. Its critic never fit the returns. What +E37 called "FPO++'s settings" was only the policy loss. FPO++'s square +fine-tuning (amazon-far/fpo-control, `manipulation_experiments/ +finetune_online_rl.py`) also has: + + two AdamW optimizers the actor at 1e-5, betas (0.9, 0.99); the critic at + 1e-4; both eps 1e-5, weight decay 1e-6. E37 gave the + random critic the actor's 1e-5. + gradient clipping actor and critic separately, at 25 on square. + advantages normalised per minibatch, not over the buffer. + a Huber CFM error d^2 within delta = 1, 2 delta |d| - delta^2 beyond, + before the mean over dimensions and the sum over + steps. + +Each is an option here, and every option defaults to off, which is FPO's +behaviour until now. +""" + +from __future__ import annotations + +import numpy as np +import pytest +import torch + +from plugrl_server.algorithm.fpo import fpo as fpo_module +from plugrl_server.algorithm.fpo.fpo import FPOAlgorithm +from plugrl_server.algorithm.fpo.fpo_config import FPOAlgoConfig +from plugrl_server.algorithm.fpo.utils import ChunkReduction, compute_cfm_loss +from test_fpo_chunk_loss_per_sample_ratio import _ZeroVelocity, _inputs +from test_fpo_tree_observations import ( + TreeObsFlowPolicy, + _collect_and_learn, + _config, + _tree_env_obs, +) + +FPO_PLUS_PLUS = dict( + learning_rate=1e-5, + critic_learning_rate=1e-4, + adam_eps=1e-5, + weight_decay=1e-6, + actor_adam_beta2=0.99, +) + + +def _algo(**overrides) -> FPOAlgorithm: + torch.manual_seed(0) + np.random.seed(0) + config = _config() + for key, value in overrides.items(): + setattr(config, key, value) + algo = FPOAlgorithm(config=config, policy=TreeObsFlowPolicy()) + algo.init_optimizers() + return algo + + +def _state(module) -> dict[str, torch.Tensor]: + return {k: v.detach().clone() for k, v in module.state_dict().items()} + + +def _moved(before: dict[str, torch.Tensor], module) -> list[str]: + return [k for k, v in module.state_dict().items() if not torch.equal(v, before[k])] + + +class TestTheOptimizer: + def test_unset_it_is_the_one_adam_fpo_always_had(self): + algo = _algo() + + assert type(algo.optimizer) is torch.optim.Adam + (group,) = algo.optimizer.param_groups + assert group["lr"] == algo.config.learning_rate + assert group["eps"] == 1e-8 and group["weight_decay"] == 0 + + def test_fpo_plus_plus_gives_actor_and_critic_their_own_adamw_groups(self): + algo = _algo(**FPO_PLUS_PLUS) + + assert isinstance(algo.optimizer, torch.optim.AdamW) + actor, critic = algo.optimizer.param_groups + assert {id(p) for p in actor["params"]} == { + id(p) for p in algo.policy.actor.parameters() + } + assert {id(p) for p in critic["params"]} == { + id(p) for p in algo.policy.critic.parameters() + } + assert (actor["lr"], actor["betas"]) == (1e-5, (0.9, 0.99)) + assert (critic["lr"], critic["betas"]) == (1e-4, (0.9, 0.999)) + for group in (actor, critic): + assert (group["eps"], group["weight_decay"]) == (1e-5, 1e-6) + + def test_a_critic_rate_alone_keeps_the_actor_at_the_shared_rate(self): + algo = _algo(critic_learning_rate=1e-3) + + actor, critic = algo.optimizer.param_groups + assert (actor["lr"], critic["lr"]) == (algo.config.learning_rate, 1e-3) + + def test_a_warmup_leaves_the_actor_exactly_as_it_was_despite_weight_decay(self): + """AdamW decays every parameter it steps, gradient or not. So the actor's + group is stepped at a rate of zero during the warmup, not only with its + gradients zeroed.""" + algo = _algo(n_critic_warmup_itrs=1, **FPO_PLUS_PLUS) + actor, critic = _state(algo.policy.actor), _state(algo.policy.critic) + + _collect_and_learn(algo, _tree_env_obs) + + assert _moved(actor, algo.policy.actor) == [] + assert _moved(critic, algo.policy.critic) != [] + assert algo.optimizer.param_groups[0]["lr"] == 1e-5, "the rate was not restored" + + def test_after_the_warmup_the_actor_moves(self): + algo = _algo(n_critic_warmup_itrs=1, **FPO_PLUS_PLUS) + _collect_and_learn(algo, _tree_env_obs) + actor = _state(algo.policy.actor) + + _collect_and_learn(algo, _tree_env_obs) + + assert _moved(actor, algo.policy.actor) != [] + + def test_a_checkpoint_round_trips_both_groups(self): + algo = _algo(**FPO_PLUS_PLUS) + _collect_and_learn(algo, _tree_env_obs) + checkpoint = algo.create_checkpoint() + + resumed = _algo(**FPO_PLUS_PLUS) + resumed.load_checkpoint(checkpoint) + + assert [g["lr"] for g in resumed.optimizer.param_groups] == [1e-5, 1e-4] + for a, b in zip(algo.optimizer.state_dict()["state"].values(), + resumed.optimizer.state_dict()["state"].values()): # fmt: skip + assert torch.equal(a["exp_avg"], b["exp_avg"]) + + +class TestGradientClipping: + def test_off_by_default(self, monkeypatch): + calls = [] + monkeypatch.setattr( + fpo_module.torch.nn.utils, + "clip_grad_norm_", + lambda params, max_norm, *a, **k: calls.append(max_norm), + ) + _collect_and_learn(_algo(), _tree_env_obs) + + assert calls == [] + + def test_actor_and_critic_are_clipped_separately(self, monkeypatch): + calls = [] + real = torch.nn.utils.clip_grad_norm_ + + def recording(params, max_norm, *args, **kwargs): + params = list(params) + calls.append((len(params), max_norm)) + return real(params, max_norm, *args, **kwargs) + + monkeypatch.setattr(fpo_module.torch.nn.utils, "clip_grad_norm_", recording) + algo = _algo(max_grad_norm=25.0) + n_actor = len(list(algo.policy.actor.parameters())) + n_critic = len(list(algo.policy.critic.parameters())) + + metrics = _collect_and_learn(algo, _tree_env_obs) + + steps = algo.config.num_updates_per_batch * ( + algo.config.buffer_size // algo.config.batch_size + ) + assert calls == [(n_actor, 25.0), (n_critic, 25.0)] * steps + assert np.isfinite(metrics["fpo"]["actor_grad_norm_max"]) + assert np.isfinite(metrics["fpo"]["critic_grad_norm_max"]) + + def test_a_tight_clip_bounds_the_step(self): + """With a clip far below the gradient, every step is the clip's size.""" + algo = _algo(max_grad_norm=1e-6, learning_rate=1.0) + seen = [] + real_step = algo.optimizer.step + + def step(*args, **kwargs): + params = [p for g in algo.optimizer.param_groups for p in g["params"]] + actor = params[: len(list(algo.policy.actor.parameters()))] + norm = torch.linalg.vector_norm( + torch.stack([p.grad.norm() for p in actor if p.grad is not None]) + ) + seen.append(float(norm)) + return real_step(*args, **kwargs) + + algo.optimizer.step = step + _collect_and_learn(algo, _tree_env_obs) + + assert seen and max(seen) <= 1e-6 * (1 + 1e-4) + + +class TestAdvantageNormalisation: + def test_per_minibatch_leaves_every_minibatch_at_unit_spread(self): + metrics = _collect_and_learn( + _algo(normalize_advantage_per_minibatch=True), _tree_env_obs + ) + + assert metrics["fpo"]["advantages_std"] == pytest.approx(1.0, abs=1e-5) + assert metrics["fpo"]["advantages_mean"] == pytest.approx(0.0, abs=1e-5) + + def test_over_the_buffer_a_minibatch_is_not_at_unit_spread(self): + """The control: without it, the test above would pass either way.""" + metrics = _collect_and_learn(_algo(), _tree_env_obs) + + assert metrics["fpo"]["advantages_std"] != pytest.approx(1.0, abs=1e-5) + + def test_it_needs_normalisation_on(self): + with pytest.raises(ValueError, match="normalize_advantage"): + FPOAlgoConfig( + normalize_advantage=False, normalize_advantage_per_minibatch=True + ) + + +class TestTheHuberError: + def test_unset_the_error_is_squared(self): + assert ChunkReduction().huber_delta is None + + def test_within_delta_squared_beyond_it_linear_and_continuous(self): + action, loss_eps, loss_t = _inputs() + delta = 0.5 + reduction = ChunkReduction( + steps=4, dims=6, sum_over_steps=True, huber_delta=delta + ) + + got = compute_cfm_loss( + _ZeroVelocity(), None, action, output_mode="u", + loss_eps=loss_eps, loss_t=loss_t, reduction=reduction, + ) # fmt: skip + + d = (action.unsqueeze(1) - loss_eps).abs() + error = torch.where(d <= delta, d.pow(2), 2 * delta * d - delta**2) + torch.testing.assert_close(got, error.mean(-1).sum(-1)) + # The two pieces meet at delta. + assert 2 * delta * delta - delta**2 == pytest.approx(delta**2) + + def test_the_algorithm_reads_it_from_its_config(self): + algo = _algo( + output_mode="u", cfm_loss_steps=1, cfm_loss_dims=1, + cfm_loss_sum_over_steps=True, cfm_loss_huber_delta=1.0, + ) # fmt: skip + + assert algo.chunk_reduction == ChunkReduction(1, 1, True, 1.0) + + +class TestTheCriticIsVisible: + """E37 could not say how much of the return its critic explained: FPO + logged no explained variance, where DPPO always has.""" + + @pytest.mark.parametrize("playground", [True, False]) + def test_learn_reports_the_critics_explained_variance(self, playground): + metrics = _collect_and_learn( + _algo(fpo_playground_trick=playground), _tree_env_obs + ) + + assert np.isfinite(metrics["rollout"]["explained_variance"])