diff --git a/tests/test_finetune_last_n_layers.py b/tests/test_finetune_last_n_layers.py new file mode 100644 index 0000000000..fa0e2af5bc --- /dev/null +++ b/tests/test_finetune_last_n_layers.py @@ -0,0 +1,92 @@ +# Unsloth - 2x faster, 70% less memory LLM finetuning +# Tests for the `finetune_last_n_layers` parity knob (CUDA side). +# +# Mirrors unsloth-zoo's `FastMLXModel.get_peft_model` parameter. +# mlx-lm CLI's CONFIG_DEFAULTS['num_layers']=16 applies LoRA to the +# last 16 transformer blocks only. On the CUDA path, PEFT exposes +# `layers_to_transform` to do the same. This convenience knob fills +# `layers_to_transform` for the user when set, matching mlx-lm CLI +# AND unsloth-zoo's MLX path with a single config value. +# +# The tests intentionally avoid pulling in CUDA / a real model +# checkpoint — they exercise only the helper that translates +# `finetune_last_n_layers` into `layers_to_transform`. + +from __future__ import annotations + +import pytest + + +def test_get_total_transformer_layers_reads_num_hidden_layers(): + from unsloth.models.vision import _get_total_transformer_layers + + class FakeConfig: + num_hidden_layers = 18 + + class FakeModel: + config = FakeConfig() + + assert _get_total_transformer_layers(FakeModel()) == 18 + + +def test_get_total_transformer_layers_reads_text_config(): + from unsloth.models.vision import _get_total_transformer_layers + + class TextConfig: + num_hidden_layers = 24 + + class FakeConfig: + text_config = TextConfig() + + class FakeModel: + config = FakeConfig() + + # No num_hidden_layers at top level — should fall through to text_config. + assert _get_total_transformer_layers(FakeModel()) == 24 + + +def test_get_total_transformer_layers_handles_alternative_attr_names(): + from unsloth.models.vision import _get_total_transformer_layers + + for attr in ("n_layer", "n_layers", "num_layers"): + cfg = type("Cfg", (), {attr: 12})() + model = type("M", (), {"config": cfg})() + assert _get_total_transformer_layers(model) == 12 + + +def test_get_total_transformer_layers_returns_none_when_unknown(): + from unsloth.models.vision import _get_total_transformer_layers + + class FakeConfig: + pass + + class FakeModel: + config = FakeConfig() + + assert _get_total_transformer_layers(FakeModel()) is None + + +def test_get_total_transformer_layers_returns_none_for_missing_config(): + from unsloth.models.vision import _get_total_transformer_layers + + class FakeModel: + pass + + assert _get_total_transformer_layers(FakeModel()) is None + + +def test_finetune_last_n_layers_signature_present_on_llama_and_vision(): + """Both entry points must expose the new parameter with default None.""" + import inspect + from unsloth.models.llama import FastLlamaModel + from unsloth.models.vision import FastBaseModel + + for cls in (FastLlamaModel, FastBaseModel): + sig = inspect.signature(cls.get_peft_model) + assert ( + "finetune_last_n_layers" in sig.parameters + ), f"{cls.__name__}.get_peft_model missing finetune_last_n_layers" + assert sig.parameters["finetune_last_n_layers"].default is None, ( + f"{cls.__name__}.get_peft_model: finetune_last_n_layers default " + f"must be None to preserve historical behavior" + ) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6ddfe04d21..20b515c711 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -2831,6 +2831,7 @@ class FastLlamaModel: bias = "none", layers_to_transform = None, layers_pattern = None, + finetune_last_n_layers = None, use_gradient_checkpointing = "unsloth", random_state = 3407, max_seq_length = 2048, # not used anymore @@ -2863,6 +2864,7 @@ class FastLlamaModel: bias = bias, layers_to_transform = layers_to_transform, layers_pattern = layers_pattern, + finetune_last_n_layers = finetune_last_n_layers, use_gradient_checkpointing = use_gradient_checkpointing, random_state = random_state, max_seq_length = max_seq_length, @@ -3160,6 +3162,14 @@ class FastLlamaModel: if target_parameters is None: target_parameters = get_moe_target_parameters(model, target_modules) + if finetune_last_n_layers is not None and layers_to_transform is None: + from .vision import _get_total_transformer_layers + + _total_layers = _get_total_transformer_layers(model) + if _total_layers is not None and _total_layers > 0: + _n = max(1, min(int(finetune_last_n_layers), _total_layers)) + layers_to_transform = list(range(_total_layers - _n, _total_layers)) + arguments = dict( r = r, lora_alpha = lora_alpha, diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index df371e00c8..73ef1db7e3 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -547,6 +547,35 @@ def _construct_vlm_processor_fallback( return None +def _get_total_transformer_layers(model): + """Best-effort total transformer block count across HF model shapes. + Returns None if not determinable; caller should skip the conversion.""" + cfg = getattr(model, "config", None) + if cfg is None: + return None + for name in ( + "num_hidden_layers", + "n_layer", + "n_layers", + "num_layers", + ): + v = getattr(cfg, name, None) + if isinstance(v, int) and v > 0: + return v + text_cfg = getattr(cfg, "text_config", None) + if text_cfg is not None: + for name in ( + "num_hidden_layers", + "n_layer", + "n_layers", + "num_layers", + ): + v = getattr(text_cfg, name, None) + if isinstance(v, int) and v > 0: + return v + return None + + class FastBaseModel: @staticmethod def from_pretrained( @@ -1319,6 +1348,7 @@ class FastBaseModel: finetune_language_layers = True, finetune_attention_modules = True, finetune_mlp_modules = True, + finetune_last_n_layers = None, layers_to_transform = None, layers_pattern = None, use_gradient_checkpointing = "unsloth", @@ -1417,6 +1447,12 @@ class FastBaseModel: if target_parameters is None: target_parameters = get_moe_target_parameters(model, target_modules) + if finetune_last_n_layers is not None and layers_to_transform is None: + _total_layers = _get_total_transformer_layers(model) + if _total_layers is not None and _total_layers > 0: + n = max(1, min(int(finetune_last_n_layers), _total_layers)) + layers_to_transform = list(range(_total_layers - n, _total_layers)) + # Get only allowed parameters for LoraConfig local_variables = { **locals(),