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
348 changes: 341 additions & 7 deletions megatron/core/fusions/fused_mla_yarn_rope_apply.py

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -41,13 +41,9 @@
from megatron.core.utils import deprecate_inference_params, get_pg_size, is_te_min_version

try:
from megatron.core.fusions.fused_mla_yarn_rope_apply import (
fused_apply_mla_rope_for_kv,
fused_apply_mla_rope_for_q,
)
from megatron.core.fusions.fused_mla_yarn_rope_apply import fused_mla_rope_concat
except ImportError:
fused_apply_mla_rope_for_kv = None
fused_apply_mla_rope_for_q = None
fused_mla_rope_concat = None

if HAVE_TE:
from megatron.core.extensions.transformer_engine import (
Expand Down Expand Up @@ -465,21 +461,19 @@ def get_query_key_value_tensors(
rotary_pos_cos = None
rotary_pos_sin = None
packed_seq = packed_seq_params is not None and packed_seq_params.qkv_format == 'thd'
if self.config.rope_type == "rope":
if self.config.apply_rope_fusion:
rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb.get_cached_cos_sin(
rotary_seq_len, dtype=hidden_states.dtype, packed_seq=packed_seq
)
rotary_pos_emb = None
assert inference_context is None, "Inference with MLA RoPE fusion is not supported"
assert (
fused_mla_rope_concat is not None
), "Fused MLA RoPE apply is not imported successfully"
elif self.config.rope_type == "rope":
rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len, packed_seq=packed_seq)
else:
if self.config.apply_rope_fusion:
rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb.get_cached_cos_sin(
rotary_seq_len, dtype=hidden_states.dtype, packed_seq=packed_seq
)
rotary_pos_emb = None
assert inference_context is None, "Inference with MLA RoPE fusion is not supported"
assert (
fused_apply_mla_rope_for_q is not None
and fused_apply_mla_rope_for_kv is not None
), "Fused MLA RoPE apply is not imported successfully"
else:
rotary_pos_emb, mscale = self.rotary_pos_emb(rotary_seq_len, packed_seq=packed_seq)
rotary_pos_emb, mscale = self.rotary_pos_emb(rotary_seq_len, packed_seq=packed_seq)

if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd':
if packed_seq_params.cu_seqlens_q_padded is not None:
Expand Down Expand Up @@ -609,34 +603,26 @@ def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_po
# Absorb k_up_weight into q_no_pe
# q_absorbed: [num_tokens, n, kv_lora_rank]
q_absorbed = torch.einsum("...nd,ndk->...nk", q_no_pe, k_up_weight)
q_absorbed = q_absorbed.contiguous()
assert q_absorbed.ndim == q.ndim
assert q_absorbed.shape[:-1] == q.shape[:-1]
assert q_absorbed.size(-1) == self.config.kv_lora_rank

# q_absorbed: [num_tokens, n, (kv_lora_rank + qk_pos_emb_head_dim)]
q_absorbed = torch.cat([q_absorbed, q_pos_emb], dim=-1)
# kv_compressed: [num_tokens, 1, (kv_lora_rank + qk_pos_emb_head_dim)]
kv_compressed = torch.cat([kv_compressed, k_pos_emb], dim=-1)

cp_rank = self.pg_collection.cp.rank()
cp_size = self.pg_collection.cp.size()
q_absorbed = fused_apply_mla_rope_for_q(
q_absorbed = fused_mla_rope_concat(
q_absorbed,
q_pos_emb,
rotary_pos_cos,
rotary_pos_sin,
self.config.kv_lora_rank,
self.config.qk_pos_emb_head_dim,
cu_seqlens_q,
cp_rank,
cp_size,
)
kv_compressed = fused_apply_mla_rope_for_q(
kv_compressed = fused_mla_rope_concat(
kv_compressed,
k_pos_emb,
rotary_pos_cos,
rotary_pos_sin,
self.config.kv_lora_rank,
self.config.qk_pos_emb_head_dim,
cu_seqlens_kv,
cp_rank,
cp_size,
Expand Down
63 changes: 59 additions & 4 deletions megatron/core/transformer/experimental_attention_variant/dsa.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,11 @@
from megatron.core.transformer.transformer_config import TransformerConfig
from megatron.core.utils import get_pg_size

try:
from megatron.core.fusions.fused_mla_yarn_rope_apply import fused_mla_rope_inplace
except ImportError:
fused_mla_rope_inplace = None

try:
from fast_hadamard_transform import hadamard_transform
except ImportError:
Expand Down Expand Up @@ -1347,12 +1352,40 @@ def backward_dw(self):
def _apply_rope(
self,
x: torch.Tensor,
rotary_pos_emb: torch.Tensor,
rotary_pos_emb: Optional[torch.Tensor],
mscale: float,
cu_seqlens: Optional[torch.Tensor] = None,
max_seqlen: Optional[int] = None,
rotary_pos_cos: Optional[torch.Tensor] = None,
rotary_pos_sin: Optional[torch.Tensor] = None,
):
"""Apply RoPE to the input tensor."""
if rotary_pos_cos is not None or rotary_pos_sin is not None:
assert rotary_pos_cos is not None and rotary_pos_sin is not None
assert fused_mla_rope_inplace is not None, "Fused MLA RoPE is not available"
if cu_seqlens is not None and cu_seqlens.device != x.device:
cu_seqlens = cu_seqlens.to(device=x.device)
squeezed_batch_dim = False
# THD RoPE expects [t, h, d], while indexer tensors are [t, 1, h, d].
if cu_seqlens is not None and x.ndim == 4 and x.size(1) == 1:
x = x.squeeze(1)
squeezed_batch_dim = True
x = fused_mla_rope_inplace(
x,
rotary_pos_cos,
rotary_pos_sin,
nope_dim=self.index_head_dim - self.qk_pos_emb_head_dim,
emb_dim=self.qk_pos_emb_head_dim,
cu_seqlens_q=cu_seqlens,
cp_rank=self.pg_collection.cp.rank(),
cp_size=self.pg_collection.cp.size(),
rope_first=True,
)
if squeezed_batch_dim:
x = x.unsqueeze(1)
return x

assert rotary_pos_emb is not None
# x_pe [seqlen, batch, *, qk_pos_emb_head_dim]
# x_nope [seqlen, batch, *, index_head_dim - qk_pos_emb_head_dim]
# To align with DeepSeek's implementation,
Expand Down Expand Up @@ -1396,7 +1429,17 @@ def forward_before_topk(
rotary_seq_len = self.rotary_pos_emb.get_rotary_seq_len(
None, None, x, self.config, packed_seq_params
)
if self.config.rope_type == "rope":
fused_indexer_rope = (
self.config.apply_rope_fusion and self.config.dsa_indexer_rope_interleaved
)
rotary_pos_cos = rotary_pos_sin = None
if fused_indexer_rope:
rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb.get_cached_cos_sin(
rotary_seq_len, dtype=x.dtype, packed_seq=packed_seq
)
rotary_pos_emb = None
mscale = 1.0
elif self.config.rope_type == "rope":
rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len, packed_seq=packed_seq)
mscale = 1.0
else:
Expand Down Expand Up @@ -1430,7 +1473,13 @@ def forward_before_topk(
# -> [seqlen, batch, index_n_heads, index_head_dim]
q = q.reshape(seqlen, bsz, self.index_n_heads, self.index_head_dim)
q = self._apply_rope(
q, rotary_pos_emb, mscale, cu_seqlens=cu_seqlens_q, max_seqlen=max_seqlen_q
q,
rotary_pos_emb,
mscale,
cu_seqlens=cu_seqlens_q,
max_seqlen=max_seqlen_q,
rotary_pos_cos=rotary_pos_cos,
rotary_pos_sin=rotary_pos_sin,
)

# =========================================
Expand All @@ -1446,7 +1495,13 @@ def forward_before_topk(
# [seqlen, batch, index_head_dim] -> [seqlen, batch, 1, index_head_dim]
k = k.reshape(seqlen, bsz, 1, self.index_head_dim)
k = self._apply_rope(
k, rotary_pos_emb, mscale, cu_seqlens=cu_seqlens_kv, max_seqlen=max_seqlen_kv
k,
rotary_pos_emb,
mscale,
cu_seqlens=cu_seqlens_kv,
max_seqlen=max_seqlen_kv,
rotary_pos_cos=rotary_pos_cos,
rotary_pos_sin=rotary_pos_sin,
)
# [seqlen, batch, 1, index_head_dim] -> [seqlen, batch, index_head_dim]
k = k.reshape(seqlen, bsz, self.index_head_dim)
Expand Down
26 changes: 11 additions & 15 deletions megatron/core/transformer/multi_latent_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -704,23 +704,19 @@ def get_query_key_value_tensors(
rotary_pos_cos = None
rotary_pos_sin = None
thd_packed_seq = packed_seq_params is not None and packed_seq_params.qkv_format == 'thd'
if self.config.rope_type == "rope":
if self.config.apply_rope_fusion:
rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb.get_cached_cos_sin(
rotary_seq_len, dtype=hidden_states.dtype, packed_seq=thd_packed_seq
)
rotary_pos_emb = None
assert inference_context is None, "Inference with MLA RoPE fusion is not supported"
assert (
fused_apply_mla_rope_for_q is not None and fused_apply_mla_rope_for_kv is not None
), "Fused MLA RoPE apply is not imported successfully"
elif self.config.rope_type == "rope":
rotary_pos_emb = self.rotary_pos_emb(rotary_seq_len, packed_seq=thd_packed_seq)
else:
if self.config.apply_rope_fusion:
rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb.get_cached_cos_sin(
rotary_seq_len, dtype=hidden_states.dtype, packed_seq=thd_packed_seq
)
rotary_pos_emb = None
assert inference_context is None, "Inference with MLA RoPE fusion is not supported"
assert (
fused_apply_mla_rope_for_q is not None
and fused_apply_mla_rope_for_kv is not None
), "Fused MLA RoPE apply is not imported successfully"
else:
rotary_pos_emb, mscale = self.rotary_pos_emb(
rotary_seq_len, packed_seq=thd_packed_seq
)
rotary_pos_emb, mscale = self.rotary_pos_emb(rotary_seq_len, packed_seq=thd_packed_seq)

if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd':
if packed_seq_params.cu_seqlens_q_padded is not None:
Expand Down
8 changes: 0 additions & 8 deletions megatron/core/transformer/transformer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -1680,7 +1680,6 @@ def __post_init__(self):
"dsa_indexer_skip_topk_offset must be non-negative, got "
f"{self.dsa_indexer_skip_topk_offset}."
)
assert not self.apply_rope_fusion, "RoPE fusion is not supported for DSAttention"
if self.context_parallel_size > 1:
cp_comm_types = (
self.cp_comm_type
Expand Down Expand Up @@ -3591,13 +3590,6 @@ class MLATransformerConfig(TransformerConfig):

def __post_init__(self):
super().__post_init__()
if (
self.multi_latent_attention
and self.apply_rope_fusion
and self.rope_type != "yarn"
and self.experimental_attention_variant != "dsv4_hybrid"
):
raise ValueError("apply_rope_fusion for MLA only works with YARN RoPE.")

if self.attention_output_gate:
raise NotImplementedError("Output gate is not supported for MLA yet.")
Expand Down
Loading
Loading