diff --git a/docs/development.md b/docs/development.md index ed3bcfe021..5f3c002995 100644 --- a/docs/development.md +++ b/docs/development.md @@ -70,6 +70,8 @@ Bringing up a new model family is three steps: implement the modeling code, regi Drop the modeling code under `src/prime_rl/trainer/models//` (HF-compatible config, modeling, and weight conversion). Mirror the layout of an existing family — `glm4_moe/` or `qwen3_moe/` are good starting points. +**Buffer contract.** Models are constructed on the meta device, so any module that registers a buffer (`register_buffer`) must implement `init_buffers_post_meta()` giving it a reasonable value. Modules without this method will cause a runtime failure. + ### Register a Mini Preset Add an entry to [`scripts/mini_moe.py`](https://github.com/PrimeIntellect-ai/prime-rl/blob/main/scripts/mini_moe.py) so the smoke-test workflow can build a ~0.5B test model in your architecture. The preset names the config class, picks small dimensions, and wires up the HF + prime-rl model classes plus a tokenizer source: @@ -135,6 +137,7 @@ Don't expect reward to climb meaningfully in 20 steps on a random model. Before merging a new model, you need to ensure the following: - The model is correctly registered and defines and all the required methods - such as `convert_hf_layer_to_tt` and `convert_tt_layer_to_hf`. +- Every buffer-owning module the new model introduces implements `init_buffers_post_meta()` (see above). - The small smoke test passes. In the PR that adds the new model, you also need to provide a table covering the KL mismatch across 20 steps on `math` environment with `batch_size=64`. All the entries in the table must lower than 0.015. If this is not met, the PR will not be merged (unless reasonable justification is provided). This is to ensure all our models are consistent and their implementations match the implementations in the inference framework. diff --git a/src/prime_rl/trainer/models/afmoe/modeling_afmoe.py b/src/prime_rl/trainer/models/afmoe/modeling_afmoe.py index f81fec99d9..d350ad8a3e 100644 --- a/src/prime_rl/trainer/models/afmoe/modeling_afmoe.py +++ b/src/prime_rl/trainer/models/afmoe/modeling_afmoe.py @@ -494,15 +494,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - __all__ = [ "AfmoeForCausalLM", diff --git a/src/prime_rl/trainer/models/base.py b/src/prime_rl/trainer/models/base.py index 363702b1f2..ae389815c6 100644 --- a/src/prime_rl/trainer/models/base.py +++ b/src/prime_rl/trainer/models/base.py @@ -1,7 +1,37 @@ +from typing import Protocol, runtime_checkable + +import torch.nn as nn from torch import Tensor from transformers.modeling_utils import PreTrainedModel +@runtime_checkable +class PostMetaBufferInitModule(Protocol): + def init_buffers_post_meta(self) -> None: ... + + +def run_init_buffers_post_meta(module: nn.Module, exempt: tuple[type[nn.Module], ...] = ()) -> None: + """Walk every submodule of `module` (module itself excluded) and either call its + `init_buffers_post_meta()` hook or, if it owns buffers but doesn't implement the hook, raise. + + Standalone so it's usable both as the real post-meta-init dispatch (called from + `PreTrainedModelPrimeRL.init_buffers_post_meta`) and directly in tests, against either a real + model or a plain `nn.Module` tree with no HF/PreTrainedModel machinery involved. + """ + for submodule in module.modules(): + if submodule is module or isinstance(submodule, exempt): + continue + if isinstance(submodule, PostMetaBufferInitModule): + submodule.init_buffers_post_meta() + elif next(submodule.buffers(recurse=False), None) is not None: + raise TypeError( + f"{type(submodule).__name__} owns buffers " + f"{[n for n, _ in submodule.named_buffers(recurse=False)]} but doesn't implement " + "init_buffers_post_meta() -- implement it (even a documented no-op) so these " + "buffers don't silently hold undefined values after meta-device materialization." + ) + + class PreTrainedModelPrimeRL(PreTrainedModel): """ Base class for all PrimeRL models that extends HuggingFace PreTrainedModel. @@ -132,17 +162,21 @@ def convert_layer_to_vllm_kernel( """ raise NotImplementedError(f"convert_layer_to_vllm_kernel is not implemented for {cls.__name__}") + _init_buffers_post_meta_exempt: tuple[type[nn.Module], ...] = () + def init_buffers_post_meta(self) -> None: """ Initialize buffers that are not in the state dict after loading with meta device. Some models have buffers (non-trainable tensors) that are not saved in the state dict but need to be properly initialized after loading the model on meta device and then - moving to the actual device. This method should initialize such buffers. + moving to the actual device. Dispatches to each submodule's own `init_buffers_post_meta` + so every layer owns the reinitialization of the buffers it registers; submodules that own + buffers but don't implement the hook cause this to raise (see `run_init_buffers_post_meta`). This is called after loading the model from a checkpoint with meta device. """ - raise NotImplementedError(f"init_buffers_post_meta is not implemented for {self.__class__.__name__}") + run_init_buffers_post_meta(self, exempt=self._init_buffers_post_meta_exempt) __all__ = ["PreTrainedModelPrimeRL"] diff --git a/src/prime_rl/trainer/models/glm4_moe/modeling_glm4_moe.py b/src/prime_rl/trainer/models/glm4_moe/modeling_glm4_moe.py index de87734764..937b9f8c6c 100644 --- a/src/prime_rl/trainer/models/glm4_moe/modeling_glm4_moe.py +++ b/src/prime_rl/trainer/models/glm4_moe/modeling_glm4_moe.py @@ -322,18 +322,5 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - # HF standard transformer model - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - - # TODO: Init TT MoE buffers - # I think .to_empty() on gpu by default fills 0 so we are ok but this might not be guaranteed behavior - __all__ = ["Glm4MoeConfig", "Glm4MoePreTrainedModel", "Glm4MoeModel", "Glm4MoeForCausalLM"] diff --git a/src/prime_rl/trainer/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/prime_rl/trainer/models/glm_moe_dsa/modeling_glm_moe_dsa.py index 547e61ba55..559bbfe626 100644 --- a/src/prime_rl/trainer/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/prime_rl/trainer/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -348,14 +348,5 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - __all__ = ["GlmMoeDsaConfig", "GlmMoeDsaPreTrainedModel", "GlmMoeDsaModel", "GlmMoeDsaForCausalLM"] diff --git a/src/prime_rl/trainer/models/gpt_oss/modeling_gpt_oss.py b/src/prime_rl/trainer/models/gpt_oss/modeling_gpt_oss.py index 5b6dd19d42..b1395f4209 100644 --- a/src/prime_rl/trainer/models/gpt_oss/modeling_gpt_oss.py +++ b/src/prime_rl/trainer/models/gpt_oss/modeling_gpt_oss.py @@ -15,10 +15,13 @@ from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask from transformers.modeling_layers import GradientCheckpointingLayer from transformers.modeling_outputs import MoeModelOutputWithPast +from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS from transformers.models.gpt_oss.modeling_gpt_oss import ( GptOssAttention, GptOssRMSNorm, - GptOssRotaryEmbedding, +) +from transformers.models.gpt_oss.modeling_gpt_oss import ( + GptOssRotaryEmbedding as HFGptOssRotaryEmbedding, ) from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, auto_docstring, can_return_tuple @@ -34,6 +37,18 @@ from prime_rl.trainer.models.layers.moe import GptOssGroupedExperts +class GptOssRotaryEmbedding(HFGptOssRotaryEmbedding): + """HF's GptOssRotaryEmbedding, plus the buffer-reinit hook it doesn't define upstream.""" + + def init_buffers_post_meta(self) -> None: + rope_init_fn = ( + self.compute_default_rope_parameters if self.rope_type == "default" else ROPE_INIT_FUNCTIONS[self.rope_type] + ) + inv_freq, self.attention_scaling = rope_init_fn(self.config, self.inv_freq.device) + self.inv_freq.copy_(inv_freq) + self.original_inv_freq.copy_(inv_freq) + + class GptOssTopKRouter(nn.Module): """Token-choice top-k router matching HF's GptOssTopKRouter parameter naming. @@ -336,22 +351,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS - - rope_init_fn = ( - ROPE_INIT_FUNCTIONS[rotary_emb.rope_type] - if rotary_emb.rope_type != "default" - else rotary_emb.compute_default_rope_parameters - ) - inv_freq, rotary_emb.attention_scaling = rope_init_fn(rotary_emb.config, rotary_emb.inv_freq.device) - rotary_emb.inv_freq.copy_(inv_freq) - if "model.rotary_emb.original_inv_freq" in buffer_names: - rotary_emb.original_inv_freq.copy_(inv_freq) - __all__ = [ "GptOssForCausalLM", diff --git a/src/prime_rl/trainer/models/laguna/modeling_laguna.py b/src/prime_rl/trainer/models/laguna/modeling_laguna.py index 27e5a5d1b7..51dce70d55 100644 --- a/src/prime_rl/trainer/models/laguna/modeling_laguna.py +++ b/src/prime_rl/trainer/models/laguna/modeling_laguna.py @@ -82,6 +82,20 @@ def forward( sin = emb.sin() * attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) + def init_buffers_post_meta(self) -> None: + for layer_type in self.layer_types: + rope_init_fn = self.compute_default_rope_parameters + if self.rope_type[layer_type] != "default": + rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type[layer_type]] + inv_freq, attention_scaling = rope_init_fn( + self.config, + getattr(self, f"{layer_type}_inv_freq").device, + layer_type=layer_type, + ) + getattr(self, f"{layer_type}_inv_freq").copy_(inv_freq) + getattr(self, f"{layer_type}_original_inv_freq").copy_(inv_freq) + setattr(self, f"{layer_type}_attention_scaling", attention_scaling) + def _laguna_attention_config(config: LagunaConfig, num_heads: int) -> AttentionConfig: return AttentionConfig( @@ -384,27 +398,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self) -> None: - rotary_emb = self.model.rotary_emb - for layer_type in rotary_emb.layer_types: - rope_init_fn = rotary_emb.compute_default_rope_parameters - if rotary_emb.rope_type[layer_type] != "default": - rope_init_fn = ROPE_INIT_FUNCTIONS[rotary_emb.rope_type[layer_type]] - inv_freq, attention_scaling = rope_init_fn( - rotary_emb.config, - getattr(rotary_emb, f"{layer_type}_inv_freq").device, - layer_type=layer_type, - ) - getattr(rotary_emb, f"{layer_type}_inv_freq").copy_(inv_freq) - getattr(rotary_emb, f"{layer_type}_original_inv_freq").copy_(inv_freq) - setattr(rotary_emb, f"{layer_type}_attention_scaling", attention_scaling) - - for module in self.modules(): - if isinstance(module, MoE) and module.tokens_per_expert.device.type != "meta": - module.tokens_per_expert.zero_() - if module.expert_bias is not None: - module.expert_bias.zero_() - __all__ = [ "LagunaForCausalLM", diff --git a/src/prime_rl/trainer/models/layers/moe.py b/src/prime_rl/trainer/models/layers/moe.py index 4f44cb5031..f8f6b22153 100644 --- a/src/prime_rl/trainer/models/layers/moe.py +++ b/src/prime_rl/trainer/models/layers/moe.py @@ -1069,6 +1069,12 @@ def init_weights( if self.load_balance_coeff is not None: self.expert_bias = torch.zeros(self.experts.num_experts, dtype=torch.float32) + def init_buffers_post_meta(self) -> None: + self.tokens_per_expert.zero_() + self.routing_confidence_sum.zero_() + if self.expert_bias is not None: + self.expert_bias.zero_() + @torch.compile(dynamic=True) def relu2(x: torch.Tensor) -> torch.Tensor: @@ -1287,6 +1293,9 @@ def forward( def init_weights(self, init_std: float): nn.init.trunc_normal_(self.gate, mean=0.0, std=init_std) + def init_buffers_post_meta(self) -> None: + self.e_score_correction_bias.zero_() + class BCNonGatedFeedForward(nn.Module): """Non-gated feed-forward network used as the shared expert in NemotronH. @@ -1514,3 +1523,9 @@ def init_weights(self, init_std: float, buffer_device: torch.device): self.routing_confidence_sum = torch.tensor(0.0, dtype=torch.float32) if self.load_balance_coeff is not None: self.expert_bias = torch.zeros(self.experts.num_experts, dtype=torch.float32) + + def init_buffers_post_meta(self) -> None: + self.tokens_per_expert.zero_() + self.routing_confidence_sum.zero_() + if self.expert_bias is not None: + self.expert_bias.zero_() diff --git a/src/prime_rl/trainer/models/layers/rotary_emb.py b/src/prime_rl/trainer/models/layers/rotary_emb.py index 8f7bc454fb..5df58d4ddc 100644 --- a/src/prime_rl/trainer/models/layers/rotary_emb.py +++ b/src/prime_rl/trainer/models/layers/rotary_emb.py @@ -57,6 +57,10 @@ def compute_default_rope_parameters(self, config=None, device=None, seq_len=None """Required by transformers 5.0.0 for weight initialization when rope_type is 'default'.""" return _compute_default_rope_parameters(config or self.config, device, seq_len, layer_type) + def init_buffers_post_meta(self) -> None: + inv_freq, self.attention_scaling = self.rope_init_fn(self.config, self.inv_freq.device) + self.inv_freq.copy_(inv_freq) + @torch.no_grad() @dynamic_rope_update # power user: used with advanced RoPE types (e.g. dynamic rope) def forward(self, x, position_ids): diff --git a/src/prime_rl/trainer/models/llama/modeling_llama.py b/src/prime_rl/trainer/models/llama/modeling_llama.py index 10610ac50d..a2bc78c151 100644 --- a/src/prime_rl/trainer/models/llama/modeling_llama.py +++ b/src/prime_rl/trainer/models/llama/modeling_llama.py @@ -312,13 +312,3 @@ def forward( labels[:, slice_indices] if labels is not None else None, temperature=temperature, ) - - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - # HF standard transformer model - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) diff --git a/src/prime_rl/trainer/models/minimax_m2/modeling_minimax_m2.py b/src/prime_rl/trainer/models/minimax_m2/modeling_minimax_m2.py index febb814abe..e4ecf23b6c 100644 --- a/src/prime_rl/trainer/models/minimax_m2/modeling_minimax_m2.py +++ b/src/prime_rl/trainer/models/minimax_m2/modeling_minimax_m2.py @@ -262,15 +262,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - __all__ = [ "MiniMaxM2ForCausalLM", diff --git a/src/prime_rl/trainer/models/nemotron_h/modeling_nemotron_h.py b/src/prime_rl/trainer/models/nemotron_h/modeling_nemotron_h.py index 3f48dfc541..1a371e3438 100644 --- a/src/prime_rl/trainer/models/nemotron_h/modeling_nemotron_h.py +++ b/src/prime_rl/trainer/models/nemotron_h/modeling_nemotron_h.py @@ -514,9 +514,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - pass - __all__ = [ "NemotronHForCausalLM", diff --git a/src/prime_rl/trainer/models/qwen3/modeling_qwen3.py b/src/prime_rl/trainer/models/qwen3/modeling_qwen3.py index a568bc89bd..30a9f83b54 100644 --- a/src/prime_rl/trainer/models/qwen3/modeling_qwen3.py +++ b/src/prime_rl/trainer/models/qwen3/modeling_qwen3.py @@ -268,14 +268,5 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - __all__ = ["Qwen3ForCausalLM", "Qwen3Model", "Qwen3PreTrainedModel"] diff --git a/src/prime_rl/trainer/models/qwen3_5/modeling_qwen3_5.py b/src/prime_rl/trainer/models/qwen3_5/modeling_qwen3_5.py index ee933bc7e1..1fe944d0bf 100644 --- a/src/prime_rl/trainer/models/qwen3_5/modeling_qwen3_5.py +++ b/src/prime_rl/trainer/models/qwen3_5/modeling_qwen3_5.py @@ -12,7 +12,7 @@ from transformers.models.qwen3_5.modeling_qwen3_5 import ( Qwen3_5PreTrainedModel as HFQwen3_5PreTrainedModel, ) -from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5VisionModel +from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5VisionModel, Qwen3_5VisionRotaryEmbedding from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs @@ -32,6 +32,19 @@ from prime_rl.utils.sequence import get_cu_seqlens_from_seq_lens +def _init_vision_rope_buffers_post_meta(self: Qwen3_5VisionRotaryEmbedding) -> None: + inv_freq = 1.0 / ( + self.theta ** (torch.arange(0, self.dim, 2, dtype=torch.float32, device=self.inv_freq.device) / self.dim) + ) + self.inv_freq.copy_(inv_freq) + + +# Qwen3_5VisionRotaryEmbedding is upstream transformers code; Qwen3_5VisionModel.__init__ +# hardcodes its construction, so we can't subclass-and-inject like elsewhere. Attach the +# buffer-init hook to the class directly instead. +Qwen3_5VisionRotaryEmbedding.init_buffers_post_meta = _init_vision_rope_buffers_post_meta + + class Qwen3_5GatedFlashAttention(Qwen3_5MoeGatedFlashAttention): pass @@ -452,26 +465,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - if self._is_vlm: - lm_rope = self.model.language_model.rotary_emb - else: - lm_rope = self.model.rotary_emb - - if hasattr(lm_rope, "rope_init_fn"): - inv_freq, lm_rope.attention_scaling = lm_rope.rope_init_fn(lm_rope.config, lm_rope.inv_freq.device) - lm_rope.inv_freq.copy_(inv_freq) - - if self._is_vlm: - vis_rope = self.model.visual.rotary_pos_emb - if hasattr(vis_rope, "inv_freq"): - dim = vis_rope.inv_freq.shape[0] - inv_freq = 1.0 / ( - 10000.0 - ** (torch.arange(0, dim * 2, 2, dtype=torch.float32, device=vis_rope.inv_freq.device) / (dim * 2)) - ) - vis_rope.inv_freq.copy_(inv_freq) - __all__ = [ "Qwen3_5ForCausalLM", diff --git a/src/prime_rl/trainer/models/qwen3_5_moe/modeling_qwen3_5_moe.py b/src/prime_rl/trainer/models/qwen3_5_moe/modeling_qwen3_5_moe.py index 6d703241a0..4b2a03921b 100644 --- a/src/prime_rl/trainer/models/qwen3_5_moe/modeling_qwen3_5_moe.py +++ b/src/prime_rl/trainer/models/qwen3_5_moe/modeling_qwen3_5_moe.py @@ -15,7 +15,10 @@ from transformers.modeling_layers import GradientCheckpointingLayer from transformers.modeling_outputs import MoeModelOutputWithPast from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update -from transformers.models.qwen3_5_moe.modeling_qwen3_5_moe import Qwen3_5MoeVisionModel +from transformers.models.qwen3_5_moe.modeling_qwen3_5_moe import ( + Qwen3_5MoeVisionModel, + Qwen3_5MoeVisionRotaryEmbedding, +) from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, logging @@ -35,6 +38,19 @@ from .converting_qwen3_5_moe import conversion_chain from .mrope import build_qwen3_5_mrope_position_ids + +def _init_vision_rope_buffers_post_meta(self: Qwen3_5MoeVisionRotaryEmbedding) -> None: + inv_freq = 1.0 / ( + self.theta ** (torch.arange(0, self.dim, 2, dtype=torch.float32, device=self.inv_freq.device) / self.dim) + ) + self.inv_freq.copy_(inv_freq) + + +# Qwen3_5MoeVisionRotaryEmbedding is upstream transformers code; Qwen3_5MoeVisionModel.__init__ +# hardcodes its construction, so we can't subclass-and-inject like elsewhere. Attach the +# buffer-init hook to the class directly instead. +Qwen3_5MoeVisionRotaryEmbedding.init_buffers_post_meta = _init_vision_rope_buffers_post_meta + logger = logging.get_logger(__name__) @@ -601,6 +617,10 @@ def _apply_interleaved_mrope(self, freqs: torch.Tensor) -> torch.Tensor: freqs_t[..., idx] = freqs[dim, ..., idx] return freqs_t + def init_buffers_post_meta(self) -> None: + inv_freq, self.attention_scaling = self.rope_init_fn(self.config, self.inv_freq.device) + self.inv_freq.copy_(inv_freq) + def _create_rotary_emb(config: Qwen3_5MoeConfig) -> Qwen3_5MoeRotaryEmbedding: return Qwen3_5MoeRotaryEmbedding(config) @@ -994,30 +1014,6 @@ def forward( temperature=temperature, ) - # ------------------------------------------------------------------ - # Buffer init after meta-device loading - # ------------------------------------------------------------------ - - def init_buffers_post_meta(self): - if self._is_vlm: - lm_rope = self.model.language_model.rotary_emb - else: - lm_rope = self.model.rotary_emb - - if hasattr(lm_rope, "rope_init_fn"): - inv_freq, lm_rope.attention_scaling = lm_rope.rope_init_fn(lm_rope.config, lm_rope.inv_freq.device) - lm_rope.inv_freq.copy_(inv_freq) - - if self._is_vlm: - vis_rope = self.model.visual.rotary_pos_emb - if hasattr(vis_rope, "inv_freq"): - dim = vis_rope.inv_freq.shape[0] - inv_freq = 1.0 / ( - 10000.0 - ** (torch.arange(0, dim * 2, 2, dtype=torch.float32, device=vis_rope.inv_freq.device) / (dim * 2)) - ) - vis_rope.inv_freq.copy_(inv_freq) - __all__ = [ "Qwen3_5MoeForCausalLM", diff --git a/src/prime_rl/trainer/models/qwen3_moe/modeling_qwen3_moe.py b/src/prime_rl/trainer/models/qwen3_moe/modeling_qwen3_moe.py index 15737c2817..48be7bde85 100644 --- a/src/prime_rl/trainer/models/qwen3_moe/modeling_qwen3_moe.py +++ b/src/prime_rl/trainer/models/qwen3_moe/modeling_qwen3_moe.py @@ -333,19 +333,6 @@ def forward( temperature=temperature, ) - def init_buffers_post_meta(self): - buffer_names = [name for name, _ in self.named_buffers()] - # HF standard transformer model - if "model.rotary_emb.inv_freq" in buffer_names: - rotary_emb = self.model.rotary_emb - inv_freq, rotary_emb.attention_scaling = rotary_emb.rope_init_fn( - rotary_emb.config, rotary_emb.inv_freq.device - ) - rotary_emb.inv_freq.copy_(inv_freq) - - # TODO: Init TT MoE buffers - # I think .to_empty() on gpu by default fills 0 so we are ok but this might not be guaranteed behavior - __all__ = [ "Qwen3MoeForCausalLM", diff --git a/tests/unit/train/models/test_glm4_moe.py b/tests/unit/train/models/test_glm4_moe.py index 2ee4bee01b..7779f00fd9 100644 --- a/tests/unit/train/models/test_glm4_moe.py +++ b/tests/unit/train/models/test_glm4_moe.py @@ -124,5 +124,35 @@ def test_glm4_moe() -> None: assert torch.allclose(grad_diff, torch.zeros_like(grad_diff), atol=1000), f"Max grad diff: {grad_diff.abs().max()}" +def test_glm4_moe_init_buffers_post_meta(): + config = Glm4MoeConfig( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + max_position_embeddings=128, + moe_intermediate_size=32, + norm_topk_prob=True, + num_attention_heads=4, + num_key_value_heads=2, + n_routed_experts=4, + num_experts_per_tok=2, + n_shared_experts=1, + num_hidden_layers=2, + rope_theta=1000000.0, + first_k_dense_replace=1, + partial_rotary_factor=0.5, + use_grouped_mm=False, + vocab_size=64, + ) + with torch.device("meta"): + model = PrimeRLGlm4MoeForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" + + if __name__ == "__main__": test_glm4_moe_mlp_only() diff --git a/tests/unit/train/models/test_glm_moe_dsa.py b/tests/unit/train/models/test_glm_moe_dsa.py new file mode 100644 index 0000000000..e64bdfddfa --- /dev/null +++ b/tests/unit/train/models/test_glm_moe_dsa.py @@ -0,0 +1,35 @@ +import pytest +import torch + +from prime_rl.trainer.models.glm_moe_dsa import GlmMoeDsaConfig, GlmMoeDsaForCausalLM + +pytestmark = [pytest.mark.gpu] + + +def test_glm_moe_dsa_init_buffers_post_meta(): + config = GlmMoeDsaConfig( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + max_position_embeddings=128, + moe_intermediate_size=32, + norm_topk_prob=True, + num_attention_heads=4, + num_key_value_heads=2, + n_routed_experts=4, + num_experts_per_tok=2, + n_shared_experts=1, + num_hidden_layers=2, + rope_theta=1000000.0, + first_k_dense_replace=1, + use_grouped_mm=False, + vocab_size=64, + ) + with torch.device("meta"): + model = GlmMoeDsaForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_gpt_oss.py b/tests/unit/train/models/test_gpt_oss.py new file mode 100644 index 0000000000..5d30e174c4 --- /dev/null +++ b/tests/unit/train/models/test_gpt_oss.py @@ -0,0 +1,32 @@ +import pytest +import torch +from transformers.models.gpt_oss.configuration_gpt_oss import GptOssConfig + +from prime_rl.trainer.models.gpt_oss import GptOssForCausalLM + +pytestmark = [pytest.mark.gpu] + + +def test_gpt_oss_init_buffers_post_meta(): + config = GptOssConfig( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + num_local_experts=4, + num_experts_per_tok=2, + vocab_size=64, + sliding_window=16, + layer_types=["sliding_attention", "full_attention"], + ) + with torch.device("meta"): + model = GptOssForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_laguna.py b/tests/unit/train/models/test_laguna.py new file mode 100644 index 0000000000..4f1845b569 --- /dev/null +++ b/tests/unit/train/models/test_laguna.py @@ -0,0 +1,33 @@ +import pytest +import torch + +from prime_rl.trainer.models.laguna import LagunaConfig, LagunaForCausalLM + +pytestmark = [pytest.mark.gpu] + + +def test_laguna_init_buffers_post_meta(): + config = LagunaConfig( + pad_token_id=0, + vocab_size=64, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + moe_intermediate_size=16, + shared_expert_intermediate_size=16, + num_experts_per_tok=2, + num_experts=4, + sliding_window=None, + layer_types=["full_attention", "full_attention"], + ) + with torch.device("meta"): + model = LagunaForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_llama.py b/tests/unit/train/models/test_llama.py index b1c94fe971..30c3a55b33 100644 --- a/tests/unit/train/models/test_llama.py +++ b/tests/unit/train/models/test_llama.py @@ -127,5 +127,25 @@ def test_llama(): assert torch.allclose(grad_diff, torch.zeros_like(grad_diff), atol=1000), f"Max grad diff: {grad_diff.abs().max()}" +def test_llama_init_buffers_post_meta(): + config = LlamaConfig( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + num_attention_heads=4, + num_key_value_heads=2, + num_hidden_layers=2, + vocab_size=64, + ) + with torch.device("meta"): + model = PrimeRLLlamaForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" + + if __name__ == "__main__": test_llama_mlp_only() diff --git a/tests/unit/train/models/test_minimax_m2.py b/tests/unit/train/models/test_minimax_m2.py new file mode 100644 index 0000000000..bedf5081e7 --- /dev/null +++ b/tests/unit/train/models/test_minimax_m2.py @@ -0,0 +1,31 @@ +import pytest +import torch + +from prime_rl.trainer.models.minimax_m2 import MiniMaxM2Config, MiniMaxM2ForCausalLM + +pytestmark = [pytest.mark.gpu] + + +def test_minimax_m2_init_buffers_post_meta(): + config = MiniMaxM2Config( + pad_token_id=0, + vocab_size=64, + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + max_position_embeddings=128, + num_local_experts=4, + num_experts_per_tok=2, + use_grouped_mm=False, + ) + with torch.device("meta"): + model = MiniMaxM2ForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_nemotron_h.py b/tests/unit/train/models/test_nemotron_h.py index 9ba3b0e641..b661ba513a 100644 --- a/tests/unit/train/models/test_nemotron_h.py +++ b/tests/unit/train/models/test_nemotron_h.py @@ -268,3 +268,19 @@ def test_nemotron_h_no_latent_projection(): for name, p in model.named_parameters(): if "experts.w1" in name and p.numel() > 0: assert p.grad is not None and p.grad.norm().item() > 0, f"Zero grad for {name}" + + +def test_nemotron_h_init_buffers_post_meta(): + config = NemotronHConfig( + **_BASE, + layers_block_type=["mamba", "moe", "attention", "moe"], + use_grouped_mm=False, + ) + with torch.device("meta"): + model = NemotronHForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_qwen3.py b/tests/unit/train/models/test_qwen3.py new file mode 100644 index 0000000000..70ed3cf9f1 --- /dev/null +++ b/tests/unit/train/models/test_qwen3.py @@ -0,0 +1,28 @@ +import pytest +import torch +from transformers.models.qwen3.configuration_qwen3 import Qwen3Config + +from prime_rl.trainer.models.qwen3 import Qwen3ForCausalLM + +pytestmark = [pytest.mark.gpu] + + +def test_qwen3_init_buffers_post_meta(): + config = Qwen3Config( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + num_hidden_layers=2, + vocab_size=64, + ) + with torch.device("meta"): + model = Qwen3ForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" diff --git a/tests/unit/train/models/test_qwen3_5.py b/tests/unit/train/models/test_qwen3_5.py index 03fa0ef0d2..dce05bab0e 100644 --- a/tests/unit/train/models/test_qwen3_5.py +++ b/tests/unit/train/models/test_qwen3_5.py @@ -93,6 +93,19 @@ def test_qwen3_5_dense_matches_hf_state_keys_on_meta(): assert tensor.shape == hf_model.state_dict()[name].shape, name +@pytest.mark.gpu +def test_qwen3_5_init_buffers_post_meta(): + config = _tiny_text_config() + with torch.device("meta"): + model = Qwen3_5ForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" + + @pytest.mark.parametrize("attn_impl", ["flash_attention_3", "kernels-community/vllm-flash-attn3"]) def test_qwen3_5_full_attention_uses_custom_class(attn_impl: str): config = _tiny_text_config(attn_impl=attn_impl) diff --git a/tests/unit/train/models/test_qwen3_5_moe.py b/tests/unit/train/models/test_qwen3_5_moe.py index c9604194cd..fa922897b4 100644 --- a/tests/unit/train/models/test_qwen3_5_moe.py +++ b/tests/unit/train/models/test_qwen3_5_moe.py @@ -186,5 +186,36 @@ def test_qwen3_5_moe_context_parallel_setup_hook(): assert linear_layer.linear_attn.cp_world_size == 2 +def test_qwen3_5_moe_init_buffers_post_meta(): + config = Qwen3_5MoeConfig( + vocab_size=256, + hidden_size=256, + num_hidden_layers=4, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=64, + moe_intermediate_size=128, + shared_expert_intermediate_size=128, + num_experts=8, + num_experts_per_tok=2, + max_position_embeddings=512, + rms_norm_eps=1e-6, + linear_conv_kernel_dim=4, + linear_key_head_dim=32, + linear_value_head_dim=32, + linear_num_key_heads=4, + linear_num_value_heads=8, + use_grouped_mm=False, + ) + with torch.device("meta"): + model = PrimeRLQwen3_5MoeForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" + + if __name__ == "__main__": test_qwen3_5_moe() diff --git a/tests/unit/train/models/test_qwen3_moe.py b/tests/unit/train/models/test_qwen3_moe.py index 38e602ab88..598ee823a0 100644 --- a/tests/unit/train/models/test_qwen3_moe.py +++ b/tests/unit/train/models/test_qwen3_moe.py @@ -161,5 +161,33 @@ def test_qwen3_moe_router_replay(): assert prime_model.model.embed_tokens.weight.grad is not None +def test_qwen3_moe_init_buffers_post_meta(): + config = Qwen3MoeConfig( + pad_token_id=0, + head_dim=16, + hidden_size=32, + max_position_embeddings=128, + moe_intermediate_size=32, + norm_topk_prob=True, + num_attention_heads=4, + num_experts=4, + num_experts_per_tok=2, + num_hidden_layers=2, + rope_theta=1000000.0, + use_qk_norm=True, + mlp_only_layers=[1], + use_grouped_mm=False, + vocab_size=64, + ) + with torch.device("meta"): + model = PrimeRLQwen3MoeForCausalLM(config) + model.to_empty(device="cuda") + + model.init_buffers_post_meta() + + for name, buffer in model.named_buffers(): + assert torch.isfinite(buffer).all(), f"buffer {name} is not finite after init_buffers_post_meta" + + if __name__ == "__main__": test_qwen3_moe_mlp_only() diff --git a/tests/unit/train/models/test_state_loading.py b/tests/unit/train/models/test_state_loading.py index 1d4237c3ea..cc6180b82d 100644 --- a/tests/unit/train/models/test_state_loading.py +++ b/tests/unit/train/models/test_state_loading.py @@ -3,8 +3,9 @@ import pytest import torch -from prime_rl.configs.trainer import ModelConfig +from prime_rl.configs.trainer import DebugModelConfig, ModelConfig from prime_rl.trainer.model import load_dcp_from_hf +from prime_rl.trainer.models.glm4_moe import Glm4MoeConfig, Glm4MoeForCausalLM from prime_rl.trainer.models.laguna.configuration_laguna import LagunaConfig from prime_rl.trainer.models.laguna.modeling_laguna import LagunaForCausalLM @@ -48,3 +49,54 @@ def fake_dcp_load(state_dict, storage_reader=None): expert_bias = model.model.layers[1].mlp.expert_bias torch.testing.assert_close(expert_bias.cpu(), expected.to(expert_bias.dtype)) + + +@pytest.fixture +def glm4_moe_model() -> Glm4MoeForCausalLM: + config = Glm4MoeConfig( + pad_token_id=0, + hidden_size=32, + intermediate_size=64, + max_position_embeddings=128, + moe_intermediate_size=32, + norm_topk_prob=True, + num_attention_heads=4, + num_key_value_heads=2, + n_routed_experts=4, + num_experts_per_tok=2, + n_shared_experts=1, + num_hidden_layers=2, + rope_theta=1000000.0, + first_k_dense_replace=1, + partial_rotary_factor=0.5, + use_grouped_mm=False, + vocab_size=64, + ) + with torch.device("meta"): + return Glm4MoeForCausalLM(config) + + +def test_load_dcp_from_hf_random_init_zeros_expert_bias(glm4_moe_model, tmp_path, monkeypatch): + """Under debug.random_init=True, dcp_load is skipped entirely, so init_buffers_post_meta is the + only thing that can overwrite to_empty()'s undefined expert_bias contents. Regression test for + glm4_moe, which (unlike laguna) never zeroed expert_bias before this fix.""" + monkeypatch.setattr("torch.distributed.barrier", lambda *args, **kwargs: None) + + real_to_empty = glm4_moe_model.to_empty + + def poisoning_to_empty(*, device): + real_to_empty(device=device) + with torch.no_grad(): + for layer in glm4_moe_model.model.layers: + if getattr(layer.mlp, "expert_bias", None) is not None: + layer.mlp.expert_bias.fill_(1.0) + return glm4_moe_model + + monkeypatch.setattr(glm4_moe_model, "to_empty", poisoning_to_empty) + + config = ModelConfig(name=str(tmp_path), debug=DebugModelConfig(random_init=True)) + load_dcp_from_hf(glm4_moe_model, config, parallel_dims=MagicMock()) + + for layer in glm4_moe_model.model.layers: + if getattr(layer.mlp, "expert_bias", None) is not None: + torch.testing.assert_close(layer.mlp.expert_bias.cpu(), torch.zeros_like(layer.mlp.expert_bias.cpu()))