AnonMLuser's picture
Refresh artefacts and code for the second review release
5c30113 verified
|
Raw
History Blame Contribute Delete
23.6 kB

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