Skip to content

Restore the training loss for the DeepSeek V4 indexer #4802

Description

@HaoranYu111

The TODO in DSV4 attention still says that indexer auxiliary loss was temporarily removed until the general aux-loss mechanism landed. That mechanism is now available through #3864.

On main 6c2dadb, the indexer produces scores, but Indexer.select() only returns integer top-k indices. Those indices do not carry gradients back to the indexer's query/key/weight projections. The main attention loss therefore does not replace the missing auxiliary training path.

A CPU check of the current Indexer.select() confirms that a loss on the selected values gives gradients to those values, while idx_q, idx_k and idx_w retain grad=None. This is a function-level check on PyTorch 2.12, not a full-model training result. The loss is intentionally absent, so this is a request to restore functionality rather than a newly introduced regression.

Is someone already restoring this alongside #4452? A useful scope would be:

  • Inject the indexer loss through the existing auxiliary-loss mechanism.
  • Gather teacher KV by row selection, avoiding the expanded four-dimensional gather intermediate.
  • Test teacher detach, repeated/invalid selections, TP/CP normalization and gradients, and activation-checkpoint recomputation.

There is an implementation and local validation from our NPU adaptation, but it still needs migration to the current flattened-token/SPMD layout and representative training validation. #3800 has related indexer work for DeepSeek V3.2; it does not restore the V4 path.

CPU check of the top-k gradient path
import torch
from torchtitan.models.deepseek_v4.compressor import Indexer

torch.manual_seed(42)
idx_q = torch.randn(16, 2, 8, requires_grad=True)
idx_k = torch.randn(4, 8, requires_grad=True)
idx_w = torch.randn(16, 2, requires_grad=True)
indices = Indexer.select(idx_q, idx_k, idx_w, seqlen=16, ratio=4, topk=2)
values = torch.randn(4, 8, requires_grad=True)
values[indices].square().sum().backward()
assert indices.dtype == torch.int64 and not indices.requires_grad
assert all(t.grad is None for t in (idx_q, idx_k, idx_w))
assert values.grad is not None
print([t.grad for t in (idx_q, idx_k, idx_w)])  # [None, None, None]

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions