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.

8 min read • View on GitHub • More from SakanaAI

A wide editorial illustration showing a compact machine with a task switch, a split output head, and two paths for handling valid answers. The scene explains how one extra input bit can change whether the model collapses many valid labels into one narrow choice or preserves a broader distribution.
The whole experiment hangs on a tiny control input. That small change makes label shape, not just optimization, the thing to watch.
Key Takeaways

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.

The experiment is really about target shape. Once you see that, RL stops looking mystical and starts looking like a distribution-preserving update.

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.

MethodTarget shapeFreedom among valid labelsExpected forgettingWhat it preserves
SFT1Single hard labelNoneHighestOnly the chosen label
SFT2Small fixed subsetLimitedLower than SFT1, still artificialA partial answer set
RL / GRPOSamples guided by rewardBroadLowest in the experimentMore of the base model prior
OracleKL-minimal distribution over valid answersPrincipled and broadMatches RL's behavior closelyThe 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.

A close technical illustration of two competing routes from the same base model. One route is a rigid funnel that forces many valid outputs into one fixed target. The other route moves through a small set of candidates and settles near the prior distribution instead of crushing it.
The oracle turns the paper’s strongest intuition into a control experiment. RL looks good because it is closer to the right distributional target.
QuestionSFTOracleRL / GRPO
What target does it learn?A narrow chosen labelThe KL-minimal distribution over valid labelsA reward-shaped distribution close to the oracle
Does it preserve prior preference?PoorlyYesLargely yes
Is the answer space respected?NoYesMostly yes through sampling and reward
Does it explain forgetting?Only partiallyYesYes, 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.

SignalWhat it measuresWhy it matters
KL divergenceHow far the new policy moves from the base modelBest predictor of forgetting in this setup
L2 weight changeHow much parameters movedUseful, but weaker as a forgetting signal
CKNNAHow internal representations driftShows 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.

ProjectPurposeScopeWhy it matters here
rl-razor-mnistReplicate one paper's MNIST experimentVery narrowLets you inspect a single scientific claim end to end
AvalancheGeneral continual-learning libraryBroadUseful if you want many algorithms and datasets
SequoiaBenchmarking platform for continual RL and CLBroadUseful 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.

David Ha, Co-founder, Sakana AI · David Ha's X/Twitter Announcement