Checklist
Describe the bug
TargetHead (the frozen target LM head used by offline EAGLE3 training and by
online disaggregated consumers) always loads cfg.model.lm_head_key, which
defaults to "lm_head.weight". Targets with tie_word_embeddings: true don't
store that tensor: the checkpoint only contains model.embed_tokens.weight,
and transformers reuses it as the output head at load time. TargetHead
never checks the tie flag, so loading fails for any tied target unless the user
knows to set model.lm_head_key=model.embed_tokens.weight by hand.
TargetEmbeddingsAndHead in target_utils.py already handles this case
(tie_word_embeddings → load only the embedding and share it), so the two
target loaders disagree. The examples already contain a manual workaround for
a tied target (qwen3.5-4b-mtp-disaggregated-npu.yaml sets
lm_head_key: "model.language_model.embed_tokens.weight").
Additional note: the fix for this might integrate nicely into PR #875 but I could also make a separate PR.
Reproduction
I happened to reproduce this issue on #875's branch (17f86f2). #875 fixed the issue that prevented Qwen2.5-0.5B-Instruct's single-file checkpoint from loading, but because the model has tied embeddings, there was still a failure due to lm_head.weight not being found:
git fetch origin pull/875/head:pr-875 && git checkout pr-875
python -c "from specforge.modeling.target.target_head import TargetHead; TargetHead.from_pretrained('Qwen/Qwen2.5-0.5B-Instruct')"
File "specforge/modeling/target/target_head.py", line 71, in load_weights
raise RuntimeError(
RuntimeError: Target head key 'lm_head.weight' is missing from Qwen/Qwen2.5-0.5B-Instruct
Qwen2.5-0.5B-Instruct's config.json has "tie_word_embeddings": true, and
its model.safetensors has no lm_head.weight.
Workaround: model.lm_head_key=model.embed_tokens.weight.
Proposed fix: in TargetHead.load_weights, if lm_head_key is missing and
the target config (top level or text_config) has tie_word_embeddings=True,
load cfg.model.embedding_key instead.
Environment
AI acknowledgement
This issue was written with the assistance of Claude Code
Checklist
Describe the bug
TargetHead(the frozen target LM head used by offline EAGLE3 training and byonline disaggregated consumers) always loads
cfg.model.lm_head_key, whichdefaults to
"lm_head.weight". Targets withtie_word_embeddings: truedon'tstore that tensor: the checkpoint only contains
model.embed_tokens.weight,and
transformersreuses it as the output head at load time.TargetHeadnever checks the tie flag, so loading fails for any tied target unless the user
knows to set
model.lm_head_key=model.embed_tokens.weightby hand.TargetEmbeddingsAndHeadintarget_utils.pyalready handles this case(
tie_word_embeddings→ load only the embedding and share it), so the twotarget loaders disagree. The examples already contain a manual workaround for
a tied target (
qwen3.5-4b-mtp-disaggregated-npu.yamlsetslm_head_key: "model.language_model.embed_tokens.weight").Additional note: the fix for this might integrate nicely into PR #875 but I could also make a separate PR.
Reproduction
I happened to reproduce this issue on #875's branch (
17f86f2). #875 fixed the issue that prevented Qwen2.5-0.5B-Instruct's single-file checkpoint from loading, but because the model has tied embeddings, there was still a failure due tolm_head.weightnot being found:Qwen2.5-0.5B-Instruct's
config.jsonhas"tie_word_embeddings": true, andits
model.safetensorshas nolm_head.weight.Workaround:
model.lm_head_key=model.embed_tokens.weight.Proposed fix: in
TargetHead.load_weights, iflm_head_keyis missing andthe target config (top level or
text_config) hastie_word_embeddings=True,load
cfg.model.embedding_keyinstead.Environment
main@53398a8, plus Fix target head weight load and refactor checkpoint loading #875 @17f86f2flashinfer 0.6.17
AI acknowledgement
This issue was written with the assistance of Claude Code