rl-razor-mnist: How a Tiny Task Flag Reveals Why RL Forgets Less
A minimal replication of Sakana AI’s RL’s Razor experiment, where a single extra input bit, a KL-minimal oracle, and group-relative RL turn catastrophic forgetting into something you can inspect, compare, and predict.
- The repo argues that forgetting is often a label-shape problem, not just an optimization problem.
- SFT can overfit by collapsing many valid answers into one arbitrary target.
- The KL-minimal oracle is the control that makes RL’s advantage falsifiable instead of mystical.
- The codebase is less a benchmark suite than a measurement instrument for watching distribution drift.
The extra bit that changes the experiment
The most interesting thing in this repo is not MNIST itself. It is the 785th input dimension, a task indicator that tells the model whether it is solving the old task or the new one.
That sounds tiny because it is tiny. But it changes the whole geometry of the experiment: the model cannot hide behind separate heads or architectural isolation. It has to keep both tasks alive inside one shared policy, so forgetting becomes visible instead of routed around.
Why SFT forgets when the labels are too narrow
`TaskIndicatorDataset` is where the trap is set. Parity can admit multiple acceptable labels, but the supervised variants collapse that space into a narrower target, especially in the `sft1` and `sft2` modes.
That matters because SFT does not just learn the right answer. It learns the exact answer it was given. When several answers are valid, forcing one arbitrary label makes the model move farther than it needs to, which is another way of saying it forgets more.
| Method | Target shape | Freedom among valid labels | Expected forgetting | What it preserves |
|---|---|---|---|---|
| SFT1 | Single hard label | None | Highest | Only the chosen label |
| SFT2 | Small fixed subset | Limited | Lower than SFT1, still artificial | A partial answer set |
| RL / GRPO | Samples guided by reward | Broad | Lowest in the experiment | More of the base model prior |
| Oracle | KL-minimal distribution over valid answers | Principled and broad | Matches RL's behavior closely | The base model's preferences |
# Conceptual shape of the label problem
# Hard labels force one arbitrary answer.
# Oracle targets preserve the base model prior over all valid answers.
valid_labels = {0, 2, 4, 6, 8}
# SFT1 collapses to a single label
sft_target = 0
# Oracle keeps a distribution over all valid labels
oracle_target = {
y: pi0[y] / sum(pi0[k] for k in valid_labels)
for y in valid_labels
}
# The key difference is not correctness.
# It is how much of the original distribution survives.
Our key insight is that the online nature of RL, where the agent continuously interacts with the environment and receives feedback, is crucial for mitigating catastrophic forgetting. This feedback loop ensures that the agent's policy remains consistent with the environment's dynamics, even as it learns new tasks.
The oracle that makes the claim falsifiable
`oracle.py` is the scientific control that gives the paper its bite. It asks a sharper question than “does RL work better?” It asks whether RL is simply landing near the right target distribution for a task with many valid answers.
If supervised fine-tuning is given the KL-minimal distribution over all correct labels, the gap narrows. That is the important result. It says the advantage is not magic in the algorithm. It is a consequence of preserving the base model’s prior when the task does not require a single answer.
| Question | SFT | Oracle | RL / GRPO |
|---|---|---|---|
| What target does it learn? | A narrow chosen label | The KL-minimal distribution over valid labels | A reward-shaped distribution close to the oracle |
| Does it preserve prior preference? | Poorly | Yes | Largely yes |
| Is the answer space respected? | No | Yes | Mostly yes through sampling and reward |
| Does it explain forgetting? | Only partially | Yes | Yes, through a practical training path |
GRPO, stripped down to the part that matters
`grpo.py` is the cleanest engineering choice in the repo. Instead of dragging in a heavyweight actor-critic stack, it uses a group-relative idea: sample a small batch of candidate actions, compare them against each other, and update the policy using relative advantage.
That matters here because the goal is not to build a giant RL system. The goal is to reproduce a specific claim in a tiny setting. GRPO keeps the mechanics legible, which makes it easier to see why online RL can stay closer to the base model while SFT drifts.
# High-level GRPO logic
samples = policy.sample(inputs, group_size=k)
rewards = reward_fn(samples)
advantages = rewards - rewards.mean()
loss = -(log_probs(samples) * advantages).mean()
loss.backward()
The measurements that make forgetting legible
The analysis scripts are where this repo becomes more than a training loop. `plot.py` checks whether KL divergence predicts forgetting better than plain parameter drift, while `analyze_drift_trajectory.py` adds a representation-level view through CKNNA.
That is the deeper contribution. The repository does not just say RL forgets less. It gives you a way to see why: the more the learned distribution stays aligned with the base model, the less the model destabilizes its earlier task behavior.
| Signal | What it measures | Why it matters |
|---|---|---|
| KL divergence | How far the new policy moves from the base model | Best predictor of forgetting in this setup |
| L2 weight change | How much parameters moved | Useful, but weaker as a forgetting signal |
| CKNNA | How internal representations drift | Shows topology-level change, not just output change |
In other words, the repo is a measurement instrument as much as a codebase. It is built to make a scientific relationship visible, not to maximize benchmark glamour.
What this repo is really for
The comparison is not really against a single rival. It is against a category of larger continual-learning systems that are built to cover many scenarios at once.
`rl-razor-mnist` is narrow on purpose. It is closer to a lab instrument than a framework like Avalanche or a broader benchmarking platform like Sequoia. That narrowness is the point: the repo is trying to make one claim auditable, not one platform extensible.
| Project | Purpose | Scope | Why it matters here |
|---|---|---|---|
| rl-razor-mnist | Replicate one paper's MNIST experiment | Very narrow | Lets you inspect a single scientific claim end to end |
| Avalanche | General continual-learning library | Broad | Useful if you want many algorithms and datasets |
| Sequoia | Benchmarking platform for continual RL and CL | Broad | Useful if you want standardized evaluation across tasks |
We present RL's Razor: Why Online Reinforcement Learning Forgets Less. We show that online RL naturally mitigates catastrophic forgetting, a key challenge in continual learning, where agents lose pre-existing skills when learning new ones.