Skip to content

Fix target head weight load and refactor checkpoint loading - #875

Open
cih9088 wants to merge 4 commits into
sgl-project:mainfrom
cih9088:feature/robust-weight-load
Open

cih9088 wants to merge 4 commits into
sgl-project:mainfrom
cih9088:feature/robust-weight-load

Conversation

@cih9088

@cih9088 cih9088 commented Sep 11, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

When loading the weights for TargetHead, the logic looks for model.safetensors.index.json and raises an error if it is not found. However model.safetensors.index.json exits only if the model weights are shards. (index creation, index saving in transformers).

Some of checkpoint loading logics implemented this feature and they are scattered around. I factored them out to a single utility function

Modifications

  • added unified checkpoint load function
  • replaced the checkpoint load logic with the unified load function

Related Issues

Closes #927

Accuracy Test

Benchmark & Profiling

Checklist

Targets with tie_word_embeddings=True (e.g. Qwen2.5-0.5B-Instruct) store
only model.embed_tokens.weight, so TargetHead failed with "Target head key
'lm_head.weight' is missing". When the head key is absent and the target
config (top level or text_config) is tied, load cfg.model.embedding_key
instead. A stored head still takes precedence and untied targets keep the
existing error.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@lru49

lru49 commented Oct 4, 2026 •

Copy link
Copy Markdown

@cih9088 thank you for this PR. I filed issue #927 because even with your fix, Qwen2.5-0.5B still fails because from_pretrained() is looking for lm_head weights but the model is using tied embeddings. It's such a small fix to have from_pretrained() reference the embedding weights instead of lm_head for tied embedding models, that I thought you might want to integrate it into this PR. I implemented the fix & confirmed that the test suite passes on my branch (except for 4 tests that I cannot run on my single GPU setup), you can see the fix here:

cih9088/SpecForge@feature/robust-weight-load...lru49:SpecForge:fix/target-head-tied-embeddings

Tell me if you think this makes more sense as a separate PR.

edit: my changes have been merged into this PR branch, so the diff link above no longer shows any changes

@cih9088

cih9088 commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor Author

@lru49
Oh, I have missed that. I don't mind that you stack up additional commits on top of mine. I gave you a write permission on my fork temporarily. Please feel free to add commits of yours or open a separate PR.

@lru49

lru49 commented Oct 6, 2026

Copy link
Copy Markdown

Thanks @cih9088 ! I just merged my changes

@lru49

lru49 commented Oct 7, 2026 •

Copy link
Copy Markdown

@cih9088 Have you posted this PR in sglang slack for review? That could help get eyes on it

@cih9088

cih9088 commented Oct 7, 2026

Copy link
Copy Markdown
Contributor Author

@cih9088 Have you posted this PR in sglang slack for review? That could help get eyes on it

@lru49 no, I didn't join the slack. If you happened to be in the slack, be my guest.

I have no idea who to ping. Could you take a look at this PR? @curnane-lab @jiapingW

cache_dir=cache_dir,
)
except KeyError as exc:
if not self._ties_word_embeddings():

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

For a tied checkpoint that still stores lm_head.weight, this prefers the stored lm_head while TargetEmbeddingsAndHead (and transformers, which re-ties after load) ignores it — worth picking one answer for both.

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

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

3 participants