[fix]: respect auto_model param in FastLlamaModel.from_pretrained for embedding models - #6886
InfoSage05 wants to merge 5 commits into
Conversation
FastLlamaModel.from_pretrained() hardcoded AutoModelForCausalLM when loading models, ignoring any `auto_model` parameter passed through kwargs. This broke FastSentenceTransformer for decoder-style embedding models (Qwen3-Embedding, etc.) because the sentence-transformer integration already correctly sets `kwargs["auto_model"] = AutoModel` to request the base model class without LM head. When AutoModelForCausalLM was always force-loaded, the sentence-transformers Transformer module received a Qwen3ForCausalLM instead of Qwen3Model. The CausalLM forward returns CausalLMOutputWithPast (with logits) instead of BaseModelOutputWithPast (with last_hidden_state), causing structurally different hidden states that degraded embedding quality by ~14 recall points on retrieval benchmarks. The fix pops `auto_model` from kwargs and uses it as the model class, defaulting to AutoModelForCausalLM for backward compatibility. Fixes unslothai#6881
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
There was a problem hiding this comment.
Code Review
This pull request introduces the ability to dynamically specify the model class used for loading by popping an auto_model argument from kwargs, defaulting to AutoModelForCausalLM. The review feedback correctly identifies a potential AttributeError if auto_model is explicitly passed as None, and provides a robust code suggestion to handle this case safely.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
|
Root cause matches everything we measured, the swallowed auto_model explains the missing final-norm module and the pre-norm-looking hidden states (we'd traced the fused epilogue in the compiler). I'll validate the branch against our recall benchmark: plain-ST baseline is 54.3 @50; if the branch load scores ≈54.3 I'll confirm here. One question: do the 4-bit / fast_inference load branches also hardcode AutoModelForCausalLM, or is this the only swallow site? |
|
Validated this branch ( Results:
So this PR fixes one real bug, but a second divergence source remains in the forward path. For what it's worth, on the pre-fix load we also tried Happy to re-run this gate on any follow-up commit. |
…sion The Qwen3 attention forward uses `fast_rms_layernorm` (Triton kernel) for QK normalization in the training/eval path. While functionally equivalent to the stock HuggingFace RMSNorm, tiny kernel-level differences compound across 28+ layers because the QK norm operates on narrow per-head tensors (head_dim ~96) before the attention softmax — small perturbations there amplify. Gate the QK norm in Qwen3Attention_fast_forward with `UNSLOTH_EMBEDDING_HIGH_PRECISION`: when set, use the pure-PyTorch `fast_rms_layernorm_inference` variant instead of the Triton kernel. The inference path already uses the PyTorch variant. Set this flag in `FastSentenceTransformer.from_pretrained()` so every embedding-model load through that entrypoint automatically gets bit-exact numerical parity with stock HF. Follow-up to unslothai#6886 — addresses the remaining ~0.09 cosine gap after the auto_model structural fix. See unslothai#6881
0c50077 to
169e603
Compare
|
@Afnisse Thanks for validating on T4 which only that confirms the structural fix was real but only part of the picture. I've pushed a follow-up commit (4006f07) that addresses what I believe is the remaining divergence source. Root cause #2: In (The per-layer input/post-attention norms and final norm also use Triton, but they operate on full hidden_size (1024-dim) where precision loss is negligible.) The fix:When |
|
@codex review |
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: 169e603497
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
…ly paths Two changes: 1. Move UNSLOTH_EMBEDDING_HIGH_PRECISION=1 outside the finally block in FastSentenceTransformer.from_pretrained(). The flag was being restored to "0" before encode() was ever called, neutering the QK-norm fix. 2. Guard two code paths in FastLlamaModel.from_pretrained() that assume a ForCausalLM wrapper (model.generate, model.model.layers). When auto_model=AutoModel loads a bare Qwen3Model/LlamaModel, these attributes don't exist. Use hasattr guards to skip or route correctly. Fixes: unslothai#6881 (follow-up)
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
|
I am afraid that this does not the user issue really from my testings , still happen on colab even with this pr on |
Problem
Fixes #6881 —
FastSentenceTransformer.from_pretrained()produces wrong embeddings for decoder-style embedding models (e.g.Qwen/Qwen3-Embedding-0.6B), costing ~14 recall points on retrieval benchmarks vs plainSentenceTransformer.SentenceTransformer('Qwen/Qwen3-Embedding-0.6B')FastSentenceTransformer(same checkpoint)No error or warning is emitted where the model loads, runs, and returns embeddings, just materially worse ones. The wrapped backbone (
model[0].auto_model) had no final RMSNorm module in its tree, and hidden states diverged well before the final norm.Root Cause
FastLlamaModel.from_pretrained()hardcodedAutoModelForCausalLMwhen loading models, ignoring anyauto_modelparameter passed through**kwargs.The sentence-transformer integration (
FastSentenceTransformer.from_pretrained()) already correctly setskwargs["auto_model"] = AutoModelto request the base model class (without LM head), but this parameter was swallowed by**kwargsand never used.When
AutoModelForCausalLMwas always force-loaded:Transformermodule received aQwen3ForCausalLMinstead ofQwen3ModelCausalLMOutputWithPast(withlogits) instead ofBaseModelOutputWithPast(withlast_hidden_state)Fix
One-line logic change: pop
auto_modelfromkwargsand use it as the model class, defaulting toAutoModelForCausalLMfor backward compatibility.The change affects two
from_pretrainedcalls (one withuser_config, one without) in thenot fast_inferencebranch ofFastLlamaModel.from_pretrained().All model-specific
from_pretrainedmethods (Qwen2, Qwen3, Mistral, GLM4-MoE, FalconH1, etc.) delegate toFastLlamaModel.from_pretrained(), so this single fix covers them all. TheFastModel.from_pretrained()path inloader.pyalready handlesauto_modelcorrectly.Testing / Verification
AutoModelForCausalLM)FastSentenceTransformerpath now correctly receives the base model (e.g.Qwen3Model) whenauto_model=AutoModelis passed, which returnsBaseModelOutputWithPastcontaininglast_hidden_stateReported-by: Issue #6881