From 8450025c84ee0351fcff42275016c015d5cacefb Mon Sep 17 00:00:00 2001 From: lizhenxing02 Date: Fri, 18 Sep 2026 17:14:07 +0800 Subject: [PATCH 1/2] GLM5 accuracy alignment --- src/mcore_bridge/model/modules/dsa_indexer.py | 9 ++++- .../model/modules/multi_latent_attention.py | 38 ++++++++++++++++++- src/mcore_bridge/model/register.py | 8 +++- 3 files changed, 50 insertions(+), 5 deletions(-) diff --git a/src/mcore_bridge/model/modules/dsa_indexer.py b/src/mcore_bridge/model/modules/dsa_indexer.py index d069e1ce..201674df 100644 --- a/src/mcore_bridge/model/modules/dsa_indexer.py +++ b/src/mcore_bridge/model/modules/dsa_indexer.py @@ -101,12 +101,18 @@ def _apply_rope(self, x: torch.Tensor, rotary_pos_emb: torch.Tensor, cu_seqlens: x_pe, x_nope = torch.split( x, [self.index_head_dim - self.qk_pos_emb_head_dim, self.qk_pos_emb_head_dim], dim=-1) origin_multi_latent_attention = self.config.multi_latent_attention + origin_rotary_interleaved = self.config.rotary_interleaved + use_align = getattr(self.config, 'use_accuracy_compatible', False) squeezed_batch_dim = False if cu_seqlens is not None and x_pe.ndim == 4 and x_pe.size(1) == 1: x_pe = x_pe.squeeze(1) squeezed_batch_dim = True try: - self.config.multi_latent_attention = self.config.dsa_indexer_rotary_interleaved + if use_align: + self.config.rotary_interleaved = self.config.dsa_indexer_rotary_interleaved + self.config.multi_latent_attention = False + else: + self.config.multi_latent_attention = self.config.dsa_indexer_rotary_interleaved x_pe = apply_rotary_pos_emb( x_pe, rotary_pos_emb, @@ -116,6 +122,7 @@ def _apply_rope(self, x: torch.Tensor, rotary_pos_emb: torch.Tensor, cu_seqlens: ) finally: self.config.multi_latent_attention = origin_multi_latent_attention + self.config.rotary_interleaved = origin_rotary_interleaved if squeezed_batch_dim: x_pe = x_pe.unsqueeze(1) # [seqlen, batch, *, index_head_dim] diff --git a/src/mcore_bridge/model/modules/multi_latent_attention.py b/src/mcore_bridge/model/modules/multi_latent_attention.py index 890f4252..121ea3e7 100644 --- a/src/mcore_bridge/model/modules/multi_latent_attention.py +++ b/src/mcore_bridge/model/modules/multi_latent_attention.py @@ -1,3 +1,5 @@ +import os + import torch import torch.nn.functional as F from megatron.core import parallel_state, tensor_parallel @@ -9,6 +11,31 @@ from megatron.core.utils import deprecate_inference_params +class _AlignedHeadExpand(torch.autograd.Function): + + @staticmethod + def forward(ctx, x, num_heads, head_axis): + ctx.num_heads = int(num_heads) + ctx.head_axis = int(head_axis) + shape = [-1] * x.dim() + shape[ctx.head_axis] = ctx.num_heads + return x.expand(*shape) + + @staticmethod + def backward(ctx, grad_output): + axis = ctx.head_axis + grad = grad_output.float() + acc = grad.narrow(axis, 0, 1) + for i in range(1, ctx.num_heads): + acc = acc + grad.narrow(axis, i, 1) + return acc.to(grad_output.dtype), None, None + + +def _align_head_expand_enabled(config) -> bool: + return bool(getattr(config, 'dsa_accuracy_compatible', False)) or \ + os.environ.get('USE_ACCURACY_COMPATIBLE', '0') == '1' + + class MLASelfAttention(McoreMLASelfAttention): def get_query_key_value_tensors( @@ -165,11 +192,18 @@ def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_po query = torch.cat([q_no_pe, q_pos_emb], dim=-1) # key: [num_tokens, n, (qk_head_dim + v_head_dim)] + _align_expand = _align_head_expand_enabled(self.config) if k_pos_emb.ndim == 4: - k_pos_emb = k_pos_emb.expand(-1, -1, self.num_attention_heads_per_partition, -1) + if _align_expand: + k_pos_emb = _AlignedHeadExpand.apply(k_pos_emb, self.num_attention_heads_per_partition, 2) + else: + k_pos_emb = k_pos_emb.expand(-1, -1, self.num_attention_heads_per_partition, -1) else: assert k_pos_emb.ndim == 3 - k_pos_emb = k_pos_emb.expand(-1, self.num_attention_heads_per_partition, -1) + if _align_expand: + k_pos_emb = _AlignedHeadExpand.apply(k_pos_emb, self.num_attention_heads_per_partition, 1) + else: + k_pos_emb = k_pos_emb.expand(-1, self.num_attention_heads_per_partition, -1) key = torch.cat([k_no_pe, k_pos_emb], dim=-1) query = query.contiguous() diff --git a/src/mcore_bridge/model/register.py b/src/mcore_bridge/model/register.py index 7b9db92b..63c755fc 100644 --- a/src/mcore_bridge/model/register.py +++ b/src/mcore_bridge/model/register.py @@ -105,8 +105,12 @@ def _replace_spec_dsa(self, layer_spec): if self.config.qk_layernorm: linear_q_up_proj = backend.column_parallel_linear() # fix megatron-core - dsa_spec.submodules.q_layernorm = backend.layer_norm(for_qk=True) - dsa_spec.submodules.kv_layernorm = backend.layer_norm(for_qk=True) + if getattr(self.config, 'use_accuracy_compatible', False): + dsa_spec.submodules.q_layernorm = WrappedTorchNorm + dsa_spec.submodules.kv_layernorm = WrappedTorchNorm + else: + dsa_spec.submodules.q_layernorm = backend.layer_norm(for_qk=True) + dsa_spec.submodules.kv_layernorm = backend.layer_norm(for_qk=True) dsa_spec.submodules.linear_q_up_proj = linear_q_up_proj dsa_spec.submodules.linear_kv_up_proj = linear_q_up_proj layer_spec.submodules.self_attention = dsa_spec From 5b2927725cb081d64f145876ece56fbb55cc50a8 Mon Sep 17 00:00:00 2001 From: ZhenxingLi Date: Thu, 8 Oct 2026 10:40:09 +0800 Subject: [PATCH 2/2] rerun