@NVIDIA/mcore-oncall
Describe the bug
The DeepSeek-V4 unfused / small-operator CSA path computes the indexer-loss teacher by recomputing attention over compressed KV only. This teacher is not equivalent to the real CSA attention distribution described by the DeepSeek-V4 paper, because the real CSA core attention also includes the original-KV sliding window and the attention sink in a shared softmax denominator.
This is a semantic mismatch in the teacher distribution, not a floating-point precision difference.
Static analysis scope:
- Megatron-LM dev commit: fd1121b
- Primary path: unfused CSA indexer loss, both SBHD and THD
- No runtime or convergence claim is required to establish the mismatch below
Related DSV4 tracking issue: #4468
Paper reference
DeepSeek-V3.2 defines the Lightning Indexer teacher from the real main-attention distribution: per-head attention probabilities are aggregated across heads and then L1-normalized (Section 2.1, equations (3) and (4)):
https://arxiv.org/html/2512.02556v1#S2.SS1.SSS1
DeepSeek-V4 keeps the Lightning Indexer objective while changing the CSA core attention. The CSA attention contains all of the following components:
- Compressed KV selected by the Lightning Indexer:
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS1.Px2
- Original KV from the sliding window:
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS3.Px3
- An attention sink in the same normalization denominator:
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS3.Px4
For a query and attention head h, let c(h,j) be a compressed-key logit, w(h,r) a window-key logit, and a(h) the sink logit. The compressed part of the paper-equivalent teacher must retain the real per-head CSA denominator:
$$Z_h = \exp(a_h) + \sum_{j \in C}\exp(c_{h,j}) + \sum_{r \in W}\exp(w_{h,r})$$
$$t_j =
\frac{
\sum_h \frac{\exp(c_{h,j})}{Z_h}
}{
\sum_{k \in C}\sum_h \frac{\exp(c_{h,k})}{Z_h}
}$$
Consequently, changing only the window logits or sink logits can change the compressed teacher distribution because it changes the relative contribution of each attention head.
Megatron-LM unfused / small-operator implementation
1. The loss receives compressed KV only
In the SBHD unfused path, key_for_loss is constructed only from compressed_kv and passed to FusedDSAIndexerLoss:
|
causal_mask = ( |
|
torch.arange(n_compressed, device=x.device).unsqueeze(0).expand(sq, -1) |
|
) |
|
positions = torch.arange(1, sq + 1, device=x.device).unsqueeze(1) |
|
causal_mask = ( |
|
torch.where(causal_mask >= positions // self.compress_ratio, float("-inf"), 0.0) |
|
.unsqueeze(0) |
|
.expand(b, -1, -1) |
|
) # [b, sq, n_compressed] |
|
|
|
if self.training and torch.is_grad_enabled(): |
|
q_indexer, k_indexer, weights_indexer = self.indexer.forward_before_topk( |
|
x_det, qr_det, packed_seq_params |
|
) |
|
indexer_loss_coeff = getattr(self.config, 'dsa_indexer_loss_coeff', 0.0) |
|
key_for_loss = compressed_kv.unsqueeze(2).expand(-1, -1, np, -1) |
|
# ``FusedDSAIndexerLoss`` does not accept a separate |
|
# indexer_softmax_scale; apply it here via the |
|
# weights-scaling trick so the effective weights match |
|
# the pre-scale-split behaviour. |
|
weights_for_unfused = weights_indexer.float() * self.indexer.softmax_scale |
|
topk_indices_compressed, indexer_loss = FusedDSAIndexerLoss.apply( |
|
q_indexer, |
|
weights_for_unfused, |
|
k_indexer, |
|
query.detach(), |
|
key_for_loss.detach(), |
|
self.softmax_scale, |
|
min(self.indexer.index_topk, n_compressed), |
|
indexer_loss_coeff, |
|
causal_mask, |
|
getattr(self.config, "dsa_indexer_use_sparse_loss", True), |
|
self.indexer.pg_collection, |
|
self.config.calculate_per_token_loss, |
|
) |
Relevant data flow:
key_for_loss = compressed_kv.unsqueeze(2).expand(-1, -1, np, -1)
topk_indices_compressed, indexer_loss = FusedDSAIndexerLoss.apply(
...,
query.detach(),
key_for_loss.detach(),
...
)
Neither window KV/logits nor attn_sink is an input to this loss call.
The THD unfused path does the same with key_for_loss_thd:
|
if self.training and torch.is_grad_enabled(): |
|
q_indexer, k_indexer, weights_indexer, cu_seqlens_compressed_idx = ( |
|
self.indexer.forward_before_topk(x_det, qr_det, packed_seq_params) |
|
) |
|
if k_indexer is None: |
|
raise RuntimeError( |
|
"CompressedSparseAttention THD unfused Path B requires " |
|
"at least one segment with compressed indexer K." |
|
) |
|
q_thd = q_indexer.squeeze(1) |
|
w_thd = weights_indexer.squeeze(1) |
|
k_thd = k_indexer.squeeze(1) |
|
|
|
key_for_loss_thd = compressed_kv.unsqueeze(1).expand(-1, np_, -1) |
|
weights_for_unfused = w_thd * self.indexer.softmax_scale |
|
indexer_loss_coeff = getattr(self.config, 'dsa_indexer_loss_coeff', 0.0) |
|
|
|
# ``_forward_thd`` (caller) absorbs trailing padded |
|
# tokens into the last ``cu_seqlens_q[-1]`` bucket |
|
# so ``batch_of_row`` doesn't OOB; the Compressor |
|
# does the same to ``cu_seqlens_compressed_idx[-1]``. |
|
# Both are correct for the sparse-attention path, |
|
# but the per-segment indexer-loss loop in |
|
# ``fwd/bwd_fused_indexer_loss_naive_thd`` would |
|
# then iterate a "fake" absorbed-padding segment |
|
# with ``seqlen_k_b < topk`` — triggering a write |
|
# shape mismatch and a downstream ``scatter_`` OOB |
|
# on its ``-1`` entries. |
|
# |
|
# Restore the original (pre-absorption) cu_seqlens |
|
# for the loss path so the segment loop's |
|
# ``if seqlen_k_b == 0: continue`` guard skips the |
|
# padding-only iteration. When no padding exists, |
|
# ``packed_seq_params`` already equals the absorbed |
|
# version and this is a no-op. |
|
# Rebuild compressed cu_seqlens from *unpadded* Q lengths. |
|
# This may disagree with the Indexer's cu_seqlens_compressed_idx |
|
# (which uses padded lengths) for the last segment — that's |
|
# intentional: the extra compressed tokens from padding sit at |
|
# the tail of k_thd and are simply never visited by the loss |
|
# loop, which is correct since they don't represent real data. |
|
# Non-last segments are unaffected (padding absorption only |
|
# extends the final segment). |
|
cu_seqlens_q_for_loss = packed_seq_params.cu_seqlens_q |
|
seg_lens_q = cu_seqlens_q_for_loss[1:] - cu_seqlens_q_for_loss[:-1] |
|
cu_seqlens_compressed_idx_for_loss = torch.cat( |
|
[ |
|
torch.zeros( |
|
1, |
|
dtype=cu_seqlens_q_for_loss.dtype, |
|
device=cu_seqlens_q_for_loss.device, |
|
), |
|
(seg_lens_q // self.compress_ratio) |
|
.cumsum(0) |
|
.to(cu_seqlens_q_for_loss.dtype), |
|
] |
|
) |
|
topk_indices_cmp, indexer_loss = FusedDSAIndexerLoss.apply( |
|
q_thd, |
|
weights_for_unfused, |
|
k_thd, |
|
query.detach(), |
|
key_for_loss_thd.detach(), |
|
self.softmax_scale, |
|
min(self.indexer.index_topk, max_seqlen_compressed_idx), |
|
indexer_loss_coeff, |
|
None, |
|
getattr(self.config, "dsa_indexer_use_sparse_loss", True), |
|
self.indexer.pg_collection, |
|
self.config.calculate_per_token_loss, |
|
cu_seqlens_q_for_loss, |
|
cu_seqlens_compressed_idx_for_loss, |
|
self.compress_ratio, |
|
) |
2. The small operators recompute a compressed-only softmax
compute_dsa_indexer_loss calculates QK scores from the supplied query and key, applies the compressed causal/Top-k masks, softmaxes over that key axis, sums attention heads, and L1-normalizes:
|
# [sk, b, np, hn] -> [b, np, hn, sk] -> [b * np, hn, sk] |
|
key = key.permute(1, 2, 3, 0).reshape(b * np, hn, sk) |
|
# Compute attention scores [b * np, sq, sk] |
|
attention_scores = torch.bmm(query.float(), key.float()) * softmax_scale |
|
# Reshape to [b, np, sq, sk] |
|
attention_scores = attention_scores.reshape(b, np, sq, sk) |
|
|
|
# causal_mask: use caller-provided mask when available (handles compressed KV), |
|
# otherwise fall back to standard upper-triangular causal mask. |
|
if causal_mask_override is not None: |
|
causal_mask = causal_mask_override.to(dtype=torch.float32) # [b, sq, sk] |
|
else: |
|
causal_mask = torch.triu( |
|
torch.full( |
|
(sq, sk), float('-inf'), dtype=torch.float32, device=attention_scores.device |
|
), |
|
diagonal=1, |
|
) |
|
# index_mask [b, sq, sk] |
|
index_mask = torch.full( |
|
(b, sq, sk), float("-inf"), dtype=torch.float32, device=causal_mask.device |
|
).scatter_(-1, topk_indices, 0) |
|
|
|
# Apply causal mask to attention_scores |
|
# causal_mask: [b, sq, sk] (from causal_mask_override) or [sq, sk] (from triu) |
|
if causal_mask.dim() == 3: |
|
attention_scores = attention_scores + causal_mask.unsqueeze(1) # [b,1,sq,sk] |
|
else: |
|
attention_scores = attention_scores + causal_mask.view(1, 1, sq, sk) |
|
if sparse_loss: |
|
# [b, np, sq, sk] + [b, 1, sq, sk] -> [b, np, sq, sk] |
|
attention_scores += index_mask.view(b, 1, sq, sk) |
|
# [b, sq, sk] + [b, sq, sk] -> [b, sq, sk] |
|
index_scores += index_mask |
|
|
|
# Identify rows where all KV positions are masked (e.g., early query positions with |
|
# compress_ratio=4 have zero valid compressed KV entries). These rows would produce NaN |
|
# from softmax(all -inf). We zero out their logits before softmax and mask out their |
|
# contributions after, so NaN is never produced. |
|
# row_valid: [b, sq] or [sq] — True if the row has at least one unmasked position. |
|
row_valid = (causal_mask > float('-inf')).any(dim=-1) |
|
if row_valid.dim() == 1: |
|
# [sq] -> broadcast for attention_scores [b, np, sq, sk] and index_scores [b, sq, sk] |
|
attn_row_mask = row_valid.view(1, 1, sq, 1) # [1, 1, sq, 1] |
|
idx_row_mask = row_valid.view(1, sq, 1) # [1, sq, 1] |
|
else: |
|
# [b, sq] |
|
attn_row_mask = row_valid.view(b, 1, sq, 1) # [b, 1, sq, 1] |
|
idx_row_mask = row_valid.view(b, sq, 1) # [b, sq, 1] |
|
|
|
# Zero out fully-masked rows before softmax so it produces valid uniform distribution |
|
attention_scores = attention_scores.masked_fill(~attn_row_mask, 0.0) |
|
index_scores = index_scores.masked_fill(~idx_row_mask, 0.0) |
|
|
|
# [b, np, sq, sk] -> [b, np, sq, sk] |
|
attention_scores = torch.nn.functional.softmax(attention_scores, dim=-1, dtype=torch.float32) |
|
# [b, sq, sk] -> [b, sq, sk] |
|
index_scores = torch.nn.functional.softmax(index_scores, dim=-1, dtype=torch.float32) |
|
|
|
# Zero out invalid rows so they contribute nothing to loss/gradients |
|
attention_scores = attention_scores * attn_row_mask.float() |
|
index_scores = index_scores * idx_row_mask.float() |
|
|
|
# Sum attention scores across heads. |
|
# [batch, heads, seqlen_q, seqlen_k] -> [batch, seqlen_q, seqlen_k] |
|
attention_scores = attention_scores.sum(dim=1) |
|
if pg_collection.tp.size() > 1: |
|
# attention scores are scattered to TP ranks in head dimension. |
|
torch.distributed.all_reduce(attention_scores.contiguous(), group=pg_collection.tp) |
|
# L1 normalize target on the last dimension. Doesn't use abs() because attention_scores are |
|
# obtained from softmax so they are already non-negative. |
|
attention_scores = attention_scores / ( |
|
attention_scores.sum(dim=-1, keepdim=True).clamp(min=1e-10) |
|
) |
|
|
|
# Compute KL divergence: KL(target || index) = target(x) * log(target(x) / index(x)) |
|
# kl_per_element [b, sq, sk] |
|
kl_per_element = attention_scores * ( |
|
torch.log(attention_scores + 1e-10) - torch.log(index_scores + 1e-10) |
Therefore, ignoring masks for clarity, the implemented teacher is:
$$\hat{t}_j =
\frac{
\sum_h
\frac{\exp(c_{h,j})}
{\sum_{l \in C}\exp(c_{h,l})}
}{
\sum_{k \in C}\sum_h
\frac{\exp(c_{h,k})}
{\sum_{l \in C}\exp(c_{h,l})}
}$$
With sparse loss enabled, the Top-k mask is applied before this per-head softmax, so C above becomes the selected Top-k subset.
This differs from the paper teacher because the window and sink terms have been removed from every per-head denominator. It forces every head to contribute unit compressed probability mass before head aggregation, even when a head's real CSA probability mass is mostly assigned to its sliding window or sink.
3. The real unfused CSA attention uses a different denominator
Only after the indexer loss has already been computed, the implementation concatenates the window indices with the compressed Top-k and calls the actual sparse attention with attn_sink:
|
n_valid_per_pos = positions // self.compress_ratio # [sq, 1] |
|
valid = topk_indices_compressed < n_valid_per_pos |
|
compress_topk_idxs = torch.where( |
|
valid, topk_indices_compressed + offset, torch.tensor(-1, device=x.device) |
|
) |
|
else: |
|
compress_topk_idxs = get_compress_topk_idxs( |
|
self.compress_ratio, b, sq, offset, query.device |
|
) |
|
|
|
topk_idxs = torch.cat([window_idxs, compress_topk_idxs], dim=-1) |
|
nvtx_range_pop("compressed_indices") |
|
else: |
|
topk_idxs = window_idxs |
|
|
|
topk_idxs = topk_idxs.int() |
|
|
|
nvtx_range_push("sparse_attn_kernel") |
|
output = unfused_compressed_sparse_attn( |
|
query, kv_full, self.attn_sink.float(), topk_idxs, self.softmax_scale |
|
) |
|
nvtx_range_pop("sparse_attn_kernel") |
topk_idxs = torch.cat([window_idxs, compress_topk_idxs], dim=-1)
output = unfused_compressed_sparse_attn(
query, kv_full, self.attn_sink.float(), topk_idxs, self.softmax_scale
)
Thus the output attention and the indexer-loss teacher do not use the same attention distribution:
| Component |
Real unfused CSA output |
Unfused indexer-loss teacher |
| Compressed KV |
Yes |
Yes |
| Original-KV sliding window |
Yes |
No |
| Attention sink |
Yes |
No |
| Shared CSA softmax denominator |
Yes |
No |
Static reproduction / invariant
Hold query and compressed_kv fixed, and perturb only either:
- the original-KV window values/logits; or
- attn_sink for one attention head.
Expected from the paper equations:
- the per-head CSA denominator changes;
- the relative head weights in the compressed teacher can change;
- therefore the final L1-normalized teacher can change.
Current unfused implementation:
- FusedDSAIndexerLoss receives exactly the same query and key_for_loss;
- the recomputed compressed-only teacher is exactly unchanged.
This invariant violation follows directly from the function inputs and does not depend on a particular GPU kernel or numerical tolerance.
Additional Top-k concern
When dsa_indexer_topk is smaller than the number of visible compressed keys, sparse loss applies the Top-k mask before the per-head attention softmax. This is generally different from:
- computing the teacher with the full real CSA denominator;
- gathering the selected compressed entries; and
- applying one final L1 normalization.
The current short/full-selection cases can hide this additional difference.
Expected behavior
The indexer-loss teacher should be derived from the real CSA attention probabilities, or from an equivalent recomputation that retains the same per-head denominator containing compressed attention, window attention, and sink.
For sparse loss, selected compressed entries should be derived from that teacher distribution and then normalized according to the Lightning Indexer objective, rather than rebuilding independent per-head softmax denominators over the selected compressed subset.
Additional context: fused path
The fused path appears to have the same semantic issue. It places compressed Top-k entries before window entries, obtains both lse and lse_indexer from FlashMLA, but passes lse_indexer to the target recomputation:
|
) |
|
combined_local = torch.cat([compress_topk_idxs, window_idxs], dim=-1) |
|
global_idxs = local_to_global_flat(combined_local, b) |
|
|
|
# ---- 4. FlashMLA forward (flat layout for both SBHD and THD). -------- |
|
if is_thd: |
|
q_flat = query |
|
kv_flat = kv_full |
|
else: |
|
q_flat = query.reshape(sq * b, np_, d) |
|
kv_flat = kv_full.reshape(skv * b, d) |
|
out_flat, lse, lse_indexer = _dsa_fwd_flash_mla( |
|
q_flat, |
|
kv_flat, |
|
global_idxs, |
|
softmax_scale, |
|
attn_sink=attn_sink, |
|
topk_length=None, |
|
indexer_topk=indexer_topk, |
|
) |
|
|
|
# ---- 4b. Derive padding-row mask for loss exclusion. ----------------- |
|
# When CUDA-graph padding makes cu_seqlens_q cover all total_q rows |
|
# (including padding), cu_seqlens_q_unpadded supplies the true |
|
# boundaries. Padding rows must not contribute to the indexer KL |
|
# loss or backward gradients — only the sparse-attention output |
|
# needs them for static-shape compatibility. |
|
# The caller only passes cu_seqlens_q_unpadded when it differs from |
|
# cu_seqlens_q (checked via data_ptr), so no GPU→CPU sync is needed. |
|
padding_row_mask: Optional[Tensor] = None # True = padding (excluded from loss) |
|
if is_thd and cu_seqlens_q_unpadded is not None: |
|
real_seg_lens = cu_seqlens_q_unpadded[1:] - cu_seqlens_q_unpadded[:-1] |
|
row_idx = torch.arange(total_q, device=query.device, dtype=torch.int32) |
|
row_batch_ids = batch_of_row(cu_seqlens_q, total_q=total_q) |
|
pos_in_seg = row_idx - cu_seqlens_q[row_batch_ids].to(torch.int32) |
|
# Rows whose intra-segment position >= real segment length |
|
# are padding (including rows in a dummy trailing segment |
|
# whose real length is 0). |
|
real_len_per_row = real_seg_lens[row_batch_ids].to(torch.int32) |
|
padding_row_mask = pos_in_seg >= real_len_per_row |
|
|
|
# ---- 5. Derive predict from indexer_scores, compute target. ---------- |
|
# Layout-specific attn tensors (detached — loss is not differentiable |
|
# through them). |
|
if is_thd: |
|
assert compressed_kv is not None, "compressed_kv is required for THD" |
|
q_attn_det = query.detach() |
|
k_attn_compressed_det = compressed_kv.detach() |
|
lse_indexer_det = lse_indexer.detach() |
|
else: |
|
q_attn_det = query.detach().permute(1, 0, 2, 3).contiguous() |
|
k_attn_compressed_det = kv_full[kv_offset:].detach().permute(1, 0, 2).contiguous() |
|
lse_indexer_det = lse_indexer.reshape(sq, b, np_).permute(1, 0, 2) |
|
|
|
# Invalidate padding rows for the loss/backward path. The sparse |
|
# attention (steps 3-4) has already built global_idxs from the |
|
# original topk_indices_cmp, so this mutation only affects steps 5-7. |
|
if padding_row_mask is not None: |
|
topk_indices_cmp = topk_indices_cmp.clone() |
|
topk_indices_cmp[padding_row_mask] = -1 |
|
indexer_scores = indexer_scores.clone() |
|
indexer_scores[padding_row_mask] = float('-inf') |
|
|
|
if sparse_loss: |
|
# Derive predict: gather topk scores from indexer_scores → softmax. |
|
safe_indices = topk_indices_cmp.clamp(min=0).long() |
|
gathered_scores = torch.gather(indexer_scores, dim=-1, index=safe_indices) |
|
gathered_scores = torch.where( |
|
topk_indices_cmp >= 0, gathered_scores, torch.finfo(torch.float32).min |
|
) |
|
predict = torch.softmax(gathered_scores, dim=-1) |
|
|
|
# THD: _compute_attn_target's kernel addresses K by flat ids over |
|
# the packed (total_k, D) buffer, so promote per-segment-local |
|
# indices to flat-global against cu_seqlens_compressed_idx. |
|
if is_thd: |
|
topk_for_target = local_to_global_flat( |
|
topk_indices_cmp, |
|
batch_size=-1, |
|
cu_seqlens_q=cu_seqlens_q, |
|
cu_seqlens_kv=cu_seqlens_compressed_idx, |
|
) |
|
else: |
|
topk_for_target = topk_indices_cmp |
|
|
|
target = _compute_attn_target( |
|
q_attn_det, |
|
k_attn_compressed_det, |
|
lse_indexer_det, |
|
topk_for_target, |
|
softmax_scale, |
|
qhead_per_kv_head=np_, |
|
topk_indices_global=is_thd, |
|
) |
|
|
|
if loss_coeff > 0: |
|
indexer_loss = _kl_loss_from_target_predict( |
|
target, predict, topk_indices_cmp, loss_coeff, calculate_per_token_loss |
|
) |
|
else: |
|
indexer_loss = torch.zeros((), device=query.device, dtype=torch.float32) |
|
else: |
|
index_score = indexer_scores |
|
index_lse = torch.logsumexp(indexer_scores, dim=-1) |
|
|
|
k_unsqueeze_dim = 1 if is_thd else 2 |
|
dense_attn_kwargs = {} |
|
if is_thd: |
|
dense_attn_kwargs = dict( |
|
cu_seqlens_q=cu_seqlens_q, |
|
cu_seqlens_kv=cu_seqlens_compressed_idx, |
|
max_seqlen_q=int(max_seqlen_q), |
|
max_seqlen_kv=int(max_seqlen_compressed_idx), |
|
) |
|
attn_score, attn_l1norm = _compute_dense_attn_score( |
|
q_attn_det, |
|
k_attn_compressed_det.unsqueeze(k_unsqueeze_dim), |
|
lse_indexer_det, |
|
qhead_per_kv_head=np_, |
|
softmax_scale=softmax_scale, |
|
ratio=ratio, |
|
**dense_attn_kwargs, |
|
) |
FlashMLA defines lse_indexer as the LSE over the first indexer_topk entries only, while attn_sink does not affect lse/lse_indexer:
https://github.com/deepseek-ai/FlashMLA/blob/b7643bd54521f563b839b98289b5cd048c062ba2/flash_mla/flash_mla_interface.py#L176-L215
Therefore the fused target also appears to omit the sliding-window and sink competition from the teacher denominator. The primary bug reported here is the statically explicit unfused/small-operator path, but the fused path likely needs the same semantic correction.
@NVIDIA/mcore-oncall
Describe the bug
The DeepSeek-V4 unfused / small-operator CSA path computes the indexer-loss teacher by recomputing attention over compressed KV only. This teacher is not equivalent to the real CSA attention distribution described by the DeepSeek-V4 paper, because the real CSA core attention also includes the original-KV sliding window and the attention sink in a shared softmax denominator.
This is a semantic mismatch in the teacher distribution, not a floating-point precision difference.
Static analysis scope:
Related DSV4 tracking issue: #4468
Paper reference
DeepSeek-V3.2 defines the Lightning Indexer teacher from the real main-attention distribution: per-head attention probabilities are aggregated across heads and then L1-normalized (Section 2.1, equations (3) and (4)):
https://arxiv.org/html/2512.02556v1#S2.SS1.SSS1
DeepSeek-V4 keeps the Lightning Indexer objective while changing the CSA core attention. The CSA attention contains all of the following components:
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS1.Px2
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS3.Px3
https://arxiv.org/html/2606.19348v1#S2.SS3.SSS3.Px4
For a query and attention head h, let c(h,j) be a compressed-key logit, w(h,r) a window-key logit, and a(h) the sink logit. The compressed part of the paper-equivalent teacher must retain the real per-head CSA denominator:
Consequently, changing only the window logits or sink logits can change the compressed teacher distribution because it changes the relative contribution of each attention head.
Megatron-LM unfused / small-operator implementation
1. The loss receives compressed KV only
In the SBHD unfused path, key_for_loss is constructed only from compressed_kv and passed to FusedDSAIndexerLoss:
Megatron-LM/megatron/core/transformer/experimental_attention_variant/csa.py
Lines 1606 to 1640 in fd1121b
Relevant data flow:
Neither window KV/logits nor attn_sink is an input to this loss call.
The THD unfused path does the same with key_for_loss_thd:
Megatron-LM/megatron/core/transformer/experimental_attention_variant/csa.py
Lines 1927 to 2000 in fd1121b
2. The small operators recompute a compressed-only softmax
compute_dsa_indexer_loss calculates QK scores from the supplied query and key, applies the compressed causal/Top-k masks, softmaxes over that key axis, sums attention heads, and L1-normalizes:
Megatron-LM/megatron/core/transformer/experimental_attention_variant/dsa.py
Lines 275 to 353 in fd1121b
Therefore, ignoring masks for clarity, the implemented teacher is:
With sparse loss enabled, the Top-k mask is applied before this per-head softmax, so C above becomes the selected Top-k subset.
This differs from the paper teacher because the window and sink terms have been removed from every per-head denominator. It forces every head to contribute unit compressed probability mass before head aggregation, even when a head's real CSA probability mass is mostly assigned to its sliding window or sink.
3. The real unfused CSA attention uses a different denominator
Only after the indexer loss has already been computed, the implementation concatenates the window indices with the compressed Top-k and calls the actual sparse attention with attn_sink:
Megatron-LM/megatron/core/transformer/experimental_attention_variant/csa.py
Lines 1652 to 1673 in fd1121b
Thus the output attention and the indexer-loss teacher do not use the same attention distribution:
Static reproduction / invariant
Hold query and compressed_kv fixed, and perturb only either:
Expected from the paper equations:
Current unfused implementation:
This invariant violation follows directly from the function inputs and does not depend on a particular GPU kernel or numerical tolerance.
Additional Top-k concern
When dsa_indexer_topk is smaller than the number of visible compressed keys, sparse loss applies the Top-k mask before the per-head attention softmax. This is generally different from:
The current short/full-selection cases can hide this additional difference.
Expected behavior
The indexer-loss teacher should be derived from the real CSA attention probabilities, or from an equivalent recomputation that retains the same per-head denominator containing compressed attention, window attention, and sink.
For sparse loss, selected compressed entries should be derived from that teacher distribution and then normalized according to the Lightning Indexer objective, rather than rebuilding independent per-head softmax denominators over the selected compressed subset.
Additional context: fused path
The fused path appears to have the same semantic issue. It places compressed Top-k entries before window entries, obtains both lse and lse_indexer from FlashMLA, but passes lse_indexer to the target recomputation:
Megatron-LM/megatron/core/transformer/experimental_attention_variant/dsa_kernels.py
Lines 1135 to 1257 in fd1121b
FlashMLA defines lse_indexer as the LSE over the first indexer_topk entries only, while attn_sink does not affect lse/lse_indexer:
https://github.com/deepseek-ai/FlashMLA/blob/b7643bd54521f563b839b98289b5cd048c062ba2/flash_mla/flash_mla_interface.py#L176-L215
Therefore the fused target also appears to omit the sliding-window and sink competition from the teacher denominator. The primary bug reported here is the statically explicit unfused/small-operator path, but the fused path likely needs the same semantic correction.