diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index 4b01ba561b6..37778aad4ee 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -10,6 +10,10 @@ AbsorbedMLASelfAttention, AbsorbedMLASelfAttentionSubmodules, ) +from megatron.core.transformer.multi_latent_attention import ( + MLASelfAttention, + MLASelfAttentionSubmodules, +) from megatron.core.transformer.experimental_attention_variant.dsa import ( DSAIndexer, DSAIndexerSubmodules, @@ -115,6 +119,25 @@ def get_dsa_module_spec_for_backend( _get_standalone_norm(config, backend, for_qk=True) if config.qk_layernorm else IdentityOp ) + if getattr(config, "use_accuracy_compatible", False): + attention = ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=backend.column_parallel_linear(), + linear_q_down_proj=backend.linear(), + linear_q_up_proj=backend.column_parallel_linear(), + linear_kv_down_proj=backend.linear(), + linear_kv_up_proj=backend.column_parallel_linear(), + core_attention=core_attention, + linear_proj=backend.row_parallel_linear(), + q_layernorm=qk_norm, + kv_layernorm=qk_norm, + ), + metainfo={"fuse_input_layernorm": False}, + ) + return attention + attention = ModuleSpec( module=AbsorbedMLASelfAttention, params={"attn_mask_type": AttnMaskType.causal}, diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index 8e0c135d725..1de334cb3a3 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -22,14 +22,28 @@ dsa_layout, dsa_masking, ) -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_config import TransformerConfig try: from fast_hadamard_transform import hadamard_transform except ImportError: - hadamard_transform = None + def hadamard_transform(x, scale=1.0): + n = x.shape[-1] + assert (n & (n - 1)) == 0, f"hadamard fallback needs power-of-2 last dim, got {n}" + orig_dtype = x.dtype + shape = x.shape + y = x.to(torch.float32).reshape(-1, n) + h = 1 + while h < n: + y = y.reshape(-1, n // (2 * h), 2, h) + a = y[:, :, 0, :] + b = y[:, :, 1, :] + y = torch.stack((a + b, a - b), dim=2).reshape(-1, n) + h *= 2 + return (y * scale).reshape(shape).to(orig_dtype) + def is_dsa_skip_topk_layer(layer_number: int, skip_topk_offset: int, topk_freq: int) -> bool: @@ -116,6 +130,22 @@ def _unfused_absorbed_dsa_fn( return output.permute(2, 0, 1, 3).contiguous() +def _aligned_sum_dim(x: torch.Tensor, dim: int) -> torch.Tensor: + xd = x.double() + acc = xd.narrow(dim, 0, 1) + for i in range(1, xd.size(dim)): + acc = acc + xd.narrow(dim, i, 1) + return acc + + +def _aligned_sum_last_dim(x: torch.Tensor) -> torch.Tensor: + xd = x.double() + acc = xd.narrow(-1, 0, 1) + for i in range(1, xd.size(-1)): + acc = acc + xd.narrow(-1, i, 1) + return acc + + class _AccuracyCompatibleSoftmax(torch.autograd.Function): """Masked softmax with an explicit backward formula for DSA alignment.""" @@ -129,9 +159,8 @@ def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor @staticmethod def backward(ctx, grad_output: torch.Tensor): probabilities, valid_mask = ctx.saved_tensors - grad_logits = probabilities * ( - grad_output - (grad_output * probabilities).sum(dim=-1, keepdim=True) - ) + row_sum = _aligned_sum_last_dim(grad_output * probabilities).float() + grad_logits = probabilities * (grad_output - row_sum) return grad_logits.masked_fill(~valid_mask, 0.0), None @@ -811,7 +840,10 @@ def bwd_fused_indexer_loss_naive( del index_scores # Sum attention scores across heads: [b, np, sq, sk] -> [b, sq, sk] - attention_scores_sum = attention_scores_softmax.sum(dim=1) + if _use_accuracy_compatible(): + attention_scores_sum = _aligned_sum_dim(attention_scores_softmax, 1).squeeze(1).float() + else: + attention_scores_sum = attention_scores_softmax.sum(dim=1) # Free attention_scores_softmax del attention_scores_softmax @@ -822,9 +854,11 @@ def bwd_fused_indexer_loss_naive( # L1 normalize. Fully masked packed/varlen rows can have zero summed # attention mass; clamp the denominator so those rows stay finite and are # later zeroed by the row-valid loss mask. - attention_scores_normalized = attention_scores_sum / attention_scores_sum.sum( - dim=-1, keepdim=True - ).clamp_min(1e-10) + if _use_accuracy_compatible(): + _denom = _aligned_sum_last_dim(attention_scores_sum).float() + else: + _denom = attention_scores_sum.sum(dim=-1, keepdim=True) + attention_scores_normalized = attention_scores_sum / _denom.clamp_min(1e-10) # Free attention_scores_sum - no longer needed after normalization del attention_scores_sum @@ -863,7 +897,12 @@ def bwd_fused_indexer_loss_naive( del attention_scores_normalized # Backward through softmax: ∂L/∂x = softmax * (∂L/∂softmax - sum(∂L/∂softmax * softmax)) - sum_grad = (grad_index_scores_softmax * index_scores_softmax).sum(dim=-1, keepdim=True) + if _use_accuracy_compatible(): + sum_grad = _aligned_sum_last_dim( + grad_index_scores_softmax * index_scores_softmax + ).float() + else: + sum_grad = (grad_index_scores_softmax * index_scores_softmax).sum(dim=-1, keepdim=True) grad_index_scores_logits = index_scores_softmax * (grad_index_scores_softmax - sum_grad) # Free intermediate tensors del index_scores_softmax, grad_index_scores_softmax, sum_grad @@ -905,7 +944,14 @@ def bwd_fused_indexer_loss_naive( del scores # ∂L/∂weights = grad * scores_for_weights (sum over sk) - grad_weights = (grad_weighted_scores * scores_for_weights).sum(dim=-1) # [sq, b, h] + if _use_accuracy_compatible(): + grad_weights = ( + _aligned_sum_last_dim(grad_weighted_scores * scores_for_weights) + .squeeze(-1) + .float() + ) # [sq, b, h] + else: + grad_weights = (grad_weighted_scores * scores_for_weights).sum(dim=-1) # [sq, b, h] # ∂L/∂scores = grad * weights grad_scores = grad_weighted_scores * weights.unsqueeze(-1) # [sq, b, h, sk] @@ -917,10 +963,16 @@ def bwd_fused_indexer_loss_naive( del relu_mask # Backward through einsum 'sbhd,tbd->sbht' - # ∂L/∂q = einsum('sbht,tbd->sbhd', grad_scores, k) - grad_q = torch.einsum('sbht,tbd->sbhd', grad_scores, k.float()) # [sq, b, h, d] - # ∂L/∂k = einsum('sbht,sbhd->tbd', grad_scores, q) - grad_k = torch.einsum('sbht,sbhd->tbd', grad_scores, q.float()) # [sk, b, d] + if _use_accuracy_compatible(): + # ∂L/∂q = einsum('sbht,tbd->sbhd', grad_scores, k) + grad_q = torch.einsum('sbht,tbd->sbhd', grad_scores.double(), k.double()).float() + # ∂L/∂k = einsum('sbht,sbhd->tbd', grad_scores, q) + grad_k = torch.einsum('sbht,sbhd->tbd', grad_scores.double(), q.double()).float() + else: + # ∂L/∂q = einsum('sbht,tbd->sbhd', grad_scores, k) + grad_q = torch.einsum('sbht,tbd->sbhd', grad_scores, k.float()) # [sq, b, h, d] + # ∂L/∂k = einsum('sbht,sbhd->tbd', grad_scores, q) + grad_k = torch.einsum('sbht,sbhd->tbd', grad_scores, q.float()) # [sk, b, d] del grad_scores return grad_q.to(q.dtype), grad_weights.to(weights.dtype), grad_k.to(k.dtype) @@ -1728,6 +1780,9 @@ def forward( Returns: output: Output tensor [sq, b, hidden_size] """ + _indexer_packed_seq_params = packed_seq_params + if isinstance(packed_seq_params, tuple): + packed_seq_params = packed_seq_params[0] query, _ = dsa_layout.ensure_sbhd(query, "query") key, _ = dsa_layout.ensure_sbhd(key, "key") if value is not None: @@ -1908,7 +1963,7 @@ def forward( topk_length = topk_length_holder.get(self.source_layer) else: assert self.indexer is not None - q, k, weights = self.indexer.forward_before_topk(x, qr, packed_seq_params) + q, k, weights = self.indexer.forward_before_topk(x, qr, _indexer_packed_seq_params) if cp_size > 1 and k.size(0) == sq: k = gather_from_sequence_parallel_region(k, group=cp_group) if kv_reorder_idx is not None: