mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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:
+2
-1
@@ -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.
|
||||
|
||||
@@ -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"),
|
||||
|
||||
Reference in New Issue
Block a user