Skip to content

[fix]: respect auto_model param in FastLlamaModel.from_pretrained for embedding models - #6886

Closed
InfoSage05 wants to merge 5 commits into
unslothai:mainfrom
InfoSage05:fix/6881-embedding-auto-model
Closed

InfoSage05 wants to merge 5 commits into
unslothai:mainfrom
InfoSage05:fix/6881-embedding-auto-model

Conversation

@InfoSage05

@InfoSage05 InfoSage05 commented Jul 5, 2026 •

Copy link
Copy Markdown
Contributor

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 plain SentenceTransformer.

Load path recall@10 @30 @50 @100 @200 never-in-top-200
SentenceTransformer('Qwen/Qwen3-Embedding-0.6B') 35.2 49.4 54.3 62.9 74.9 25.1%
FastSentenceTransformer (same checkpoint) 22.8–23.6 33.0–34.5 39.3–40.4 47.6–49.4 57.3–60.3 ~40–42%

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() hardcoded AutoModelForCausalLM when loading models, ignoring any auto_model parameter passed through **kwargs.

The sentence-transformer integration (FastSentenceTransformer.from_pretrained()) already correctly sets kwargs["auto_model"] = AutoModel to request the base model class (without LM head), but this parameter was swallowed by **kwargs and never used.

When AutoModelForCausalLM was always force-loaded:

  1. The sentence-transformers Transformer module received a Qwen3ForCausalLM instead of Qwen3Model
  2. The CausalLM forward returns CausalLMOutputWithPast (with logits) instead of BaseModelOutputWithPast (with last_hidden_state)
  3. This caused structurally different hidden states → degraded embeddings

Fix

One-line logic change: pop auto_model from kwargs and use it as the model class, defaulting to AutoModelForCausalLM for backward compatibility.

auto_model_class = kwargs.pop("auto_model", AutoModelForCausalLM)
model = auto_model_class.from_pretrained(...)  # was: AutoModelForCausalLM.from_pretrained(...)

The change affects two from_pretrained calls (one with user_config, one without) in the not fast_inference branch of FastLlamaModel.from_pretrained().

All model-specific from_pretrained methods (Qwen2, Qwen3, Mistral, GLM4-MoE, FalconH1, etc.) delegate to FastLlamaModel.from_pretrained(), so this single fix covers them all. The FastModel.from_pretrained() path in loader.py already handles auto_model correctly.

Testing / Verification

  • Syntax/compile check passes
  • No behavioral change for standard LLM loading (defaults to AutoModelForCausalLM)
  • The FastSentenceTransformer path now correctly receives the base model (e.g. Qwen3Model) when auto_model=AutoModel is passed, which returns BaseModelOutputWithPast containing last_hidden_state

Reported-by: Issue #6881

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
@chatgpt-codex-connector

Copy link
Copy Markdown

Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits.
Repo admins can enable using credits for code reviews in their settings.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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.

Comment thread unsloth/models/llama.py Outdated
@Afnisse

Afnisse commented Jul 5, 2026

Copy link
Copy Markdown

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?

@Afnisse

Afnisse commented Jul 5, 2026 •

Copy link
Copy Markdown

Validated this branch (0c50077) against the benchmark from #6881 - Colab T4 GPU runtime, plain-ST control re-run in the same session, identical eval code and data.

Results:

  • ✅ Structural change confirmed: model[0].auto_model now has a top-level final norm module (before this PR it was missing entirely), consistent with the base model loading instead of the CausalLM - so the auto_model swallow was real and this PR addresses it.
  • ⚠️ Single-text full-vector cosine vs the plain-ST embedding of the same string: 0.912 - still far from parity (same-weights dtype jitter would be ≥0.999).
  • ❌ Retrieval parity is NOT restored: recall@50 = 39.7 vs 54.3 control, statistically identical to the pre-fix load (which scored 39.3–40.4 across configs).
Load path @10 @30 @50 @100 @200 never-in-top-200
plain SentenceTransformer (control, same session) 34.8 48.7 54.3 62.9 75.3 24.7%
FastSentenceTransformer @ this PR 23.2 32.6 39.7 49.1 58.8 41.2%

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 UNSLOTH_HIGH_PRECISION_LAYERNORM=1 + UNSLOTH_FORCE_FLOAT32=1 + dtype=torch.float32 with no effect, so the remaining gap doesn't look like simple kernel precision — though I haven't re-run those flags on top of this branch.

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
@InfoSage05
InfoSage05 force-pushed the fix/6881-embedding-auto-model branch from 0c50077 to 169e603 Compare July 5, 2026 20:23
@InfoSage05
InfoSage05 requested a review from Etherll as a code owner July 5, 2026 20:23
@InfoSage05

Copy link
Copy Markdown
Contributor Author

@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 Qwen3Attention_fast_forward, the QK normalization uses fast_rms_layernorm which is a Triton autograd kernel in the training/eval path. While mathematically equivalent to the stock HF RMSNorm, the QK norm operates on per-head tensors (head_dim = 96 for Qwen3-0.6B). At such narrow widths, even single-ULP differences in the Triton reduction amplify through the attention softmax, then cascade through all 28 decoder layers.

(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 UNSLOTH_EMBEDDING_HIGH_PRECISION is set, Qwen3Attention_fast_forward uses fast_rms_layernorm_inference (pure PyTorch, bit-exact to HF) for QK norm instead of the Triton kernel. This flag is automatically set by FastSentenceTransformer.from_pretrained().
Could you re-run the recall benchmark on 4006f07 once ? If cosine still isn't ≥0.999, the next suspects would be the attention dispatch path (SDPA backend selection) or the RoPE kernel.

@danielhanchen

Copy link
Copy Markdown
Member

@codex review

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 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".

Comment thread unsloth/models/llama.py Outdated
Comment thread unsloth/models/sentence_transformer.py Outdated
Ayushman Paul added 2 commits July 7, 2026 04:09
…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)
@chatgpt-codex-connector

Copy link
Copy Markdown

Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits.
Repo admins can enable using credits for code reviews in their settings.

@chatgpt-codex-connector

Copy link
Copy Markdown

Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits.
Repo admins can enable using credits for code reviews in their settings.

@Etherll

Etherll commented Jul 7, 2026

Copy link
Copy Markdown
Collaborator

I am afraid that this does not the user issue really from my testings , still happen on colab even with this pr on

@Etherll Etherll self-assigned this Jul 7, 2026
@Etherll

Etherll commented Jul 7, 2026

Copy link
Copy Markdown
Collaborator

#6939

@Etherll Etherll closed this Jul 10, 2026
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] FastSentenceTransformer silently degrades Qwen3-Embedding quality (recall@50: 54.3 → 39.7 on our eval)

4 participants