diff --git a/tensorrt_llm/_torch/models/modeling_minimaxm2.py b/tensorrt_llm/_torch/models/modeling_minimaxm2.py index 944f20ec77f3..754962ab5ebc 100644 --- a/tensorrt_llm/_torch/models/modeling_minimaxm2.py +++ b/tensorrt_llm/_torch/models/modeling_minimaxm2.py @@ -19,6 +19,7 @@ from torch import nn from transformers import PretrainedConfig +from tensorrt_llm._ipc_utils import can_access_peer from tensorrt_llm.functional import AllReduceStrategy, PositionEmbeddingType from tensorrt_llm.mapping import Mapping @@ -119,7 +120,13 @@ def forward( # We use all_reduce across all tp gpus to get the rms norm variance sum class MiniMaxRMSNorm(nn.Module): def __init__( - self, *, hidden_size: int, eps: float, mapping: Mapping, dtype: torch.dtype = torch.bfloat16 + self, + *, + hidden_size: int, + eps: float, + mapping: Mapping, + dtype: torch.dtype = torch.bfloat16, + head_dim: Optional[int] = None, ): super().__init__() self.mapping = mapping @@ -128,14 +135,37 @@ def __init__( self.hidden_size = hidden_size self.eps = eps self.dtype = dtype + self.head_dim = head_dim + self.is_p2p_supported = can_access_peer(mapping) self.all_reduce = AllReduce(mapping=self.mapping, strategy=AllReduceStrategy.NCCL) self.minimax_all_reduce_rms = MiniMaxAllReduceRMS(mapping=self.mapping) def load_weights(self, weights: List[Dict]): assert len(weights) == 1 + src = weights[0]["weight"] + # When num_total_heads < tp_size (e.g. 8 KV heads, tp=16), the checkpoint weight + # [num_total_heads * head_dim] is smaller than what TP sharding expects + # [tp_size * local_hidden_size]. Replicate at the head level before sharding, + # consistent with how duplicate_kv_weight handles k_proj/v_proj. + full_size = self.mapping.tp_size * self.hidden_size + if src.shape[0] < full_size and self.head_dim is not None: + assert src.shape[0] % self.head_dim == 0, ( + f"checkpoint weight size {src.shape[0]} is not divisible by head_dim {self.head_dim}" + ) + num_total_heads = src.shape[0] // self.head_dim + assert self.mapping.tp_size % num_total_heads == 0, ( + f"tp_size {self.mapping.tp_size} must be divisible by num_total_heads {num_total_heads} " + f"for head-level weight replication" + ) + reps = self.mapping.tp_size // num_total_heads + src = ( + src.reshape(num_total_heads, self.head_dim) + .repeat_interleave(reps, dim=0) + .reshape(-1) + ) weight = load_weight_shard( - weights[0]["weight"], + src, tensor_parallel_size=self.mapping.tp_size, tensor_parallel_rank=self.mapping.tp_rank, tensor_parallel_mode=TensorParallelMode.COLUMN, @@ -144,6 +174,15 @@ def load_weights(self, weights: List[Dict]): def forward(self, hidden_states: torch.Tensor): hidden_states = hidden_states.contiguous() + if not self.is_p2p_supported: + # Inter-node TP: IPC is unavailable, fall back to NCCL all-reduce of + # partial sum-of-squares followed by local RMS normalization. + hidden_f32 = hidden_states.float() + local_sum_sq = hidden_f32.pow(2).sum(-1, keepdim=True) + total_sum_sq = self.all_reduce(local_sum_sq) + total_hidden = self.hidden_size * self.mapping.tp_size + rms_inv = torch.rsqrt(total_sum_sq / total_hidden + self.eps) + return (hidden_f32 * rms_inv).to(hidden_states.dtype) * self.weight rms_norm_out = self.minimax_all_reduce_rms(hidden_states, self.weight, self.eps) return rms_norm_out @@ -188,12 +227,14 @@ def __init__( eps=config.rms_norm_eps, mapping=self.qkv_proj.mapping, dtype=config.torch_dtype, + head_dim=self.head_dim, ) self.k_norm = MiniMaxRMSNorm( hidden_size=self.kv_size, eps=config.rms_norm_eps, mapping=self.qkv_proj.mapping, dtype=config.torch_dtype, + head_dim=self.head_dim, ) else: self.q_norm = RMSNorm( @@ -209,6 +250,9 @@ def __init__( def apply_qk_norm(self, q, k): if self.qkv_proj.mapping.tp_size > 1: + if not self.q_norm.is_p2p_supported: + # Inter-node TP: fall back to separate per-tensor NCCL-based norm. + return self.q_norm(q), self.k_norm(k) q = q.contiguous() k = k.contiguous() q, k = self.q_norm.minimax_all_reduce_rms.forward_qk(