review: offline logit reconstruction casts to the lm_head dtype; log calibration as a KL

Cast the captured final hidden state to the lm_head weight dtype before reconstructing target logits offline, cover the offline path with a forward test that matches the online loss, and report the calibration term as KL(p|K || q), whose gradient is identical to the cross-entropy's.
This commit is contained in:
mrusanovsky
2026-09-30 12:07:39 -07:00
parent dc1d3363ed
commit 50d697647e
5 changed files with 72 additions and 16 deletions
+2 -1
View File
@@ -30,7 +30,7 @@ Changelog
- Add the DFlash2 draft variant, selected with ``dflash_architecture_config.projector_type="dflash2"``: DFlash's one-pass parallel backbone plus a grouped dynamic convolution around every attention/MLP sublayer (``conv_kernel_size`` / ``conv_group_size``) and a low-rank candidate selector (``selector_rank`` / ``selector_top_k``, weighted by ``dflash_selector_loss_alpha``). Exported checkpoints declare ``DFlash2DraftModel`` and load in the SGLang/vLLM DFlash2 serving path.
- LiLiCorr now trains offline and in streaming mode, not only online: the target logits its distractor penalty reads are reconstructed from the captured final hidden state, the same path self-logit distillation uses.
- Add ``dflash_lilicorr_w_cal`` (default ``0.0``, off): an optional LiLiCorr proposal-calibration term, a cross-entropy from the target's distribution over each slot's candidates onto the reranker's candidate distribution, usable in place of the distractor penalty.
- Add ``dflash_lilicorr_w_cal`` (default ``0.0``, off): an optional LiLiCorr proposal-calibration term, a KL divergence from the target's distribution over each slot's candidates to the reranker's candidate distribution, usable in place of the distractor penalty.
*Megatron Framework (M-LM / M-Bridge)*
@@ -91,6 +91,7 @@ Changelog
**Bug Fixes**
- Fix DFlash conversion on NoPE targets whose config leaves ``rope_theta`` unset.
- Fix offline DFlash training failing to reconstruct the target logits when the captured hidden states are stored in a different dtype than the target's weights.
- Fix Megatron-Core checkpoint saving for quantized grouped MoE experts when tensor and expert parallelism are both enabled.
- Fix shared ONNX export metadata and Diffusers attention policy: every ``NVFP4QuantExporter`` post-process now upgrades the default-domain opset to at least 23, all FP8 custom-op exports re-run ONNX shape/type inference after setting output metadata, and quantized SDPA derives FP8 MHA enablement from the live Q/K/V quantizers instead of honoring a caller-set ``_disable_fp8_mha`` attribute.
- Fix ONNX FP16 conversion failing to preserve public output types when type inference changes a graph output declaration before output casts are inserted.
+2 -2
View File
@@ -382,8 +382,8 @@ class DFlashConfig(ModeloptBaseConfig):
ge=0.0,
allow_inf_nan=False,
description=(
"LiLiCorr only: absolute weight of the optional calibration term, a cross-entropy "
"from the target's distribution renormalized over each slot's k candidates onto "
"LiLiCorr only: absolute weight of the optional calibration term, a KL divergence "
"from the target's distribution renormalized over each slot's k candidates to "
"the reranker's. An alternative to w_pen whose gradient does not scale with the "
"target's logit gap. Requires the target's logits, like w_pen. 0.0 (default) is a "
"no-op, and it is not part of the three weights' all-or-nothing validation. "
@@ -59,8 +59,8 @@ The objective adds weighted terms to the DFlash loss::
- ``L_pen``: the head's own probability mass on the competing candidates, each
weighted by the target model's logit gap to the ground truth — so a candidate the
target finds plausible is penalized lightly and a confident wrong one hard.
- ``L_cal`` (optional, off by default): cross-entropy from the target's distribution
renormalized over the ``k`` candidates onto the head's. An alternative to ``L_pen``
- ``L_cal`` (optional, off by default): KL divergence from the target's distribution
renormalized over the ``k`` candidates to the head's. An alternative to ``L_pen``
whose gradient does not scale with the target's logit gap.
The weights are absolute, with no outer multiplier, so
@@ -576,10 +576,12 @@ class HFLiLiCorrModel(HFDFlashModel):
supervised,
denominator,
):
"""Cross-entropy ``H(p|K, q)`` from the target's candidate distribution to the head's.
"""``KL(p|K || q)`` from the target's candidate distribution to the head's.
``p|K`` is the target renormalized over the ``k`` candidates, so the gathered
candidate logits suffice and the full-vocabulary normalizer is not needed.
candidate logits suffice and the full-vocabulary normalizer is not needed. The
gradient is the cross-entropy's, and the logged value floors at 0 rather than at
``H(p|K)``, which varies with the data and the temperature.
"""
candidate_target_logits = self._candidate_target_logits(
candidate_ids=candidate_ids,
@@ -588,10 +590,11 @@ class HFLiLiCorrModel(HFDFlashModel):
requested_by="dflash_lilicorr_w_cal",
)
target_probs = F.softmax(candidate_target_logits, dim=-1)
target_log_probs = F.log_softmax(candidate_target_logits, dim=-1)
potentials = torch.stack(node_potentials, dim=1)
head_log_probs = F.log_softmax(potentials.float(), dim=-1)
per_slot = -(target_probs * head_log_probs).sum(dim=-1)
per_slot = (target_probs * (target_log_probs - head_log_probs)).sum(dim=-1)
return (per_slot * supervised).sum() / denominator
@torch.no_grad()
@@ -149,7 +149,8 @@ class DFlashBaseModelOutput:
if out_hiddens is None:
raise KeyError("base_model_hidden_states")
out_hiddens = _maybe_apply_base_final_norm(out_hiddens, d, base_model_norm)
logits = base_model_lm_head(out_hiddens)
# A producer can store the hidden states in a wider dtype than the target's weights.
logits = base_model_lm_head(out_hiddens.to(base_model_lm_head.weight.dtype))
return cls(
target_hidden=d["aux_hidden_states"],
logits=logits,
@@ -336,7 +336,7 @@ class TestLiLiCorrForward:
)
def test_calibration_loss_is_minimal_when_the_head_matches_the_target(self):
"""Zero gradient when the reranker's candidate distribution already matches the target's."""
"""Zero loss and gradient when the reranker's candidate distribution matches the target's."""
model = _converted(dflash_lilicorr_w_cal=0.5)
torch.manual_seed(0)
num_slots = BLOCK_SIZE - 1
@@ -351,20 +351,71 @@ class TestLiLiCorrForward:
}
at_target = model._candidate_target_logits(**inputs, requested_by="test")
def grad_at(potentials):
def loss_and_grad_at(potentials):
potentials = potentials.clone().requires_grad_(True)
supervised = torch.ones(2, num_slots)
model._calibration_loss(
loss = model._calibration_loss(
node_potentials=list(potentials.unbind(1)),
supervised=supervised,
denominator=supervised.sum(),
**inputs,
).backward()
return potentials.grad
)
loss.backward()
return loss.item(), potentials.grad
# A per-slot shift leaves the softmax, and so the optimum, unchanged.
assert grad_at(at_target + 3.0).abs().max() < 1e-6
assert grad_at(torch.zeros_like(at_target)).abs().max() > 1e-3
loss, grad = loss_and_grad_at(at_target + 3.0)
assert loss == pytest.approx(0.0, abs=1e-6)
assert grad.abs().max() < 1e-6
loss, grad = loss_and_grad_at(torch.zeros_like(at_target))
assert loss > 1e-3
assert grad.abs().max() > 1e-3
def test_offline_forward_matches_online(self):
"""Target logits rebuilt from a captured hidden state give the online loss.
The hidden state is stored in float64, wider than the target's weights.
"""
models = []
for offline in (False, True):
model = get_tiny_llama(num_hidden_layers=4)
model.config.num_orig_hidden_layers = model.config.num_hidden_layers
config = _get_lilicorr_config(
topk=VOCAB_SIZE, dflash_lilicorr_w_cal=0.5, dflash_offline=offline
)
mtsp.convert(model, [("dflash", config)])
models.append(model)
online, offline = models
assert not offline.load_state_dict(online.state_dict(), strict=False).missing_keys
batch = _make_batch(VOCAB_SIZE)
online.eval()
with torch.no_grad():
hidden_states = online(
input_ids=batch["input_ids"],
attention_mask=batch["attention_mask"],
output_hidden_states=True,
).hidden_states
base_model_outputs = {
"aux_hidden_states": torch.cat(
[hidden_states[lid + 1] for lid in online.target_layer_ids], dim=-1
),
"base_model_hidden_states": hidden_states[-1].double(),
}
online.train()
offline.train()
torch.manual_seed(1)
online_out = online(**batch)
torch.manual_seed(1)
offline_out = offline(**batch, base_model_outputs=base_model_outputs)
assert offline_out.loss.item() == pytest.approx(online_out.loss.item(), abs=1e-5)
for term in ("lilicorr_penalty", "lilicorr_calibration"):
assert online_out.lilicorr_metrics[term] > 0.0
assert offline_out.lilicorr_metrics[term] == pytest.approx(
online_out.lilicorr_metrics[term], abs=1e-5
)
@pytest.mark.parametrize(
("kwargs", "degenerate"),