[BugFix] fix the bug of minimax_m2.5 model that causes repeated outputs when using tp16 (#20967)

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
kingkingleeljj
2026-04-10 22:21:19 +08:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 1ff51555f2
commit 84194c25c1
+34 -10
View File
@@ -253,27 +253,47 @@ def rms_apply_serial(
class MiniMaxM2RMSNormTP(nn.Module):
"""RMSNorm with Tensor Parallel support for QK normalization."""
def __init__(self, hidden_size: int, eps: float = 1e-6) -> None:
def __init__(self, hidden_size: int, num_heads: int, eps: float = 1e-6) -> None:
super().__init__()
self.attn_tp_size = get_attention_tp_size()
self.attn_tp_rank = get_attention_tp_rank()
# Align with QKVParallelLinear pattern
if self.attn_tp_size >= num_heads:
assert (
self.attn_tp_size % num_heads == 0
), f"attn_tp_size ({self.attn_tp_size}) must be divisible by num_heads ({num_heads})"
self.num_heads = 1
self.num_head_replicas = self.attn_tp_size // num_heads
else:
assert (
num_heads % self.attn_tp_size == 0
), f"num_heads ({num_heads}) must be divisible by attn_tp_size ({self.attn_tp_size})"
self.num_heads = num_heads // self.attn_tp_size
self.num_head_replicas = 1
self.head_dim = hidden_size // num_heads
# Weight parameter is sharded across TP ranks
self.weight = nn.Parameter(torch.ones(int(hidden_size / self.attn_tp_size)))
self.weight = nn.Parameter(torch.ones(self.num_heads * self.head_dim))
self.weight.weight_loader = self.weight_loader
self.variance_epsilon = eps
@staticmethod
def weight_loader(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
) -> None:
"""Custom weight loader that handles TP sharding."""
attn_tp_size = get_attention_tp_size()
attn_tp_rank = get_attention_tp_rank()
shard_size = loaded_weight.shape[0] // attn_tp_size
shard = slice(attn_tp_rank * shard_size, (attn_tp_rank + 1) * shard_size)
shard_id = self.attn_tp_rank // self.num_head_replicas
shard_size = param.data.shape[0]
shard_end = (shard_id + 1) * shard_size
assert shard_end <= loaded_weight.shape[0], (
f"Weight shard out of bounds: shard [{shard_id * shard_size}:{shard_end}] "
f"exceeds loaded_weight size {loaded_weight.shape[0]} "
f"(attn_tp_rank={self.attn_tp_rank}, num_head_replicas={self.num_head_replicas})"
)
shard = slice(shard_id * shard_size, shard_end)
param.data.copy_(loaded_weight[shard])
@torch.compile(dynamic=True, backend=get_compiler_backend())
@@ -641,10 +661,14 @@ class MiniMaxM2Attention(nn.Module):
# Use RMSNormTP for proper tensor parallel support
# Use total dimensions (before TP sharding) for correct normalization
self.q_norm = MiniMaxM2RMSNormTP(
self.total_num_heads * self.head_dim, eps=config.rms_norm_eps
self.total_num_heads * self.head_dim,
num_heads=self.total_num_heads,
eps=config.rms_norm_eps,
)
self.k_norm = MiniMaxM2RMSNormTP(
self.total_num_kv_heads * self.head_dim, eps=config.rms_norm_eps
self.total_num_kv_heads * self.head_dim,
num_heads=self.total_num_kv_heads,
eps=config.rms_norm_eps,
)
else:
raise ValueError(f"Unsupported qk_norm_type: {self.qk_norm_type}")