Pinned Tweet
Why is RL so sensitive to training-inference mismatch (TIM)? With each training step, the trainer is pulled toward the sampler, resulting in accumulating drift. We exploit this intuition to devise a novel correction method to stabilize RL under TIM, called Score Centering.
18
58
565
69,868
Martin Marek retweeted
We've just added score centering to slime! You can enable sc with --use-score-centering PR: github.com/THUDM/slime/pull/…
RL with LLMs is very unstable when training and sampling policies differ. Standard fixes (matching numerics, importance sampling) work around the problem. We find the root cause of this instability from first principles and propose a way to directly cancel it. Score Centering is competitive and compatible with existing approaches — while simple to implement! 🧵 [1/6]
3
12
108
13,908
Why is RL so sensitive to training-inference mismatch (TIM)? With each training step, the trainer is pulled toward the sampler, resulting in accumulating drift. We exploit this intuition to devise a novel correction method to stabilize RL under TIM, called Score Centering.
18
58
565
69,868
See also @m_ryabinin's thread
RL with LLMs is very unstable when training and sampling policies differ. Standard fixes (matching numerics, importance sampling) work around the problem. We find the root cause of this instability from first principles and propose a way to directly cancel it. Score Centering is competitive and compatible with existing approaches — while simple to implement! 🧵 [1/6]
1
1
12
2,824
Martin Marek retweeted
RL with LLMs is very unstable when training and sampling policies differ. Standard fixes (matching numerics, importance sampling) work around the problem. We find the root cause of this instability from first principles and propose a way to directly cancel it. Score Centering is competitive and compatible with existing approaches — while simple to implement! 🧵 [1/6]
34
118
1,120
232,107
Martin Marek retweeted
New paper: arxiv.org/abs/2605.26097 The main idea is that we can use an LLM to generate its own replay data to prevent forgetting, as long as we have spare capacity. Very overtrained models have to forget to learn new information.
4
26
173
14,148
New paper! "Forgetting in Language Models: Capacity, Optimization, and Self-Generated Replay"
How much does a language model forget when finetuned on new tasks? We show both model size and optimization matter and forgetting can be nearly eliminated with self-generated replay! arxiv.org/abs/2605.26097 w/@mrtnm @dongkyucho @ShikaiQiu @rumichunara @Pavel_Izmailov 1/8
1
2
28
4,292
Interestingly, TPU v7x (Ironwood) is the first generation to 𝘥𝘳𝘰𝘱 4-bit precision, an opposite trend to Nvidia. While Google Cloud docs do not list full TPU specs, they’re actually listed in the Pallas source code: github.com/jax-ml/jax/blob/m…
5
810
🎄 My holiday project – implementing Qwen3 in pure JAX in just 70 LOC – without any model libraries (Flax / Haiku / etc).
2
1
17
1,097
On a TPU v6e-8, Qwen3-8B achieves 30% training MFU and Qwen3-32B achieves over 20,000 tokens / sec sampling throughput (~50% memory bandwidth utilization).
1
1
341
I hope this can be useful for researchers who want to run both training + sampling on a single model replica or implement new models – e.g. qwen3.py and llama3.py differ in just 3 LOC! github.com/martin-marek/jax-…
2
286
Learn how small batch size enables training with just 16 bits / parameter. Happening right now, stand #908
8
379
How should we scale Adam’s hparams with batch size? I had some spare TPUs available so I remastered Figure 4 from our paper on batch size at a higher resolution. Using a 30M language model, we find a constant β₂ half-life (10M tokens) to be optimum across batch sizes.
2
9
526
We also find the optimum LR to increase much slower than sqrt(batch size). For example, as we scale the batch size from 1 to 1024, the square root rule would suggest that the LR should be scaled by a factor of 32, whereas we empirically observe only a factor of 3 scaling.
2
3
333
Getting small batch sizes to work in bfloat16 precision can be challenging. In our recent paper on batch size, we ran all experiments in float32, but memory-constrained settings demand lower precision. Here are two tricks that we used to enable bf16 training at small batch sizes:
10
26
257
22,171
After applying these two tricks to our fine-tuning experiment, Adafactor with bf16 weights still matches the baseline performance of Adam with fp32 weights but crucially its memory footprint is similar to LoRA (with bf16 weights).
1
14
1,151
We updated our codebase with a Colab notebook to finetune Gemma 3 (12B) using a TPU v6e-1 with just 32 GB of memory. We implemented everything from scratch in JAX, including sampling! We also updated our paper to be more explicit about parameter precision. github.com/martin-marek/batc…
3
36
1,132