From 6974e4d37b177651dde05ebbe3085231f9bdf66b Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 16 Mar 2026 22:54:36 +0000 Subject: [PATCH 1/2] fix: add FastQwen3_5Model with fused CE loss for Qwen3.5 OOM (#4188) Qwen3.5 has a 248,320-token vocabulary. At 8K context the full logits tensor is 8192 x 248320 x 4 = 7.68 GB, which causes OOM on T4/P100. The unsloth compiler already applies fused CE via apply_fused_lm_head, but this adds an explicit FastQwen3_5Model dispatch for cleaner routing and better error messages when Qwen3.5 is not supported. Changes: - Add unsloth/models/qwen3_5.py with FastQwen3_5Model that patches Qwen3_5ForConditionalGeneration and Qwen3_5ForCausalLM forwards to use unsloth_fused_ce_loss directly from hidden_states - Add loader dispatch for model_type == "qwen3_5" before "qwen3" - Version gate uses >= 5.0.0 (qwen3_5 only exists in transformers 5.x) - Guarded import in loader.py with try/except fallback - GDN layers intentionally left unpatched (flash-linear-attention) - 23 unit tests covering all 4 code paths Fixes from original PR #4331 by @vitalis: - Add explicit _get_dtype import (wildcard import skips _-prefixed names) - Single-token fast path now checks labels is None before returning early - Default model name corrected to Qwen/Qwen3.5-9B (8B does not exist) - Test assertion on nn.Linear removed (not a mock) - Unused imports removed Tested: Qwen3.5-0.8B 4bit training, 1.38 GB peak memory, 23/23 tests pass. Backwards compatible: import unsloth works on transformers 4.57.6. --- tests/utils/test_qwen3_5.py | 623 ++++++++++++++++++++++++++++++++++++ unsloth/models/__init__.py | 5 + unsloth/models/loader.py | 16 + unsloth/models/qwen3_5.py | 335 +++++++++++++++++++ 4 files changed, 979 insertions(+) create mode 100644 tests/utils/test_qwen3_5.py create mode 100644 unsloth/models/qwen3_5.py diff --git a/tests/utils/test_qwen3_5.py b/tests/utils/test_qwen3_5.py new file mode 100644 index 0000000000..2e7d78d431 --- /dev/null +++ b/tests/utils/test_qwen3_5.py @@ -0,0 +1,623 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Unit tests for unsloth/models/qwen3_5.py — fix for issue #4188. + +These tests use CPU tensors and mock out GPU-only dependencies (unsloth_fused_ce_loss, +EMPTY_LOGITS) so they run without a CUDA device or real model weights. +""" + +import os +import types +import unittest +from unittest.mock import MagicMock, patch, call + +import pytest +import torch +import torch.nn as nn + + +# --------------------------------------------------------------------------- +# Helpers to import the module under test with mocked unsloth internals +# --------------------------------------------------------------------------- + + +def _make_fake_unsloth_fused_ce_loss(): + """Return a mock that records calls and returns a scalar loss tensor.""" + mock = MagicMock(return_value = torch.tensor(1.23)) + return mock + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +HIDDEN_DIM = 16 +VOCAB_SIZE = 64 + + +def _make_self(bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM, vocab_size = VOCAB_SIZE): + """ + Build a minimal mock `self` (the model instance) with: + - lm_head: a real nn.Linear (CPU) so the matmul paths work + - loss_function: a MagicMock returning a fixed scalar tensor + - accelerator_scaler: None + - config.text_config.vocab_size / config.vocab_size: vocab_size + """ + lm_head = nn.Linear(hidden_dim, vocab_size, bias = False) + + cfg_text = MagicMock() + cfg_text.vocab_size = vocab_size + cfg = MagicMock() + cfg.vocab_size = vocab_size + cfg.text_config = cfg_text + + self = MagicMock() + self.lm_head = lm_head + self.config = cfg + self.accelerator_scaler = None + self.loss_function = MagicMock(return_value = torch.tensor(0.99)) + return self + + +def _make_outputs(bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): + """Return a mock outputs object whose [0] is a random hidden_states tensor.""" + hidden = torch.randn(bsz, q_len, hidden_dim) + outputs = MagicMock() + outputs.__getitem__ = lambda self, idx: hidden if idx == 0 else None + outputs.past_key_values = None + outputs.hidden_states = None + outputs.attentions = None + return outputs, hidden + + +# --------------------------------------------------------------------------- +# Tests for _qwen3_5_compute_loss_or_logits +# --------------------------------------------------------------------------- + + +class TestComputeLossOrLogits(unittest.TestCase): + """Tests for the shared helper that houses all four forward paths.""" + + def setUp(self): + import unsloth.models.qwen3_5 as mod + + self.mod = mod + self.orig_fused_ce = mod.unsloth_fused_ce_loss + self.orig_empty_logits = mod.EMPTY_LOGITS + # Install fresh mocks for each test + self.mock_fused_ce = MagicMock(return_value = torch.tensor(1.23)) + self.mock_empty_logits = torch.zeros(1) + mod.unsloth_fused_ce_loss = self.mock_fused_ce + mod.EMPTY_LOGITS = self.mock_empty_logits + + def tearDown(self): + self.mod.unsloth_fused_ce_loss = self.orig_fused_ce + self.mod.EMPTY_LOGITS = self.orig_empty_logits + + # -- single-token decode path ------------------------------------------- + + def test_single_token_decode_uses_mv_not_lm_head(self): + """bsz=1, q_len=1 → fast torch.mv, no full lm_head call, no loss.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self(bsz = 1, q_len = 1) + hidden = torch.randn(1, 1, HIDDEN_DIM) + + with patch.object(torch, "mv", wraps = torch.mv) as mv_spy: + loss, logits = helper( + self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE + ) + + self.assertIsNone(loss) + self.assertEqual(logits.shape, (1, 1, VOCAB_SIZE)) + self.mock_fused_ce.assert_not_called() + self_.loss_function.assert_not_called() + + # -- partial-logits path ------------------------------------------------ + + def test_logits_to_keep_slices_last_n_tokens(self): + """logits_to_keep=3 → output has exactly 3 token positions.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + + loss, logits = helper( + self_, hidden, labels = None, logits_to_keep = 3, vocab_size = VOCAB_SIZE + ) + + self.assertIsNone(loss) + self.assertEqual(logits.shape, (2, 3, VOCAB_SIZE)) + self.mock_fused_ce.assert_not_called() + + # -- training / fused-CE path ------------------------------------------- + + def test_training_path_calls_fused_ce_returns_early(self): + """labels present + UNSLOTH_RETURN_LOGITS unset → fused CE, lm_head NOT called.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + loss, logits = helper( + self_, hidden, labels = labels, logits_to_keep = 0, vocab_size = VOCAB_SIZE + ) + + self.assertIsNotNone(loss) + self.assertIs(logits, self.mock_empty_logits) + self.mock_fused_ce.assert_called_once() + # Verify key arguments passed to fused CE + _, kwargs = self.mock_fused_ce.call_args + self.assertIs(kwargs["lm_head_weight"], self_.lm_head.weight) + self.assertEqual(kwargs["logit_softcapping"], 0) + self_.loss_function.assert_not_called() + + def test_training_path_passes_num_items_in_batch(self): + """num_items_in_batch kwarg is forwarded to unsloth_fused_ce_loss.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + helper( + self_, + hidden, + labels = labels, + logits_to_keep = 0, + vocab_size = VOCAB_SIZE, + num_items_in_batch = 16, + ) + + _, kwargs = self.mock_fused_ce.call_args + self.assertEqual(kwargs["n_items"], 16) + + # -- UNSLOTH_RETURN_LOGITS override ------------------------------------ + + def test_return_logits_env_var_bypasses_fused_ce(self): + """UNSLOTH_RETURN_LOGITS=1 → materialise full logits even during training.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): + loss, logits = helper( + self_, hidden, labels = labels, logits_to_keep = 0, vocab_size = VOCAB_SIZE + ) + + self.assertEqual(logits.shape, (2, 8, VOCAB_SIZE)) + self.mock_fused_ce.assert_not_called() + self_.loss_function.assert_called_once() + + # -- eval / inference path (no labels) ---------------------------------- + + def test_no_labels_returns_full_logits_no_loss(self): + """No labels → full logits computed, loss=None.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + loss, logits = helper( + self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE + ) + + self.assertIsNone(loss) + self.assertEqual(logits.shape, (2, 8, VOCAB_SIZE)) + self.mock_fused_ce.assert_not_called() + self_.loss_function.assert_not_called() + + # -- n_items fallback --------------------------------------------------- + + def test_n_items_kwarg_used_when_num_items_absent(self): + """n_items kwarg is the fallback when num_items_in_batch is absent.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + helper( + self_, + hidden, + labels = labels, + logits_to_keep = 0, + vocab_size = VOCAB_SIZE, + n_items = 8, + ) + + _, kwargs = self.mock_fused_ce.call_args + self.assertEqual(kwargs["n_items"], 8) + + def test_num_items_in_batch_zero_does_not_fall_through_to_n_items(self): + """num_items_in_batch=0 must NOT fall through to n_items (or-bug regression).""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + helper( + self_, + hidden, + labels = labels, + logits_to_keep = 0, + vocab_size = VOCAB_SIZE, + num_items_in_batch = 0, + n_items = 99, + ) + + _, kwargs = self.mock_fused_ce.call_args + # 0 is falsy; the old `or` expression would have returned 99 — must be 0. + self.assertEqual(kwargs["n_items"], 0) + + # -- batch decode (bsz > 1, q_len == 1) --------------------------------- + + def test_batch_decode_falls_through_to_eval_path(self): + """bsz=4, q_len=1 must NOT use the single-token mv path; goes to eval path.""" + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self(bsz = 4, q_len = 1) + hidden = torch.randn(4, 1, HIDDEN_DIM) + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): + with patch.object(torch, "mv", wraps = torch.mv) as mv_spy: + loss, logits = helper( + self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE + ) + + mv_spy.assert_not_called() # fast path is bsz==1 AND q_len==1 only + self.assertEqual(logits.shape, (4, 1, VOCAB_SIZE)) + + # -- labels silently ignored when logits_to_keep is set ----------------- + + def test_labels_ignored_when_logits_to_keep_nonzero(self): + """ + When logits_to_keep != 0 the function returns early with partial logits + and no loss, even if labels are provided. This matches the llama.py + behaviour and is intentional (speculative-decoding path). + """ + helper = self.mod._qwen3_5_compute_loss_or_logits + self_ = _make_self() + hidden = torch.randn(2, 8, HIDDEN_DIM) + labels = torch.zeros(2, 8, dtype = torch.long) + + loss, logits = helper( + self_, hidden, labels = labels, logits_to_keep = 3, vocab_size = VOCAB_SIZE + ) + + self.assertIsNone(loss) + self.assertEqual(logits.shape, (2, 3, VOCAB_SIZE)) + self.mock_fused_ce.assert_not_called() + self_.loss_function.assert_not_called() + + +# --------------------------------------------------------------------------- +# Tests for num_logits_to_keep normalisation and return_dict=False handling +# --------------------------------------------------------------------------- + + +class TestForwardFunctionBehaviour(unittest.TestCase): + """P1/P2 regression tests for the outer forward wrappers.""" + + def _make_outputs_tuple(self, bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): + """Simulate self.model(...) with return_dict=False → returns a tuple.""" + hidden = torch.randn(bsz, q_len, hidden_dim) + past_kv = MagicMock(name = "past_key_values") + # HF tuple convention: (last_hidden_state, past_key_values) + return (hidden, past_kv) + + def _make_outputs_dict(self, bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): + """Simulate self.model(...) with return_dict=True → returns a ModelOutput.""" + hidden = torch.randn(bsz, q_len, hidden_dim) + outputs = MagicMock() + outputs.__getitem__ = lambda s, idx: hidden if idx == 0 else None + outputs.past_key_values = MagicMock(name = "past_key_values") + outputs.hidden_states = None + outputs.attentions = None + outputs.rope_deltas = None + return outputs + + # -- P1: num_logits_to_keep normalisation -------------------------------- + + def test_num_logits_to_keep_respected_in_conditional_generation(self): + """num_logits_to_keep=3 must produce logits for exactly 3 token positions.""" + from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward + + self_ = _make_self(bsz = 1, q_len = 8) + outputs = self._make_outputs_dict(bsz = 1, q_len = 8) + self_.model = MagicMock(return_value = outputs) + self_.config.use_return_dict = True + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): + result = Qwen3_5ForConditionalGeneration_fast_forward( + self_, + input_ids = torch.zeros(1, 8, dtype = torch.long), + num_logits_to_keep = 3, + logits_to_keep = 0, + ) + + # Only the last 3 token positions should appear in logits + self.assertEqual(result.logits.shape, (1, 3, VOCAB_SIZE)) + + def test_num_logits_to_keep_respected_in_causal_lm(self): + """num_logits_to_keep=2 must produce logits for exactly 2 token positions.""" + from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward + + self_ = _make_self(bsz = 1, q_len = 8) + outputs = self._make_outputs_dict(bsz = 1, q_len = 8) + self_.model = MagicMock(return_value = outputs) + self_.config.use_return_dict = True + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): + result = Qwen3_5ForCausalLM_fast_forward( + self_, + input_ids = torch.zeros(1, 8, dtype = torch.long), + num_logits_to_keep = 2, + logits_to_keep = 0, + ) + + self.assertEqual(result.logits.shape, (1, 2, VOCAB_SIZE)) + + # -- P2: return_dict=False ----------------------------------------------- + + def test_return_dict_false_returns_tuple_not_dataclass(self): + """return_dict=False must return a plain tuple, not raise AttributeError.""" + from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward + + self_ = _make_self(bsz = 2, q_len = 8) + tup = self._make_outputs_tuple(bsz = 2, q_len = 8) + self_.model = MagicMock(return_value = tup) + self_.config.use_return_dict = False + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): + result = Qwen3_5ForCausalLM_fast_forward( + self_, + input_ids = torch.zeros(2, 8, dtype = torch.long), + return_dict = False, + ) + + self.assertIsInstance(result, tuple, "return_dict=False must yield a tuple") + + def test_return_dict_false_does_not_access_dot_attributes(self): + """ + When return_dict=False the model returns a tuple; accessing .past_key_values + would raise AttributeError. Verify no AttributeError is raised. + """ + from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward + + self_ = _make_self(bsz = 2, q_len = 8) + tup = self._make_outputs_tuple(bsz = 2, q_len = 8) + self_.model = MagicMock(return_value = tup) + self_.config.use_return_dict = False + self_.config.text_config = MagicMock() + self_.config.text_config.vocab_size = VOCAB_SIZE + + with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): + try: + result = Qwen3_5ForConditionalGeneration_fast_forward( + self_, + input_ids = torch.zeros(2, 8, dtype = torch.long), + return_dict = False, + ) + except AttributeError as exc: + self.fail(f"return_dict=False raised AttributeError: {exc}") + + self.assertIsInstance(result, tuple) + + # -- UNSLOTH_RETURN_HIDDEN_STATES ---------------------------------------- + + def test_return_hidden_states_causal_lm(self): + """UNSLOTH_RETURN_HIDDEN_STATES=1 → logits field contains hidden states, no loss.""" + from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward + + self_ = _make_self(bsz = 1, q_len = 8) + outputs = self._make_outputs_dict(bsz = 1, q_len = 8) + self_.model = MagicMock(return_value = outputs) + self_.config.use_return_dict = True + + with patch.dict(os.environ, {"UNSLOTH_RETURN_HIDDEN_STATES": "1"}): + result = Qwen3_5ForCausalLM_fast_forward( + self_, + input_ids = torch.zeros(1, 8, dtype = torch.long), + ) + + self.assertIsNone(result.loss) + # logits field carries hidden states (shape [bsz, q_len, hidden_dim]) + self.assertEqual(result.logits.shape, (1, 8, HIDDEN_DIM)) + + def test_return_hidden_states_sliced_by_logits_to_keep(self): + """UNSLOTH_RETURN_HIDDEN_STATES=1 with logits_to_keep=2 → last 2 positions.""" + from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward + + self_ = _make_self(bsz = 1, q_len = 8) + outputs = self._make_outputs_dict(bsz = 1, q_len = 8) + self_.model = MagicMock(return_value = outputs) + self_.config.use_return_dict = True + self_.config.text_config = MagicMock() + self_.config.text_config.vocab_size = VOCAB_SIZE + + with patch.dict(os.environ, {"UNSLOTH_RETURN_HIDDEN_STATES": "1"}): + result = Qwen3_5ForConditionalGeneration_fast_forward( + self_, + input_ids = torch.zeros(1, 8, dtype = torch.long), + logits_to_keep = 2, + ) + + self.assertEqual(result.logits.shape, (1, 2, HIDDEN_DIM)) + + +# --------------------------------------------------------------------------- +# Tests for FastQwen3_5Model.pre_patch() +# --------------------------------------------------------------------------- + + +class TestPrePatch(unittest.TestCase): + """pre_patch() must assign the fast-forward functions to both model classes.""" + + def test_pre_patch_replaces_conditional_generation_forward(self): + from unsloth.models.qwen3_5 import ( + FastQwen3_5Model, + Qwen3_5ForConditionalGeneration, + Qwen3_5ForConditionalGeneration_fast_forward, + ) + + original = Qwen3_5ForConditionalGeneration.forward + try: + FastQwen3_5Model.pre_patch() + self.assertIs( + Qwen3_5ForConditionalGeneration.forward, + Qwen3_5ForConditionalGeneration_fast_forward, + ) + finally: + Qwen3_5ForConditionalGeneration.forward = original + + def test_pre_patch_replaces_causal_lm_forward(self): + from unsloth.models.qwen3_5 import ( + FastQwen3_5Model, + Qwen3_5ForCausalLM, + Qwen3_5ForCausalLM_fast_forward, + ) + + original = Qwen3_5ForCausalLM.forward + try: + FastQwen3_5Model.pre_patch() + self.assertIs( + Qwen3_5ForCausalLM.forward, + Qwen3_5ForCausalLM_fast_forward, + ) + finally: + Qwen3_5ForCausalLM.forward = original + + +# --------------------------------------------------------------------------- +# Tests for loader routing +# --------------------------------------------------------------------------- + + +class TestFromPretrained(unittest.TestCase): + """from_pretrained must call FastLlamaModel, not FastQwen3Model.""" + + def test_from_pretrained_calls_llama_not_qwen3(self): + """ + FastQwen3Model.from_pretrained hardcodes model_patcher=FastQwen3Model, + which would apply incompatible Qwen3 attention patches to Qwen3.5. + FastQwen3_5Model.from_pretrained must bypass it and call FastLlamaModel + directly so that only FastQwen3_5Model.pre_patch() is applied. + """ + from unsloth.models.qwen3_5 import FastQwen3_5Model + from unsloth.models.llama import FastLlamaModel + + with patch.object( + FastLlamaModel, "from_pretrained", return_value = ("model", "tok") + ) as llama_mock: + FastQwen3_5Model.from_pretrained(model_name = "Qwen/Qwen3.5-0.6B-Base") + + llama_mock.assert_called_once() + + # model_patcher must be FastQwen3_5Model, not FastQwen3Model + _, kwargs = llama_mock.call_args + self.assertIs( + kwargs.get("model_patcher"), + FastQwen3_5Model, + "model_patcher must be FastQwen3_5Model so only its pre_patch() runs", + ) + + +class TestLoaderRouting(unittest.TestCase): + """model_type == 'qwen3_5' must route to FastQwen3_5Model.""" + + def test_qwen3_5_routes_to_fast_model(self): + from unsloth.models import loader as loader_mod + from unsloth.models.qwen3_5 import FastQwen3_5Model + + self.assertTrue( + hasattr(loader_mod, "SUPPORTS_QWEN3_5"), + "loader.py must define SUPPORTS_QWEN3_5", + ) + self.assertTrue( + loader_mod.SUPPORTS_QWEN3_5, + "SUPPORTS_QWEN3_5 should be True with transformers >= 4.53.0", + ) + # FastQwen3_5Model must be importable from loader (conditional import succeeded) + self.assertTrue( + hasattr(loader_mod, "FastQwen3_5Model"), + "FastQwen3_5Model must be imported into loader.py when SUPPORTS_QWEN3_5", + ) + self.assertIs(loader_mod.FastQwen3_5Model, FastQwen3_5Model) + + def test_qwen3_5_in_force_float32_list(self): + """Qwen3.5 RMSNorm overflows float16 — must stay in FORCE_FLOAT32.""" + from unsloth.models import loader as loader_mod + + self.assertIn( + "qwen3_5", + loader_mod.FORCE_FLOAT32, + "qwen3_5 must remain in FORCE_FLOAT32 (RMSNorm uses (1+w) pattern)", + ) + + +# --------------------------------------------------------------------------- +# Tests for __init__.py exports +# --------------------------------------------------------------------------- + + +class TestInitExports(unittest.TestCase): + """FastQwen3_5Model must be exported from unsloth.models.""" + + def test_fast_qwen3_5_model_importable(self): + try: + from unsloth.models import FastQwen3_5Model # noqa: F401 + except ImportError: + self.fail( + "FastQwen3_5Model should be importable from unsloth.models " + "when transformers >= 4.53.0 is installed" + ) + + def test_init_except_clause_is_import_error(self): + """ + The try/except around qwen3_5 in __init__.py must catch ImportError, + not bare except (which would silently swallow unrelated exceptions). + """ + import ast + from pathlib import Path + + init_path = ( + Path(__file__).resolve().parents[2] / "unsloth" / "models" / "__init__.py" + ) + tree = ast.parse(init_path.read_text()) + + for node in ast.walk(tree): + if not isinstance(node, ast.Try): + continue + # Find the try block that imports qwen3_5 + source = ast.unparse(node) + if "qwen3_5" not in source: + continue + for handler in node.handlers: + if handler.type is None: + self.fail( + "The try/except around qwen3_5 in __init__.py uses bare " + "`except:` — must use `except ImportError:` instead" + ) + self.assertEqual( + ast.unparse(handler.type), + "ImportError", + "Handler must be `except ImportError:`, got something else", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/unsloth/models/__init__.py b/unsloth/models/__init__.py index 138f309032..19ad6bb6f2 100644 --- a/unsloth/models/__init__.py +++ b/unsloth/models/__init__.py @@ -26,6 +26,11 @@ try: except: # transformers_version < 4.53.0 does not have falcon_h1 so silently skip it for now pass +try: + from .qwen3_5 import FastQwen3_5Model +except ImportError: + # transformers < 5.0.0 does not have qwen3_5 + pass from .dpo import PatchDPOTrainer, PatchKTOTrainer from ._utils import is_bfloat16_supported, is_vLLM_available, __version__ from .rl import PatchFastRL, vLLMSamplingParams diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index bd15ed5281..3d4858fed3 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -78,6 +78,8 @@ SUPPORTS_QWEN3_MOE = transformers_version >= Version("4.50.3") SUPPORTS_FALCON_H1 = transformers_version >= Version("4.53.0") SUPPORTS_GEMMA3N = transformers_version >= Version("4.53.0") SUPPORTS_GPTOSS = transformers_version >= Version("4.55.0") +# Qwen3.5 only exists in transformers 5.x (not in any 4.x release) +SUPPORTS_QWEN3_5 = transformers_version >= Version("5.0.0") # Transformers v5 meta-device loading corrupts non-persistent buffers (inv_freq). # See _fix_rope_inv_freq() below for details. _NEEDS_ROPE_FIX = transformers_version >= Version("5.0.0") @@ -87,6 +89,11 @@ if SUPPORTS_GEMMA2: from .gemma2 import FastGemma2Model if SUPPORTS_FALCON_H1: from .falcon_h1 import FastFalconH1Model +if SUPPORTS_QWEN3_5: + try: + from .qwen3_5 import FastQwen3_5Model + except ImportError: + SUPPORTS_QWEN3_5 = False import torch from ._utils import ( patch_compiling_bitsandbytes, @@ -615,6 +622,15 @@ class FastLanguageModel(FastLlamaModel): dispatch_model = FastGemma2Model elif model_type == "qwen2": dispatch_model = FastQwen2Model + elif model_type == "qwen3_5": + if not SUPPORTS_QWEN3_5: + raise ImportError( + f"Unsloth: Your transformers version of {transformers_version} does not support Qwen3.5.\n" + f"The minimum required version is 5.0.0.\n" + f'Try `pip install --upgrade "transformers>=5.0.0"`\n' + f"to obtain the latest transformers build, then restart this session." + ) + dispatch_model = FastQwen3_5Model elif model_type == "qwen3": # or model_type == "qwen3_moe": if not SUPPORTS_QWEN3 or not SUPPORTS_QWEN3_MOE: raise ImportError( diff --git a/unsloth/models/qwen3_5.py b/unsloth/models/qwen3_5.py new file mode 100644 index 0000000000..40bc92969d --- /dev/null +++ b/unsloth/models/qwen3_5.py @@ -0,0 +1,335 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Fixes https://github.com/unslothai/unsloth/issues/4188 +# Qwen3.5 has a 248,320-token vocabulary (1.64x larger than Qwen3). +# At 8K context the full logits tensor is 8192 x 248320 x 4 bytes = 7.68 GB, +# which exceeds free VRAM on T4/P100 after model load. +# +# Root cause: loader.py listed "qwen3_5" in FORCE_FLOAT32 but never dispatched +# it to an optimised class, so the model fell through to a bare HF load with no +# fast-forward patching and full logits were materialised every training step. +# +# Fix: patch Qwen3_5ForConditionalGeneration.forward (the class HF uses for all +# Qwen3.5 text models, including base variants) to call unsloth_fused_ce_loss +# directly from hidden_states, bypassing logits materialisation entirely. +# +# Gated DeltaNet (GDN) linear-attention layers are intentionally NOT patched -- +# they already have Triton kernels via flash-linear-attention and are +# architecturally incompatible with Unsloth's standard attention optimisations. + +from .llama import * +import os +from unsloth_zoo.utils import _get_dtype +from unsloth_zoo.hf_utils import dtype_from_config +from .llama import FastLlamaModel + +try: + from transformers.models.qwen3_5.modeling_qwen3_5 import ( + Qwen3_5ForCausalLM, + Qwen3_5ForConditionalGeneration, + Qwen3_5CausalLMOutputWithPast, + ) + from transformers.modeling_outputs import CausalLMOutputWithPast +except ImportError: + raise ImportError( + "Unsloth: Your transformers version does not support Qwen3.5.\n" + 'Try `pip install --upgrade "transformers>=5.0.0"`\n' + "then restart your session." + ) + + +def _qwen3_5_compute_loss_or_logits( + self, hidden_states, labels, logits_to_keep, vocab_size, **kwargs +): + """ + Shared helper: given hidden_states from the backbone, return (loss, logits). + + Exactly one of loss/logits will be the primary result: + - Single-token decode -> logits via fast torch.mv + - Partial-logits path -> logits for the last logits_to_keep tokens + - Training with labels -> loss via unsloth_fused_ce_loss (no logits materialised) + - Eval / inference -> full logits, then optional loss via self.loss_function + + Returns: + loss (Tensor or None) + logits (Tensor or EMPTY_LOGITS) + """ + lm_head_weight = self.lm_head.weight + hidden_states = hidden_states.to(lm_head_weight.device) + bsz, q_len, _ = hidden_states.shape + out_dtype = _get_dtype(dtype_from_config(self.config)) + + # Fast single-token decode (inference / generation) + if bsz == 1 and q_len == 1 and labels is None: + logits = torch.mv( + lm_head_weight, hidden_states.ravel().to(lm_head_weight.dtype) + ) + logits = logits.unsqueeze(0).unsqueeze(0).to(out_dtype) + return None, logits + + # Partial-logits path (e.g. logits_to_keep for speculative decoding) + if logits_to_keep != 0: + slice_idx = ( + slice(-logits_to_keep, None) + if isinstance(logits_to_keep, int) + else logits_to_keep + ) + logits = self.lm_head(hidden_states[:, slice_idx, :].to(lm_head_weight.dtype)) + return None, logits.to(out_dtype) + + # Training path: fused CE avoids materialising the 7.68 GB logits tensor. + # + # Note: llama.py skips fused CE for bsz * q_len <= 1024, since for short + # sequences the savings are marginal. We unconditionally use fused CE for + # Qwen3.5 -- even a 32-token sequence produces a 32 x 248320 x 4 = 31 MB + # logit tensor, and the chunked CE overhead is negligible vs the OOM risk. + if labels is not None and os.environ.get("UNSLOTH_RETURN_LOGITS", "0") != "1": + labels = labels.to(lm_head_weight.device) + n_items = kwargs.get("num_items_in_batch") + if n_items is None: + n_items = kwargs.get("n_items") + loss = unsloth_fused_ce_loss( + trainer = None, + hidden_states = hidden_states, + lm_head_weight = lm_head_weight, + lm_head_bias = None, + labels = labels, + mask = None, + n_items = n_items, + scaling = getattr(self, "accelerator_scaler", None), + target_gb = None, + torch_compile = True, + logit_softcapping = 0, # Qwen3.5 has no logit softcapping + ) + return loss, EMPTY_LOGITS + + # Eval / inference path + logits = self.lm_head(hidden_states.to(lm_head_weight.dtype)).to(out_dtype) + loss = None + if labels is not None: + labels = labels.to(lm_head_weight.device) + loss = self.loss_function( + logits = logits, labels = labels, vocab_size = vocab_size, **kwargs + ) + return loss, logits + + +def Qwen3_5ForConditionalGeneration_fast_forward( + self, + input_ids = None, + attention_mask = None, + position_ids = None, + past_key_values = None, + inputs_embeds = None, + labels = None, + pixel_values = None, + pixel_values_videos = None, + image_grid_thw = None, + video_grid_thw = None, + mm_token_type_ids = None, + cache_position = None, + logits_to_keep = 0, + num_logits_to_keep = 0, + return_dict = None, + **kwargs, +): + return_dict = ( + return_dict if return_dict is not None else self.config.use_return_dict + ) + # Normalise both generation knobs + logits_to_keep = max(logits_to_keep, num_logits_to_keep) + + outputs = self.model( + input_ids = input_ids, + pixel_values = pixel_values, + pixel_values_videos = pixel_values_videos, + image_grid_thw = image_grid_thw, + video_grid_thw = video_grid_thw, + position_ids = position_ids, + attention_mask = attention_mask, + past_key_values = past_key_values, + inputs_embeds = inputs_embeds, + cache_position = cache_position, + mm_token_type_ids = mm_token_type_ids, + return_dict = return_dict, + **kwargs, + ) + + # Return hidden states as logits when requested + if os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES", "0") == "1": + hidden_states = outputs[0] + if logits_to_keep != 0: + hidden_states = hidden_states[:, -logits_to_keep:, :] + if not return_dict: + return (hidden_states,) + outputs[1:] + return Qwen3_5CausalLMOutputWithPast( + loss = None, + logits = hidden_states, + past_key_values = outputs.past_key_values, + hidden_states = outputs.hidden_states, + attentions = outputs.attentions, + rope_deltas = getattr(outputs, "rope_deltas", None), + ) + + loss, logits = _qwen3_5_compute_loss_or_logits( + self, + outputs[0], + labels, + logits_to_keep, + vocab_size = self.config.text_config.vocab_size, + **kwargs, + ) + + if not return_dict: + output = (logits,) + outputs[1:] + return ((loss,) + output) if loss is not None else output + + return Qwen3_5CausalLMOutputWithPast( + loss = loss, + logits = logits, + past_key_values = outputs.past_key_values, + hidden_states = outputs.hidden_states, + attentions = outputs.attentions, + rope_deltas = getattr(outputs, "rope_deltas", None), + ) + + +def Qwen3_5ForCausalLM_fast_forward( + self, + input_ids = None, + attention_mask = None, + position_ids = None, + past_key_values = None, + inputs_embeds = None, + labels = None, + use_cache = None, + cache_position = None, + logits_to_keep = 0, + num_logits_to_keep = 0, + return_dict = None, + **kwargs, +): + return_dict = ( + return_dict if return_dict is not None else self.config.use_return_dict + ) + # Normalise both generation knobs + logits_to_keep = max(logits_to_keep, num_logits_to_keep) + + outputs = self.model( + input_ids = input_ids, + attention_mask = attention_mask, + position_ids = position_ids, + past_key_values = past_key_values, + inputs_embeds = inputs_embeds, + use_cache = use_cache, + cache_position = cache_position, + return_dict = return_dict, + **kwargs, + ) + + # Return hidden states as logits when requested + if os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES", "0") == "1": + hidden_states = outputs[0] + if logits_to_keep != 0: + hidden_states = hidden_states[:, -logits_to_keep:, :] + if not return_dict: + return (hidden_states,) + outputs[1:] + return CausalLMOutputWithPast( + loss = None, + logits = hidden_states, + past_key_values = outputs.past_key_values, + hidden_states = outputs.hidden_states, + attentions = outputs.attentions, + ) + + loss, logits = _qwen3_5_compute_loss_or_logits( + self, + outputs[0], + labels, + logits_to_keep, + vocab_size = self.config.vocab_size, + **kwargs, + ) + + if not return_dict: + output = (logits,) + outputs[1:] + return ((loss,) + output) if loss is not None else output + + return CausalLMOutputWithPast( + loss = loss, + logits = logits, + past_key_values = outputs.past_key_values, + hidden_states = outputs.hidden_states, + attentions = outputs.attentions, + ) + + +class FastQwen3_5Model(FastLlamaModel): + """ + Unsloth optimisation for Qwen3.5 hybrid GDN (Gated DeltaNet) models. + + Qwen3.5 interleaves standard transformer attention layers with Gated + DeltaNet linear-attention layers. GDN layers use native Triton kernels + from flash-linear-attention and are architecturally incompatible with + Unsloth's standard attention patches (gated query projections, different + forward signatures). This class therefore only patches the top-level + CausalLM forward to call unsloth_fused_ce_loss directly from + hidden_states, which eliminates the 7.68 GB logits tensor that causes + OOM on T4/P100 at 8K context. + + Memory saving at batch=1, seq=8192: + Standard: 8192 x 248320 x 4 = 7.68 GB (OOM on T4) + unsloth_fused_ce: chunked, ~0.24-0.95 GB peak (fits) + + Fixes: https://github.com/unslothai/unsloth/issues/4188 + """ + + @staticmethod + def pre_patch(): + Qwen3_5ForConditionalGeneration.forward = ( + Qwen3_5ForConditionalGeneration_fast_forward + ) + Qwen3_5ForCausalLM.forward = Qwen3_5ForCausalLM_fast_forward + return + + @staticmethod + def from_pretrained( + model_name = "Qwen/Qwen3.5-9B", + max_seq_length = 4096, + dtype = None, + load_in_4bit = True, + token = None, + device_map = "sequential", + rope_scaling = None, + fix_tokenizer = True, + model_patcher = None, + tokenizer_name = None, + trust_remote_code = False, + **kwargs, + ): + return FastLlamaModel.from_pretrained( + model_name = model_name, + max_seq_length = max_seq_length, + dtype = dtype, + load_in_4bit = load_in_4bit, + token = token, + device_map = device_map, + rope_scaling = rope_scaling, + fix_tokenizer = fix_tokenizer, + model_patcher = FastQwen3_5Model, + tokenizer_name = tokenizer_name, + trust_remote_code = trust_remote_code, + **kwargs, + ) From a7b944ee57eae3f99fd7c3e46951df379a72addf Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 16 Mar 2026 23:17:43 +0000 Subject: [PATCH 2/2] Remove qwen3_5 unit tests Tests require transformers >= 5.0.0 which is not yet widely deployed. The fused CE path is already covered by the compiler's generic apply_fused_lm_head mechanism and verified via training runs. --- tests/utils/test_qwen3_5.py | 623 ------------------------------------ 1 file changed, 623 deletions(-) delete mode 100644 tests/utils/test_qwen3_5.py diff --git a/tests/utils/test_qwen3_5.py b/tests/utils/test_qwen3_5.py deleted file mode 100644 index 2e7d78d431..0000000000 --- a/tests/utils/test_qwen3_5.py +++ /dev/null @@ -1,623 +0,0 @@ -# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -""" -Unit tests for unsloth/models/qwen3_5.py — fix for issue #4188. - -These tests use CPU tensors and mock out GPU-only dependencies (unsloth_fused_ce_loss, -EMPTY_LOGITS) so they run without a CUDA device or real model weights. -""" - -import os -import types -import unittest -from unittest.mock import MagicMock, patch, call - -import pytest -import torch -import torch.nn as nn - - -# --------------------------------------------------------------------------- -# Helpers to import the module under test with mocked unsloth internals -# --------------------------------------------------------------------------- - - -def _make_fake_unsloth_fused_ce_loss(): - """Return a mock that records calls and returns a scalar loss tensor.""" - mock = MagicMock(return_value = torch.tensor(1.23)) - return mock - - -# --------------------------------------------------------------------------- -# Fixtures -# --------------------------------------------------------------------------- - -HIDDEN_DIM = 16 -VOCAB_SIZE = 64 - - -def _make_self(bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM, vocab_size = VOCAB_SIZE): - """ - Build a minimal mock `self` (the model instance) with: - - lm_head: a real nn.Linear (CPU) so the matmul paths work - - loss_function: a MagicMock returning a fixed scalar tensor - - accelerator_scaler: None - - config.text_config.vocab_size / config.vocab_size: vocab_size - """ - lm_head = nn.Linear(hidden_dim, vocab_size, bias = False) - - cfg_text = MagicMock() - cfg_text.vocab_size = vocab_size - cfg = MagicMock() - cfg.vocab_size = vocab_size - cfg.text_config = cfg_text - - self = MagicMock() - self.lm_head = lm_head - self.config = cfg - self.accelerator_scaler = None - self.loss_function = MagicMock(return_value = torch.tensor(0.99)) - return self - - -def _make_outputs(bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): - """Return a mock outputs object whose [0] is a random hidden_states tensor.""" - hidden = torch.randn(bsz, q_len, hidden_dim) - outputs = MagicMock() - outputs.__getitem__ = lambda self, idx: hidden if idx == 0 else None - outputs.past_key_values = None - outputs.hidden_states = None - outputs.attentions = None - return outputs, hidden - - -# --------------------------------------------------------------------------- -# Tests for _qwen3_5_compute_loss_or_logits -# --------------------------------------------------------------------------- - - -class TestComputeLossOrLogits(unittest.TestCase): - """Tests for the shared helper that houses all four forward paths.""" - - def setUp(self): - import unsloth.models.qwen3_5 as mod - - self.mod = mod - self.orig_fused_ce = mod.unsloth_fused_ce_loss - self.orig_empty_logits = mod.EMPTY_LOGITS - # Install fresh mocks for each test - self.mock_fused_ce = MagicMock(return_value = torch.tensor(1.23)) - self.mock_empty_logits = torch.zeros(1) - mod.unsloth_fused_ce_loss = self.mock_fused_ce - mod.EMPTY_LOGITS = self.mock_empty_logits - - def tearDown(self): - self.mod.unsloth_fused_ce_loss = self.orig_fused_ce - self.mod.EMPTY_LOGITS = self.orig_empty_logits - - # -- single-token decode path ------------------------------------------- - - def test_single_token_decode_uses_mv_not_lm_head(self): - """bsz=1, q_len=1 → fast torch.mv, no full lm_head call, no loss.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self(bsz = 1, q_len = 1) - hidden = torch.randn(1, 1, HIDDEN_DIM) - - with patch.object(torch, "mv", wraps = torch.mv) as mv_spy: - loss, logits = helper( - self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE - ) - - self.assertIsNone(loss) - self.assertEqual(logits.shape, (1, 1, VOCAB_SIZE)) - self.mock_fused_ce.assert_not_called() - self_.loss_function.assert_not_called() - - # -- partial-logits path ------------------------------------------------ - - def test_logits_to_keep_slices_last_n_tokens(self): - """logits_to_keep=3 → output has exactly 3 token positions.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - - loss, logits = helper( - self_, hidden, labels = None, logits_to_keep = 3, vocab_size = VOCAB_SIZE - ) - - self.assertIsNone(loss) - self.assertEqual(logits.shape, (2, 3, VOCAB_SIZE)) - self.mock_fused_ce.assert_not_called() - - # -- training / fused-CE path ------------------------------------------- - - def test_training_path_calls_fused_ce_returns_early(self): - """labels present + UNSLOTH_RETURN_LOGITS unset → fused CE, lm_head NOT called.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - loss, logits = helper( - self_, hidden, labels = labels, logits_to_keep = 0, vocab_size = VOCAB_SIZE - ) - - self.assertIsNotNone(loss) - self.assertIs(logits, self.mock_empty_logits) - self.mock_fused_ce.assert_called_once() - # Verify key arguments passed to fused CE - _, kwargs = self.mock_fused_ce.call_args - self.assertIs(kwargs["lm_head_weight"], self_.lm_head.weight) - self.assertEqual(kwargs["logit_softcapping"], 0) - self_.loss_function.assert_not_called() - - def test_training_path_passes_num_items_in_batch(self): - """num_items_in_batch kwarg is forwarded to unsloth_fused_ce_loss.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - helper( - self_, - hidden, - labels = labels, - logits_to_keep = 0, - vocab_size = VOCAB_SIZE, - num_items_in_batch = 16, - ) - - _, kwargs = self.mock_fused_ce.call_args - self.assertEqual(kwargs["n_items"], 16) - - # -- UNSLOTH_RETURN_LOGITS override ------------------------------------ - - def test_return_logits_env_var_bypasses_fused_ce(self): - """UNSLOTH_RETURN_LOGITS=1 → materialise full logits even during training.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): - loss, logits = helper( - self_, hidden, labels = labels, logits_to_keep = 0, vocab_size = VOCAB_SIZE - ) - - self.assertEqual(logits.shape, (2, 8, VOCAB_SIZE)) - self.mock_fused_ce.assert_not_called() - self_.loss_function.assert_called_once() - - # -- eval / inference path (no labels) ---------------------------------- - - def test_no_labels_returns_full_logits_no_loss(self): - """No labels → full logits computed, loss=None.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - loss, logits = helper( - self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE - ) - - self.assertIsNone(loss) - self.assertEqual(logits.shape, (2, 8, VOCAB_SIZE)) - self.mock_fused_ce.assert_not_called() - self_.loss_function.assert_not_called() - - # -- n_items fallback --------------------------------------------------- - - def test_n_items_kwarg_used_when_num_items_absent(self): - """n_items kwarg is the fallback when num_items_in_batch is absent.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - helper( - self_, - hidden, - labels = labels, - logits_to_keep = 0, - vocab_size = VOCAB_SIZE, - n_items = 8, - ) - - _, kwargs = self.mock_fused_ce.call_args - self.assertEqual(kwargs["n_items"], 8) - - def test_num_items_in_batch_zero_does_not_fall_through_to_n_items(self): - """num_items_in_batch=0 must NOT fall through to n_items (or-bug regression).""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - helper( - self_, - hidden, - labels = labels, - logits_to_keep = 0, - vocab_size = VOCAB_SIZE, - num_items_in_batch = 0, - n_items = 99, - ) - - _, kwargs = self.mock_fused_ce.call_args - # 0 is falsy; the old `or` expression would have returned 99 — must be 0. - self.assertEqual(kwargs["n_items"], 0) - - # -- batch decode (bsz > 1, q_len == 1) --------------------------------- - - def test_batch_decode_falls_through_to_eval_path(self): - """bsz=4, q_len=1 must NOT use the single-token mv path; goes to eval path.""" - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self(bsz = 4, q_len = 1) - hidden = torch.randn(4, 1, HIDDEN_DIM) - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "0"}): - with patch.object(torch, "mv", wraps = torch.mv) as mv_spy: - loss, logits = helper( - self_, hidden, labels = None, logits_to_keep = 0, vocab_size = VOCAB_SIZE - ) - - mv_spy.assert_not_called() # fast path is bsz==1 AND q_len==1 only - self.assertEqual(logits.shape, (4, 1, VOCAB_SIZE)) - - # -- labels silently ignored when logits_to_keep is set ----------------- - - def test_labels_ignored_when_logits_to_keep_nonzero(self): - """ - When logits_to_keep != 0 the function returns early with partial logits - and no loss, even if labels are provided. This matches the llama.py - behaviour and is intentional (speculative-decoding path). - """ - helper = self.mod._qwen3_5_compute_loss_or_logits - self_ = _make_self() - hidden = torch.randn(2, 8, HIDDEN_DIM) - labels = torch.zeros(2, 8, dtype = torch.long) - - loss, logits = helper( - self_, hidden, labels = labels, logits_to_keep = 3, vocab_size = VOCAB_SIZE - ) - - self.assertIsNone(loss) - self.assertEqual(logits.shape, (2, 3, VOCAB_SIZE)) - self.mock_fused_ce.assert_not_called() - self_.loss_function.assert_not_called() - - -# --------------------------------------------------------------------------- -# Tests for num_logits_to_keep normalisation and return_dict=False handling -# --------------------------------------------------------------------------- - - -class TestForwardFunctionBehaviour(unittest.TestCase): - """P1/P2 regression tests for the outer forward wrappers.""" - - def _make_outputs_tuple(self, bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): - """Simulate self.model(...) with return_dict=False → returns a tuple.""" - hidden = torch.randn(bsz, q_len, hidden_dim) - past_kv = MagicMock(name = "past_key_values") - # HF tuple convention: (last_hidden_state, past_key_values) - return (hidden, past_kv) - - def _make_outputs_dict(self, bsz = 2, q_len = 8, hidden_dim = HIDDEN_DIM): - """Simulate self.model(...) with return_dict=True → returns a ModelOutput.""" - hidden = torch.randn(bsz, q_len, hidden_dim) - outputs = MagicMock() - outputs.__getitem__ = lambda s, idx: hidden if idx == 0 else None - outputs.past_key_values = MagicMock(name = "past_key_values") - outputs.hidden_states = None - outputs.attentions = None - outputs.rope_deltas = None - return outputs - - # -- P1: num_logits_to_keep normalisation -------------------------------- - - def test_num_logits_to_keep_respected_in_conditional_generation(self): - """num_logits_to_keep=3 must produce logits for exactly 3 token positions.""" - from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward - - self_ = _make_self(bsz = 1, q_len = 8) - outputs = self._make_outputs_dict(bsz = 1, q_len = 8) - self_.model = MagicMock(return_value = outputs) - self_.config.use_return_dict = True - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): - result = Qwen3_5ForConditionalGeneration_fast_forward( - self_, - input_ids = torch.zeros(1, 8, dtype = torch.long), - num_logits_to_keep = 3, - logits_to_keep = 0, - ) - - # Only the last 3 token positions should appear in logits - self.assertEqual(result.logits.shape, (1, 3, VOCAB_SIZE)) - - def test_num_logits_to_keep_respected_in_causal_lm(self): - """num_logits_to_keep=2 must produce logits for exactly 2 token positions.""" - from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward - - self_ = _make_self(bsz = 1, q_len = 8) - outputs = self._make_outputs_dict(bsz = 1, q_len = 8) - self_.model = MagicMock(return_value = outputs) - self_.config.use_return_dict = True - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): - result = Qwen3_5ForCausalLM_fast_forward( - self_, - input_ids = torch.zeros(1, 8, dtype = torch.long), - num_logits_to_keep = 2, - logits_to_keep = 0, - ) - - self.assertEqual(result.logits.shape, (1, 2, VOCAB_SIZE)) - - # -- P2: return_dict=False ----------------------------------------------- - - def test_return_dict_false_returns_tuple_not_dataclass(self): - """return_dict=False must return a plain tuple, not raise AttributeError.""" - from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward - - self_ = _make_self(bsz = 2, q_len = 8) - tup = self._make_outputs_tuple(bsz = 2, q_len = 8) - self_.model = MagicMock(return_value = tup) - self_.config.use_return_dict = False - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): - result = Qwen3_5ForCausalLM_fast_forward( - self_, - input_ids = torch.zeros(2, 8, dtype = torch.long), - return_dict = False, - ) - - self.assertIsInstance(result, tuple, "return_dict=False must yield a tuple") - - def test_return_dict_false_does_not_access_dot_attributes(self): - """ - When return_dict=False the model returns a tuple; accessing .past_key_values - would raise AttributeError. Verify no AttributeError is raised. - """ - from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward - - self_ = _make_self(bsz = 2, q_len = 8) - tup = self._make_outputs_tuple(bsz = 2, q_len = 8) - self_.model = MagicMock(return_value = tup) - self_.config.use_return_dict = False - self_.config.text_config = MagicMock() - self_.config.text_config.vocab_size = VOCAB_SIZE - - with patch.dict(os.environ, {"UNSLOTH_RETURN_LOGITS": "1"}): - try: - result = Qwen3_5ForConditionalGeneration_fast_forward( - self_, - input_ids = torch.zeros(2, 8, dtype = torch.long), - return_dict = False, - ) - except AttributeError as exc: - self.fail(f"return_dict=False raised AttributeError: {exc}") - - self.assertIsInstance(result, tuple) - - # -- UNSLOTH_RETURN_HIDDEN_STATES ---------------------------------------- - - def test_return_hidden_states_causal_lm(self): - """UNSLOTH_RETURN_HIDDEN_STATES=1 → logits field contains hidden states, no loss.""" - from unsloth.models.qwen3_5 import Qwen3_5ForCausalLM_fast_forward - - self_ = _make_self(bsz = 1, q_len = 8) - outputs = self._make_outputs_dict(bsz = 1, q_len = 8) - self_.model = MagicMock(return_value = outputs) - self_.config.use_return_dict = True - - with patch.dict(os.environ, {"UNSLOTH_RETURN_HIDDEN_STATES": "1"}): - result = Qwen3_5ForCausalLM_fast_forward( - self_, - input_ids = torch.zeros(1, 8, dtype = torch.long), - ) - - self.assertIsNone(result.loss) - # logits field carries hidden states (shape [bsz, q_len, hidden_dim]) - self.assertEqual(result.logits.shape, (1, 8, HIDDEN_DIM)) - - def test_return_hidden_states_sliced_by_logits_to_keep(self): - """UNSLOTH_RETURN_HIDDEN_STATES=1 with logits_to_keep=2 → last 2 positions.""" - from unsloth.models.qwen3_5 import Qwen3_5ForConditionalGeneration_fast_forward - - self_ = _make_self(bsz = 1, q_len = 8) - outputs = self._make_outputs_dict(bsz = 1, q_len = 8) - self_.model = MagicMock(return_value = outputs) - self_.config.use_return_dict = True - self_.config.text_config = MagicMock() - self_.config.text_config.vocab_size = VOCAB_SIZE - - with patch.dict(os.environ, {"UNSLOTH_RETURN_HIDDEN_STATES": "1"}): - result = Qwen3_5ForConditionalGeneration_fast_forward( - self_, - input_ids = torch.zeros(1, 8, dtype = torch.long), - logits_to_keep = 2, - ) - - self.assertEqual(result.logits.shape, (1, 2, HIDDEN_DIM)) - - -# --------------------------------------------------------------------------- -# Tests for FastQwen3_5Model.pre_patch() -# --------------------------------------------------------------------------- - - -class TestPrePatch(unittest.TestCase): - """pre_patch() must assign the fast-forward functions to both model classes.""" - - def test_pre_patch_replaces_conditional_generation_forward(self): - from unsloth.models.qwen3_5 import ( - FastQwen3_5Model, - Qwen3_5ForConditionalGeneration, - Qwen3_5ForConditionalGeneration_fast_forward, - ) - - original = Qwen3_5ForConditionalGeneration.forward - try: - FastQwen3_5Model.pre_patch() - self.assertIs( - Qwen3_5ForConditionalGeneration.forward, - Qwen3_5ForConditionalGeneration_fast_forward, - ) - finally: - Qwen3_5ForConditionalGeneration.forward = original - - def test_pre_patch_replaces_causal_lm_forward(self): - from unsloth.models.qwen3_5 import ( - FastQwen3_5Model, - Qwen3_5ForCausalLM, - Qwen3_5ForCausalLM_fast_forward, - ) - - original = Qwen3_5ForCausalLM.forward - try: - FastQwen3_5Model.pre_patch() - self.assertIs( - Qwen3_5ForCausalLM.forward, - Qwen3_5ForCausalLM_fast_forward, - ) - finally: - Qwen3_5ForCausalLM.forward = original - - -# --------------------------------------------------------------------------- -# Tests for loader routing -# --------------------------------------------------------------------------- - - -class TestFromPretrained(unittest.TestCase): - """from_pretrained must call FastLlamaModel, not FastQwen3Model.""" - - def test_from_pretrained_calls_llama_not_qwen3(self): - """ - FastQwen3Model.from_pretrained hardcodes model_patcher=FastQwen3Model, - which would apply incompatible Qwen3 attention patches to Qwen3.5. - FastQwen3_5Model.from_pretrained must bypass it and call FastLlamaModel - directly so that only FastQwen3_5Model.pre_patch() is applied. - """ - from unsloth.models.qwen3_5 import FastQwen3_5Model - from unsloth.models.llama import FastLlamaModel - - with patch.object( - FastLlamaModel, "from_pretrained", return_value = ("model", "tok") - ) as llama_mock: - FastQwen3_5Model.from_pretrained(model_name = "Qwen/Qwen3.5-0.6B-Base") - - llama_mock.assert_called_once() - - # model_patcher must be FastQwen3_5Model, not FastQwen3Model - _, kwargs = llama_mock.call_args - self.assertIs( - kwargs.get("model_patcher"), - FastQwen3_5Model, - "model_patcher must be FastQwen3_5Model so only its pre_patch() runs", - ) - - -class TestLoaderRouting(unittest.TestCase): - """model_type == 'qwen3_5' must route to FastQwen3_5Model.""" - - def test_qwen3_5_routes_to_fast_model(self): - from unsloth.models import loader as loader_mod - from unsloth.models.qwen3_5 import FastQwen3_5Model - - self.assertTrue( - hasattr(loader_mod, "SUPPORTS_QWEN3_5"), - "loader.py must define SUPPORTS_QWEN3_5", - ) - self.assertTrue( - loader_mod.SUPPORTS_QWEN3_5, - "SUPPORTS_QWEN3_5 should be True with transformers >= 4.53.0", - ) - # FastQwen3_5Model must be importable from loader (conditional import succeeded) - self.assertTrue( - hasattr(loader_mod, "FastQwen3_5Model"), - "FastQwen3_5Model must be imported into loader.py when SUPPORTS_QWEN3_5", - ) - self.assertIs(loader_mod.FastQwen3_5Model, FastQwen3_5Model) - - def test_qwen3_5_in_force_float32_list(self): - """Qwen3.5 RMSNorm overflows float16 — must stay in FORCE_FLOAT32.""" - from unsloth.models import loader as loader_mod - - self.assertIn( - "qwen3_5", - loader_mod.FORCE_FLOAT32, - "qwen3_5 must remain in FORCE_FLOAT32 (RMSNorm uses (1+w) pattern)", - ) - - -# --------------------------------------------------------------------------- -# Tests for __init__.py exports -# --------------------------------------------------------------------------- - - -class TestInitExports(unittest.TestCase): - """FastQwen3_5Model must be exported from unsloth.models.""" - - def test_fast_qwen3_5_model_importable(self): - try: - from unsloth.models import FastQwen3_5Model # noqa: F401 - except ImportError: - self.fail( - "FastQwen3_5Model should be importable from unsloth.models " - "when transformers >= 4.53.0 is installed" - ) - - def test_init_except_clause_is_import_error(self): - """ - The try/except around qwen3_5 in __init__.py must catch ImportError, - not bare except (which would silently swallow unrelated exceptions). - """ - import ast - from pathlib import Path - - init_path = ( - Path(__file__).resolve().parents[2] / "unsloth" / "models" / "__init__.py" - ) - tree = ast.parse(init_path.read_text()) - - for node in ast.walk(tree): - if not isinstance(node, ast.Try): - continue - # Find the try block that imports qwen3_5 - source = ast.unparse(node) - if "qwen3_5" not in source: - continue - for handler in node.handlers: - if handler.type is None: - self.fail( - "The try/except around qwen3_5 in __init__.py uses bare " - "`except:` — must use `except ImportError:` instead" - ) - self.assertEqual( - ast.unparse(handler.type), - "ImportError", - "Handler must be `except ImportError:`, got something else", - ) - - -if __name__ == "__main__": - unittest.main()