Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion src/mcore_bridge/model/modules/dsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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]
Expand Down
38 changes: 36 additions & 2 deletions src/mcore_bridge/model/modules/multi_latent_attention.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import os

import torch
import torch.nn.functional as F
from megatron.core import parallel_state, tensor_parallel
Expand All @@ -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(
Expand Down Expand Up @@ -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()
Expand Down
8 changes: 6 additions & 2 deletions src/mcore_bridge/model/register.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading