Skip to content
Merged
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
17 changes: 17 additions & 0 deletions megatron/core/pipeline_parallel/schedules.py
Original file line number Diff line number Diff line change
Expand Up @@ -321,6 +321,23 @@ def forward_step_calc_loss(
else:
MTPLossAutoScaler.set_loss_scale(loss_scale / num_microbatches)

# Set the loss scale for the DSA indexer loss.
if hasattr(config, 'dsa_indexer_loss_coeff') and config.dsa_indexer_loss_coeff is not None:
from megatron.core.transformer.experimental_attention_variant.dsa import (
DSAIndexerLossAutoScaler,
)

device = get_tensor_device(output_tensor)
loss_scale = (
config.grad_scale_func(torch.ones(1, device=device))
if config.grad_scale_func is not None
else torch.ones(1, device=device)
)
if config.calculate_per_token_loss:
DSAIndexerLossAutoScaler.set_loss_scale(loss_scale)
else:
DSAIndexerLossAutoScaler.set_loss_scale(loss_scale / num_microbatches)

return output_tensor, num_tokens


Expand Down
395 changes: 294 additions & 101 deletions megatron/core/transformer/experimental_attention_variant/csa.py

Large diffs are not rendered by default.

43 changes: 37 additions & 6 deletions megatron/core/transformer/experimental_attention_variant/dsa.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,6 +196,7 @@ def compute_dsa_indexer_loss(
sparse_loss: bool,
pg_collection: ProcessGroupCollection,
causal_mask_override: Optional[torch.Tensor] = None,
calculate_per_token_loss: bool = False,
) -> torch.Tensor:
"""
Compute KL divergence loss between index_scores and true attention_scores.
Expand All @@ -216,6 +217,10 @@ def compute_dsa_indexer_loss(
sparse_loss: bool, whether to use sparse indexer loss. If True, only the topk
indices will be used to compute the loss.
pg_collection: Process group collection, must have TP process group.
causal_mask_override: Optional mask used by compressed KV paths.
calculate_per_token_loss: If True, return a raw local sum so the global
token divisor can be applied by finalize_model_grads. If False, keep
the historical local BSHD average over ``batch * seqlen`` rows.

Returns:
index_loss: KL divergence loss (scalar).
Expand Down Expand Up @@ -308,7 +313,11 @@ def compute_dsa_indexer_loss(

# [b, sq, sk] -> [b, sq] -> [1]
# Each token has same weight in the loss.
kl_div = kl_per_element.sum(dim=-1).mean()
kl_per_row = kl_per_element.sum(dim=-1)
if calculate_per_token_loss:
kl_div = kl_per_row.sum()
else:
kl_div = kl_per_row.mean()

# Scale by coefficient.
indexer_loss = kl_div * loss_coeff
Expand Down Expand Up @@ -388,7 +397,18 @@ def fused_qk_topk_naive(


def fwd_fused_indexer_loss_naive(
q, weights, k, query, key, topk, softmax_scale, loss_coeff, mask, sparse_loss, pg_collection
q,
weights,
k,
query,
key,
topk,
softmax_scale,
loss_coeff,
mask,
sparse_loss,
pg_collection,
calculate_per_token_loss,
):
"""Naive implementation of forward pass for indexer loss."""
index_scores, topk_indices = fused_qk_topk_naive(q, k, weights, topk, mask)
Expand All @@ -403,6 +423,7 @@ def fwd_fused_indexer_loss_naive(
sparse_loss,
pg_collection,
causal_mask_override=mask,
calculate_per_token_loss=calculate_per_token_loss,
)

return topk_indices, indexer_loss
Expand All @@ -421,6 +442,7 @@ def bwd_fused_indexer_loss_naive(
grad_loss,
pg_collection,
causal_mask_override=None,
calculate_per_token_loss=False,
):
"""Naive implementation of backward pass for indexer loss."""
index_scores = _compute_index_scores(q, weights, k) # [B, Sq, Sk]
Expand Down Expand Up @@ -520,11 +542,15 @@ def bwd_fused_indexer_loss_naive(
del attention_scores_sum

# Backward through loss = kl_div * loss_coeff
# where kl_div = kl_per_element.sum(dim=-1).mean()
# where kl_div is either kl_per_element.sum(dim=-1).mean() or the raw
# local sum when calculate_per_token_loss=True.
grad_kl_div = grad_loss * loss_coeff # scalar

# Backward through mean: distribute gradient equally
grad_kl_per_row = grad_kl_div / (b * sq) # scalar value for each row
if calculate_per_token_loss:
grad_kl_per_row = grad_kl_div
else:
# Backward through mean: distribute gradient equally
grad_kl_per_row = grad_kl_div / (b * sq) # scalar value for each row

# Backward through sum(dim=-1): broadcast back to [b, sq, sk]
# Each element in a row contributes to the sum, so gradient is same for all
Expand Down Expand Up @@ -630,6 +656,7 @@ def forward(
mask,
sparse_loss,
pg_collection,
calculate_per_token_loss,
):
"""
Fused forward: index_scores never materialized in full.
Expand All @@ -646,6 +673,7 @@ def forward(
mask,
sparse_loss,
pg_collection,
calculate_per_token_loss,
)

# Save for backward (recomputation strategy)
Expand All @@ -654,6 +682,7 @@ def forward(
ctx.loss_coeff = loss_coeff
ctx.sparse_loss = sparse_loss
ctx.pg_collection = pg_collection
ctx.calculate_per_token_loss = calculate_per_token_loss

return topk_indices, loss

Expand All @@ -677,10 +706,11 @@ def backward(ctx, grad_topk_indices, grad_loss):
grad_loss,
ctx.pg_collection,
causal_mask_override=mask,
calculate_per_token_loss=ctx.calculate_per_token_loss,
)

# query and key are detached in forward, so return None for their gradients
return grad_q, grad_weights, grad_k, None, None, None, None, None, None, None, None
return grad_q, grad_weights, grad_k, None, None, None, None, None, None, None, None, None


class DSAIndexerLossAutoScaler(torch.autograd.Function):
Expand Down Expand Up @@ -1201,6 +1231,7 @@ def forward(
float_mask,
getattr(self.config, "dsa_indexer_use_sparse_loss", False),
self.indexer.pg_collection,
self.config.calculate_per_token_loss,
)
# Save indexer loss for logging
if indexer_loss_coeff > 0:
Expand Down
Loading
Loading