Skip to content

[Bug] TargetHead fails for targets with tied word embeddings (e.g. Qwen2.5-0.5B) because lm_head.weight is missing #927

Description

@lru49

Checklist

  • 1. I have searched related issues but cannot get the expected help.
  • 2. The bug has not been fixed in the latest version.
  • 3. Please note that if the bug-related issue you submitted lacks corresponding environment info and a minimal reproducible demo, it will be challenging for us to reproduce and resolve the issue, reducing the likelihood of receiving feedback.
  • 4. If the issue you raised is not a bug but a question, please raise a discussion at https://github.com/sgl-project/SpecForge/discussions/new/choose Otherwise, it will be closed.
  • 5. Please use English, otherwise it will be closed.

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

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

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions