unsloth/studio/backend/tests/test_mlx_training_worker_config.py
DoubleMathew f372da407b
MLX Training updates (#5656)
* Expose MLX grad value clipping in Studio

* update test

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* dataset ordering + wd

* fix mlx smoke step expectations

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* cast norm activation output back to original input dtype

* address mlx studio review feedback

* Fix present-but-None seed override for PR #5656

studio/backend/core/training/worker.py
  `config.get("model_random_state", random_seed)` only fills the
  default when the key is absent. When a caller passes
  `config["model_random_state"] = None` explicitly (which happens
  any time a JSON payload sends an explicit `null`), the old code
  forwarded `None` to FastMLXModel and disabled deterministic init
  silently. Same for `lora_random_state`. Treat absent and explicit
  None the same way: fall back to random_seed.

studio/backend/tests/test_training_raw_support.py
  Update the source-string assertions to match the new lines.

* Guard optional MLXTrainingConfig fields and normalize random_seed for PR #5656

The MLX worker now passes `cast_norm_output_to_input_dtype` and
`dataset_order` only when the linked unsloth-zoo dataclass actually
declares them. Released zoo trees that predate the paired PR can still
construct `MLXTrainingConfig` without raising
`TypeError: unexpected keyword argument`. Once the dependency floor is
bumped to a release that contains both fields, the feature-detect
guards become no-ops.

`random_seed = config.get("random_seed", 3407)` was unguarded against
explicit `None` from raw / backend callers. The same value seeded the
trainer and was the fallback target for `model_random_state` /
`lora_random_state`. Normalize once at the top of the function and use
the normalized value everywhere so an explicit `None` cannot reach
FastMLXModel / get_peft_model / MLXTrainingConfig.

Existing seed source-pattern test updated to match the new normalize
helper. New test asserts the feature-detection guards exist and that
the unconditional kwargs do not include the gated fields.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Normalize seed / cast / max_grad_value at TrainingBackend for PR #5656

Round-3 review consensus: the per-field guards that landed in the MLX
worker only protect the MLX path. The same `TrainingBackend.start_training`
config still reaches the CUDA/text trainer at `worker.py:2267`, the
embedding LoRA init at `worker.py:2450`, and embedding TrainingArguments
at `worker.py:2624` with raw `None` values, so an explicit
`random_seed=None` from a raw / backend caller still breaks non-MLX
training even after the previous fix.

Move the normalization into `TrainingBackend.start_training` itself,
where it runs once for every training mode:

- `_coerce_seed(value)`: explicit `None`, non-int, or absent all become
  3407. Every downstream worker now sees an int.
- `_coerce_optional_bool(value, default)`: explicit `None` falls back
  to `default` instead of `bool(None) == False`. Also normalizes the
  common raw-config / YAML string aliases ("true" / "false" / "0" /
  "1"). Used for `cast_norm_output_to_input_dtype`.
- `_coerce_optional_nonneg_float(name, value)`: rejects negative
  numerics from raw / backend callers, matching the Pydantic
  `ge=0` constraint the HTTP route already enforces. Used for
  `max_grad_value`.

worker.py MLX path: the existing `bool(config.get(key, True))` for
`cast_norm_output_to_input_dtype` was changed to also fall back on
explicit `None`, so direct worker callers (bypassing
`TrainingBackend.start_training`) are equally safe. `max_grad_value`
also raises on negative values inside the worker for the same reason.

TrainingStartRequest.random_seed default bumped from 42 to 3407 so
direct REST callers that omit the field receive the same default as
the Studio frontend and the MLX worker.

New regression test exercises the three new helpers across explicit
None, valid values, string aliases, and negative-value rejection.

* Tighten feature-detect test paren tracking for PR #5656

The block-extraction used , which stops at the
first inner closing paren (e.g. )
and would silently miss a future unconditional
/  added later in the same dict literal. Switched to
proper paren-depth tracking so the unconditional block is checked end-to-end.

* Shorten verbose comments in MLX Studio backend

* Handle MLX Studio EOS appending by mode

* Wire MLX leaf norm clipping through Studio

* Respect VLM layer filters for explicit LoRA targets

Rationale / guardrails for the local Studio/vision push:

When callers provide explicit VLM LoRA target_modules together with layer filters, FastVisionModel still needs to route the explicit targets through get_peft_regex. Otherwise the layer filters are ignored and adapters can be attached outside the requested language/vision scope.

Do not revert this to plain list(target_modules) for explicit module lists. The CUDA/Studio-facing contract is that explicit targets and layer filters compose: target_modules selects module names, while finetune_language_layers / finetune_vision_layers / finetune_attention_modules / finetune_mlp_modules constrain where those targets are allowed.

The regression test covers the language-only explicit q_proj case and source-checks that explicit targets are wrapped through get_peft_regex when filters are active.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Refresh MLX smoke clip-config note for leaf_norm default

Trim the 11-line comment block to 5 lines and correct the stale claim
that MLXTrainingConfig defaults to max_grad_value=1.0. The new default
is max_grad_leaf_norm=1.0 (same memory profile as elementwise but
direction-preserving). The smoke still pins max_grad_value=1.0
explicitly to keep the 13-seed pass-rate fixture stable.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Forward max_grad_leaf_norm through the training route and warn when layer filters constrain explicit target_modules for PR #5656

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han-Chen <info@unsloth.ai>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
2026-06-14 04:58:50 -07:00

223 lines
7.3 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
import importlib.util
import sys
import types
from pathlib import Path
import pytest
def _load_worker_module():
stub_names = (
"structlog",
"loggers",
"utils",
"utils.hardware",
"utils.wheel_utils",
)
previous_modules = {name: sys.modules.get(name) for name in stub_names}
try:
sys.modules["structlog"] = types.ModuleType("structlog")
loggers = types.ModuleType("loggers")
loggers.get_logger = lambda *_args, **_kwargs: None
sys.modules["loggers"] = loggers
utils = types.ModuleType("utils")
utils.__path__ = []
sys.modules["utils"] = utils
hardware = types.ModuleType("utils.hardware")
hardware.apply_gpu_ids = lambda *_args, **_kwargs: None
sys.modules["utils.hardware"] = hardware
wheel_utils = types.ModuleType("utils.wheel_utils")
for name in (
"direct_wheel_url",
"flash_attn_wheel_url",
"has_blackwell_gpu",
"install_wheel",
"probe_torch_wheel_env",
"url_exists",
):
setattr(wheel_utils, name, lambda *_args, **_kwargs: None)
sys.modules["utils.wheel_utils"] = wheel_utils
worker_path = Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py"
spec = importlib.util.spec_from_file_location("mlx_training_worker_under_test", worker_path)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module
finally:
for name, module in previous_modules.items():
if module is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = module
_worker = _load_worker_module()
_normalize_mlx_studio_optimizer = _worker._normalize_mlx_studio_optimizer
_normalize_mlx_studio_scheduler = _worker._normalize_mlx_studio_scheduler
_mlx_vlm_max_resized_size = _worker._mlx_vlm_max_resized_size
_mlx_vlm_resized_image_layout = _worker._mlx_vlm_resized_image_layout
_copy_mlx_vlm_image_processor = _worker._copy_mlx_vlm_image_processor
_resize_mlx_vlm_image = _worker._resize_mlx_vlm_image
_adapt_for_mlx_vlm = _worker._adapt_for_mlx_vlm
def test_mlx_studio_optimizer_aliases_are_explicit():
assert _normalize_mlx_studio_optimizer("adamw_8bit") == "adamw"
assert _normalize_mlx_studio_optimizer("paged_adamw_8bit") == "adamw"
assert _normalize_mlx_studio_optimizer("adafactor") == "adafactor"
def test_mlx_studio_rejects_unknown_optimizer():
with pytest.raises(ValueError, match = "Unsupported optimizer for MLX training"):
_normalize_mlx_studio_optimizer("adamw_typo")
def test_mlx_studio_rejects_unknown_scheduler():
with pytest.raises(ValueError, match = "Unsupported LR scheduler for MLX training"):
_normalize_mlx_studio_scheduler("linear_typo")
def test_mlx_studio_keeps_hf_style_tokenizer_dual_purpose():
source = (Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py").read_text()
assert "tokenizer = tokenizer" in source
assert "processor = tokenizer if is_vlm else None" not in source
def test_mlx_vlm_resize_uses_max_dimension_like_torch_trainer():
assert _mlx_vlm_max_resized_size(1000, 500, 512) == (512, 256)
assert _mlx_vlm_max_resized_size(500, 1000, 512) == (256, 512)
assert _mlx_vlm_max_resized_size(1000, 1000, 512) == (512, 512)
assert _mlx_vlm_max_resized_size(256, 128, 1536) == (256, 128)
assert _mlx_vlm_max_resized_size(512, 256, 512) == (512, 256)
# Half-pixel cases must match the Torch collator (not banker's round).
assert _mlx_vlm_max_resized_size(333, 1000, 500) == (167, 500)
assert _mlx_vlm_max_resized_size(1000, 333, 500) == (500, 167)
def test_mlx_vlm_resize_keeps_default_numpy_layout_hwc():
Image = pytest.importorskip("PIL.Image")
image = Image.new("RGB", (320, 200), color = (10, 20, 30))
resized = _resize_mlx_vlm_image(image, 128)
assert resized.shape == (80, 128, 3)
assert resized.flags.c_contiguous
def test_mlx_vlm_resize_uses_requested_chw_numpy_layout():
Image = pytest.importorskip("PIL.Image")
image = Image.new("RGB", (320, 200), color = (10, 20, 30))
resized = _resize_mlx_vlm_image(image, 128, image_layout = "chw")
assert resized.shape == (3, 80, 128)
assert resized.flags.c_contiguous
def test_mlx_vlm_resized_image_layout_probes_processor_contract():
class ChwOnlyImageProcessor:
def __call__(self, images = None):
image = images[0]
if image.shape[0] == 3:
return {"pixel_values": image}
raise ValueError("expected CHW")
class HwcImageProcessor:
def __call__(self, images = None):
image = images[0]
if image.shape[-1] == 3:
return {"pixel_values": image}
raise ValueError("expected HWC")
assert (
_mlx_vlm_resized_image_layout(
types.SimpleNamespace(image_processor = ChwOnlyImageProcessor())
)
== "chw"
)
assert (
_mlx_vlm_resized_image_layout(types.SimpleNamespace(image_processor = HwcImageProcessor()))
is None
)
def test_mlx_vlm_layout_probe_copies_image_processor():
class StatefulImageProcessor:
def __init__(self):
self.calls = 0
def __call__(self, images = None):
self.calls += 1
image = images[0]
if image.shape[0] == 3:
return {"pixel_values": image}
raise ValueError("expected CHW")
image_processor = StatefulImageProcessor()
layout = _mlx_vlm_resized_image_layout(types.SimpleNamespace(image_processor = image_processor))
assert layout == "chw"
assert image_processor.calls == 0
def test_mlx_vlm_image_processor_copy_refuses_uncopyable_processors():
class UncopyableImageProcessor:
def __copy__(self):
raise RuntimeError("no copy")
def __deepcopy__(self, _memo):
raise RuntimeError("no deepcopy")
image_processor = UncopyableImageProcessor()
assert _copy_mlx_vlm_image_processor(image_processor) is None
def test_mlx_vlm_layout_probe_skips_uncopyable_processors():
class UncopyableImageProcessor:
def __copy__(self):
raise RuntimeError("no copy")
def __deepcopy__(self, _memo):
raise RuntimeError("no deepcopy")
def __call__(self, images = None):
raise AssertionError("live processor should not be probed")
assert (
_mlx_vlm_resized_image_layout(
types.SimpleNamespace(image_processor = UncopyableImageProcessor())
)
is None
)
def test_mlx_vlm_adapter_applies_chw_layout_to_message_images():
Image = pytest.importorskip("PIL.Image")
image = Image.new("RGB", (320, 200), color = (10, 20, 30))
item = {
"messages": [
{
"role": "user",
"content": [
{"type": "image", "image": image},
{"type": "text", "text": "Describe it."},
],
}
]
}
adapted = _adapt_for_mlx_vlm([item], resize = 128, image_layout = "chw")
assert adapted[0]["image"].shape == (3, 80, 128)
assert adapted[0]["messages"][0]["content"][0] == {"type": "image"}