ReMDM Experiments
Research and diagnostic scripts for investigating RL fine-tuning of the ReMDM
diffusion planner. These scripts are standalone research code β they
import from src/ (model, sampling, env wrapper, evaluator) but never modify
the core training pipeline. They start from a pretrained DAgger checkpoint.
rl_finetuning/ β RL Fine-Tuning Ablation Suite
Diagnoses why RL fine-tuning of the diffusion model collapses and which interventions fix it. Implements 26 ablations: a baseline plus four groups (A: Regularisation, B: Training Signal, C: Architecture, D: Data Quality), with a comprehensive diagnostic and analysis pipeline.
Training data is on-policy, and returns are per window. Each iteration rolls the current model out under its EMA weights and trains on the resulting windows; a window's return is the reward sum over exactly the actions it trains on, not the episode total broadcast to every window.
Directory structure
rl_finetuning/
βββ run_ablations.py # CLI entry point
βββ ablations/
β βββ losses.py # 16 loss/objective factory functions + LossContext
β βββ optimizers.py # AdamW, LLRD, LoRA, frozen params, PCGrad helpers
β βββ registry.py # AblationSpec dataclass + REGISTRY (26 ablations)
β βββ training.py # run_ablation() loop, MixedReplayBuffer, RewardModel,
β # AblationHistory dataclass
βββ diagnostics/
β βββ gradient.py # Grad alignment, per-layer norms, surgery metrics
β βββ representation.py # KL drift, CKA similarity, activation norms
β βββ timestep.py # t-bin gradient norms, per-t loss decomposition
βββ analysis/
β βββ gdelta.py # Return term g_delta of the decomposition (no training)
β βββ plots.py # 16 matplotlib figure generators
β βββ tables.py # Summary tables as polars DataFrames + LaTeX export
β βββ report.py # diagnosis.md + decision tree figure
β βββ action_distribution.py # Pre/post-RL action distribution analysis
βββ configs/
βββ ablations_default.yaml # Base: all ablation hyperparameters
βββ ablations_fast.yaml # Smoke-test overlay (50 iterations)
βββ ablations_final_minihack_gpu_h200.yaml # H200 overrides only
βββ ablations_final_minihack_gpu_24gb.yaml # RTX 3090 Ti overrides only (reference)
Config layering
Two layers, always: ablations_default.yaml carries every ablation
hyperparameter, a machine config carries only what that machine changes, and
configs never inherit from one another.
Merge order, later wins:
configs/defaults.yaml # architecture, env IDs, token IDs
-> ablations_default.yaml # base: all ablation hyperparameters
-> --ablations-config FILE # machine overrides only
-> ablations_fast.yaml # --fast only, raw overlay
-> CLI flags # --max-iter --batch-size --eval-every --lr --seed
Machine configs carry only the keys they change:
# ablations_final_minihack_gpu_24gb.yaml
batch_size: 4608
cka_batch_size: 128
| Rule | Behaviour |
|---|---|
| Base | ablations_default.yaml, applied automatically; there is no extends key |
--fast |
Applied raw on top, so machine keys it does not set (use_amp, diffusion_steps_collect) survive. It does set batch_size: 128, which therefore overrides the machine value |
| Unknown key | KeyError, not a silent no-op. Valid keys are those in configs/defaults.yaml plus ablations_default.yaml |
| Restated default | Rejected by tests/test_config.py: a key whose value equals what it would inherit is redundant and must be deleted |
Key validation matters here because every ablation reads config through
getattr(cfg, key, fallback), so an unrejected typo such as batch_sze: 512
leaves the real batch_size at its inherited value with no error.
A bare --fast run (no --ablations-config) inherits num_seeds: 3 and AMP
from the base. Pass --num-seeds 1 for a single-seed smoke run.
Usage
List all ablations:
python experiments/rl_finetuning/run_ablations.py --list
Smoke test (2 ablations, fast config):
python experiments/rl_finetuning/run_ablations.py \
--checkpoint path/to/dagger_checkpoint.pth \
--ablations baseline_rl kl_penalty \
--fast
Full suite (all 26 ablations):
python experiments/rl_finetuning/run_ablations.py \
--checkpoint path/to/dagger_checkpoint.pth \
--all \
--num-seeds 3 \
--use-wandb
Full suite on a specific machine:
python experiments/rl_finetuning/run_ablations.py \
--checkpoint path/to/dagger_checkpoint.pth \
--ablations-config experiments/rl_finetuning/configs/ablations_final_minihack_gpu_24gb.yaml \
--all \
--use-wandb
Specific ablations:
python experiments/rl_finetuning/run_ablations.py \
--checkpoint path/to/dagger_checkpoint.pth \
--ablations ewc lora gradient_surgery trust_region_kl
Re-plot from saved results (no training):
python experiments/rl_finetuning/run_ablations.py \
--analyze-only --output-dir outputs/run_20260331_120000 \
--ablations baseline_rl kl_penalty ewc # optional: a subset
Spread across GPUs, then merge:
Run independent subsets on different machines or GPUs, then combine:
# GPU 0
CUDA_VISIBLE_DEVICES=0 python experiments/rl_finetuning/run_ablations.py \
--checkpoint ckpt.pth \
--ablations baseline_rl kl_penalty ewc llrd lora mixed_replay \
--output-dir outputs/gpu0
# GPU 1
CUDA_VISIBLE_DEVICES=1 python experiments/rl_finetuning/run_ablations.py \
--checkpoint ckpt.pth \
--ablations trust_region_kl low_t t_curriculum entropy_bonus \
--output-dir outputs/gpu1
# Merge and regenerate all analysis
python experiments/rl_finetuning/run_ablations.py \
--merge outputs/gpu0/results.json outputs/gpu1/results.json \
--output-dir outputs/combined
--merge accepts any number of results.json files. Where the same ablation appears
in more than one β the same arm run at --seed 0 and --seed 1000, say β the per-seed
scores are concatenated and mean/std recomputed over the union, reported as
<ablation>: <mean> +/- <std> (2 seeds). The merged file is a results.json like any
other, so --analyze-only works on it.
--merge only pools runs from configs that agree on result-affecting keys.
It compares the configs the results files recorded and refuses, naming every
diverging key with both values; a file that records no config is refused too. All
published MiniHack ablation results were produced on the RTX 3090 Ti
(ablations_final_minihack_gpu_24gb.yaml), which is the reference config.
ablations_final_minihack_gpu_h200.yaml is not poolable with it. It diverges on four
keys that change the result, not just the wall-clock:
| Key | GPU-24GB | GPU-H200 | Effect |
|---|---|---|---|
batch_size |
4608 | 512 | ~9x per-update SNR |
episodes_per_iter |
30 | 20 | 15k vs 10k total episodes |
diffusion_steps_collect |
5 | 3 | different collection policy |
eval_episodes |
20 | 10 | noisier score |
Values above are post-layering: a config's own value, or what it inherits from
ablations_default.yaml. Differences in diagnostic cadence (eval_every, cka_every,
cka_batch_size, per_layer_every, repr_drift_every, grad_align_every,
t_analysis_every) are wall-clock only and do not affect poolability.
tests/test_config.py enforces this: every ablations_final_*.yaml must be
declared poolable or not, configs declared poolable must match the reference on
the result-affecting keys, and the recorded GPU-H200 divergence must stay accurate.
Aligning GPU-H200 later fails the test until it is moved to the poolable set. The
key set itself is declared once, in run_ablations._RESULT_AFFECTING, so the
classification the tests check and the refusal --merge performs are the same
policy.
Ablations
| Group | Name | Tests |
|---|---|---|
| Baseline | baseline_rl |
Standard return-weighted ELBO |
| A: Regularisation | kl_penalty |
Soft KL constraint vs pretrained |
ewc |
Elastic Weight Consolidation (Fisher diagonal) | |
llrd |
Layer-wise Learning Rate Decay | |
lora |
Low-Rank Adaptation of attention projections | |
mixed_replay |
Self-replay: the run's own past online windows resampled into each batch | |
trust_region_kl |
Hard KL trust region via quadratic barrier | |
| B: Training Signal | t_curriculum |
Anneal t range high-to-low over training |
entropy_bonus |
Entropy regularisation for action diversity | |
gradient_surgery |
PCGrad: project conflicting RL/BC gradients | |
advantage_clip |
PPO-style advantage clipping [1-eps, 1+eps] | |
normalized_adv |
Std-normalised advantages (per-minibatch) | |
bc_wins |
Uniform ELBO on win windows (no advantage weighting) | |
bc_all |
Uniform ELBO on all rollout windows (no advantage weighting) | |
low_t |
ELBO restricted to low-t (fine-detail) regime | |
| C: Architecture | frozen_backbone |
Train the action head + token embeddings (backbone frozen) |
head_only |
Train only the final action projection | |
attention_only |
Train only the attention projections (Q/K/V/O) | |
ffn_only |
Train only the per-block FFN layers | |
layer_ablation_top1 |
Train only the top-1 transformer block + head | |
layer_ablation_top2 |
Train only the top-2 transformer blocks + head | |
layer_ablation_top3 |
Train only the top-3 transformer blocks + head | |
| D: Data Quality | reward_filtering |
Top-75th-percentile return windows only |
running_stats |
EMA running mean/std for advantage normalisation | |
action_diversity |
Discard degenerate (all-same-action) plans | |
reward_model |
MLP reward model soft-weighting of advantages |
Group C freezes parameters exactly: under an adversarial gradient on every
tensor, each frozen tensor of each arm measures a parameter delta of exactly
0.0 (tests/test_spec_ablations.py, at the production architecture). The
weights the suite evaluates are the EMA shadow, which drifts even when its
parameter does not, because decay * x + (1 - decay) * x is not exactly x in
float32: 45 of 72 tensors, by up to 4.6e-05 over 500 updates at decay 0.999.
A group-C arm's frozen parameters are therefore bit-exact in the trained weights
and approximate to that magnitude in the evaluated ones β far below the
resolution of any reported win rate, and left as it is (ModelEMA in
src/models/denoiser.py).
Output structure
experiments/rl_finetuning/outputs/{run_id}/
βββ results.json # All histories + final scores (machine-readable)
βββ diagnosis.md # Human-readable verdict + evidence + recommendations
βββ checkpoint_{name}.pth # Per-ablation fine-tuned model state dict
βββ figures/
β βββ curves_{name}.png # Per-ablation 2x3 curves (eval, loss, env score,
β β # KL drift, grad alignment, grad norms)
β βββ per_layer_grad_heatmap_{name}.png # Per-layer gradient norms over training
β βββ t_bin_grad_norms_{name}.png # Per-t-bin gradient norms over training
β βββ per_env_collapse_{name}.png # Per-env win rate over eval checkpoints
β βββ final_score_comparison.png # Bar chart of final scores across ablations
β βββ eval_scores_over_training.png # All ablation eval curves overlaid
β βββ score_delta_over_baseline_rl.png # Sorted bar chart of improvement over baseline_rl
β βββ gradient_alignment.png # Gradient cosine similarity over training
β βββ gradient_conflict_map.png # Binary heatmap of gradient conflicts (cos_sim < 0)
β βββ representation_drift.png # KL divergence drift by t-range
β βββ cka_similarity.png # CKA similarity vs pretrained over training
β βββ t_distribution_analysis.png # High/low-t norm ratio + low-high cosine alignment
β βββ t_bin_norms_heatmap.png # Heatmap of per-t-bin gradient norms (final iter)
β βββ win_rate_and_effective_batch_size.png # Online win rate + effective batch size
β βββ group_comparison.png # Boxplot of scores by ablation group
β βββ per_env_delta.png # Heatmap of per-env win rate change (end - start)
β βββ diagnosis_decision_tree.png # Hypothesis evidence bar chart
β βββ action_dist/ # Only with --action-dist (see below)
β βββ action_dist_comparison_{name}.png # Pre/post action frequency bars
β βββ probability_change_{name}.png # Per-action delta and log-ratio
β βββ distribution_metrics_{name}.png # Entropy (nats), effective actions, Gini
β βββ episode_analysis_{name}.png # Return and length histograms
β βββ cumulative_distribution_{name}.png # Cumulative sorted probability
β βββ action_transitions_{name}.png # Pre, post, diff transition matrices
β βββ action_distribution_results_{name}.json # Metrics + statistical tests
β βββ js_divergence_comparison.png # JS divergence across ablations
βββ gdelta/ # --measure-gdelta only
β βββ gdelta_seed{n}.json # Per rollout seed; +/- within is across that seed's draws
β βββ gdelta_aggregate.json # Across seeds; the dispersion the paper's table prints
βββ tables/
βββ results.tex # --emit-tex-macros only: \newcommand per headline number
βββ main_results.{csv,tex} # Per-condition table: score, seed sd, deltas, verdict
βββ significance_test.txt # Max-statistic permutation test + p floor + bootstrap CI
βββ group_summary.{csv,tex} # Group-level summary table
βββ gradient_analysis.{csv,tex} # Grad alignment (mean/final/trend) + KL drift
βββ t_distribution.{csv,tex} # High/low-t ratio, alignment, dominant regime
βββ repr_drift.{csv,tex} # KL drift values at final iteration
βββ per_env.{csv,tex} # Per-environment win rates
βββ forgetting_analysis.{csv,tex} # First collapse iter, min score, recovery
βββ hypothesis_verdict.{csv,tex} # Per-ablation hypothesis verdict + conclusion
βββ gdelta.{csv,tex} # --measure-gdelta only: the decomposition per weight transform
Action distribution analysis is opt-in via --action-dist. It costs roughly
len(id_envs) * --action-dist-episodes * (1 + n_ablations) episodes. The
pretrained baseline is rolled out once and reused across ablations. It reads
the per-ablation checkpoint_{name}.pth files, so it only covers ablations
whose checkpoint was saved.
Measuring the return term (--measure-gdelta)
Splits the return-weighted ELBO gradient into an imitation term and a return term at a single parameter point:
grad L_RW = Abar * ( grad L_BC + g_delta ),
g_delta = (1/B) sum_i delta_i grad l_i, delta_i = A_i/Abar - 1.
It loads the pretrained checkpoint, collects one on-policy batch from it, and
evaluates grad L_BC, g_delta and grad L_RW on that batch at those
parameters under a shared (z_t, t) draw, so the only difference between the
three is the weight vector. It repeats for the four weighting ablations
(baseline_rl, advantage_clip, normalized_adv, bc_wins). No training and
no optimiser step occur; it runs on a laptop CPU in minutes.
Results land in gdelta/ under the run's own output directory, beside
results.json, and the aggregate additionally produces
tables/gdelta.{csv,tex}. With --emit-tex-macros, the analysis pass picks
the aggregate up and emits the measured quantities as \mhGdelta* macros.
Those are kept separate from the \mhCvA* macros, which recover CV_A from
the ESS logged during training: the two are measured on different batches and
do not agree.
Config comes from --results-path, so the weight transforms measured are the
ones that run trained under; without it the standard layering applies.
Reproduction (three rollout seeds, aggregated in one pass):
uv run python experiments/rl_finetuning/run_ablations.py \
--measure-gdelta --gdelta-seeds 0 1 2 \
--checkpoint checkpoints/online/<run>/iterNNN.pth \
--results-path experiments/rl_finetuning/outputs/minihack_ablations/results.json \
--output-dir experiments/rl_finetuning/outputs/minihack_ablations
Seeds run on separate machines are aggregated afterwards with
--gdelta-inputs, the counterpart to --merge:
uv run python experiments/rl_finetuning/run_ablations.py --run-id gdelta \
--gdelta-inputs experiments/rl_finetuning/outputs/gdelta/gdelta_seed{0,1,2}.json
A single seed's ratio_std_draws / cos_std_draws are dispersions over that
seed's eight (z_t, t) draws. The aggregate averages the per-seed means and
reports the standard deviation across seeds, which is what the paper's
table prints.
Reported per weight transform: CV_A, Abar, Abar relative to the baseline's, ESS
as a fraction of the batch, |g_delta| / |grad L_BC|, the cosine between them, and the
same two against a shuffled-delta null β delta permuted across the batch, which
preserves CV_A and destroys the association between a window's weight and its own
gradient. Anything that survives the shuffle is batch heterogeneity, not return signal.
Scope: the objective is the ELBO term alone, excluding the trainer's unweighted
auxiliary goal loss (which would break the identity above for reasons unrelated to the
return), so the goal head carries no gradient here. The collection size is
episodes_per_iter, MiniHack rollouts being sequential rather than vectorised;
measure() takes the override under the sibling's num_envs name. --results-path
layers over configs/defaults.yaml, which supplies the structural keys a recorded
scalar-only config omits.
The sibling remdm-planner-craftax carries the same module with the same
--measure-gdelta flags and output schema; the two drivers otherwise differ on four
flags (craftax has --num-envs; this repo has --device, --wandb-resume-id,
--action-dist-episodes). Its version needs an explicit Orbax sharding to restore a
GPU-written checkpoint on CPU, where torch.load(map_location=...) does not.
results.json schema:
{
"pretrained_score": 0.1234,
"config": {"max_iter": 1000, "batch_size": 512, ...}, // keys lowercase
"merge_provenance": { ... }, // --merge only: inputs + which supplied config
"ablations": {
"kl_penalty": {
"score": 0.1456,
"score_std": 0.008,
"all_scores": [0.14, 0.15, 0.14],
"base_seed": 42, "seeds": [42, 43, 44],
"wall_clock_s": 812.4,
"per_seed_finals": [{...}], "per_seed_final_evals": [{...}],
"history": { ... }
}
}
}
results.json is written incrementally after each ablation completes β a partial file
with N of 26 ablations is fully valid and loadable by --analyze-only or --merge.
CLI reference
| Flag | Description |
|---|---|
--checkpoint PATH |
Pretrained DAgger checkpoint (.pth or wandb: artifact) |
--config PATH |
Main config override (default: configs/defaults.yaml) |
--ablations-config PATH |
Ablations config, layered on ablations_default.yaml (default: ablations_default.yaml) |
--all |
Run all 26 ablations |
--ablations NAME [NAME ...] |
Run specific ablations by name |
--list |
Print registered ablations and exit |
--fast |
Smoke-test mode (max_iter: 50, eval_episodes: 5, batch_size: 128) |
--num-seeds N |
Number of seeds per ablation (overrides num_seeds, default 3) |
--seed N |
Base random seed |
--output-dir DIR |
Output directory (default: auto-timestamped) |
--run-id ID |
Custom run ID for output directory naming |
--analyze-only |
Skip training, regenerate analysis from existing results |
--results-path PATH |
Explicit path to results.json (with --analyze-only or --measure-gdelta) |
--merge JSON [JSON ...] |
Merge multiple results.json files and regenerate analysis |
--measure-gdelta |
Measure the return term at the pretrained checkpoint; no training |
--gdelta-seeds N [N ...] |
Rollout seeds to measure (default 0); the reported +/- is across these |
--gdelta-draws N |
Independent (z_t, t) draws per seed (default 8) |
--gdelta-inputs PATH [PATH ...] |
Aggregate per-seed gdelta JSONs from separate machines |
--use-wandb / --no-use-wandb |
Enable/disable W&B logging (overrides use_wandb, default false) |
--wandb-project NAME |
W&B project (overrides wandb_project, default remdm-planner-minihack-ablations) |
--wandb-entity NAME |
W&B entity (overrides wandb_entity) |
--wandb-resume-id ID |
W&B run ID for curve continuity |
--max-iter N |
Override max training iterations |
--batch-size N |
Override batch size |
--eval-every N |
Override evaluation frequency |
--lr FLOAT |
Override learning rate |
--device DEVICE |
Torch device (default: auto-detect) |
--emit-tex-macros |
Also write tables/results.tex, one \newcommand per headline number |
--action-dist / --no-action-dist |
Pre/post action-distribution analysis (default off here; on in the craftax twin) |
--action-dist-episodes N |
Episodes per environment for that analysis (default 10) |
W&B logging
When --use-wandb is passed, all training dynamics are logged in real time:
| Namespace | Metrics | Frequency |
|---|---|---|
train/ |
loss, learning_rate, grad_norm, effective_batch_size, ablation_local_iter |
Every iteration |
online/ |
win_rate, mean_return |
Every iteration |
speed/ |
iter_time_sec, collect_time_sec, train_step_time_sec, gpu_memory_mb |
Every iteration |
model/ |
param_norm, param_drift_from_init, ema_gate_value |
Every 10 iterations |
eval/ |
id_win_rate, per_env/{env}/win_rate |
Every eval_every |
diag/ |
grad_alignment_cos, repr_drift_kl, cka_similarity, t_grad_norm_low/high |
At diagnostic intervals |
Final scores per ablation are written to wandb.summary.
Diagnostic metrics collected
| Metric | Frequency | What it answers |
|---|---|---|
| Eval score (ID win rate) | every eval_every iters |
Primary performance |
| Training loss | every iteration | Optimisation signal |
| Win rate (online) | every iteration | Online rollout quality |
| Gradient alignment (cos sim) | every grad_align_every |
Is the RL gradient useful? |
| Per-layer gradient norms | every per_layer_every |
Which layers collapse? |
| KL drift from pretrained | every repr_drift_every |
How much has the model changed? |
| CKA similarity | every cka_every |
Representational drift (activation level) |
| t-bin gradient norms | every t_analysis_every |
Is high-t gradient biased? |
| Gradient surgery fraction | every grad_align_every |
PCGrad projected mass |