feat(mlx): route trainer callbacks (#6929)
This commit is contained in:
parent
2a6abe2ff5
commit
934f879043
2 changed files with 96 additions and 6 deletions
|
|
@ -115,6 +115,9 @@ def test_mlx_training_arguments_accept_trl_style_kwargs():
|
|||
def test_mlx_training_arguments_do_not_warn_for_implemented_or_falsey_extras():
|
||||
"""Implemented and falsey inert compatibility kwargs should stay quiet."""
|
||||
unsloth = _import_mlx_unsloth()
|
||||
supported_eval_kwargs = {}
|
||||
if "eval_strategy" in unsloth._MLX_TRAINING_CONFIG_FIELDS:
|
||||
supported_eval_kwargs = {"eval_strategy": "no", "eval_delay": 1}
|
||||
|
||||
with warnings.catch_warnings(record = True) as caught:
|
||||
warnings.simplefilter("always")
|
||||
|
|
@ -125,12 +128,16 @@ def test_mlx_training_arguments_do_not_warn_for_implemented_or_falsey_extras():
|
|||
remove_unused_columns = False,
|
||||
assistant_only_loss = False,
|
||||
completion_only_loss = False,
|
||||
**supported_eval_kwargs,
|
||||
)
|
||||
|
||||
assert args.warmup_steps == 2
|
||||
assert args.padding_free is False
|
||||
assert args.remove_unused_columns is False
|
||||
assert args.completion_only_loss is False
|
||||
if supported_eval_kwargs:
|
||||
assert args.eval_strategy == "no"
|
||||
assert args.eval_delay == 1
|
||||
assert caught == []
|
||||
|
||||
|
||||
|
|
@ -705,17 +712,10 @@ def test_mlx_trainer_rejects_unsafe_unsupported_sft_kwargs():
|
|||
)
|
||||
|
||||
|
||||
def test_mlx_trainer_rejects_metrics_and_callbacks():
|
||||
"""Trainer hooks should fail because MLXTrainer cannot honor them yet."""
|
||||
def test_mlx_trainer_rejects_compute_metrics():
|
||||
"""compute_metrics is still unsupported by MLXTrainer."""
|
||||
unsloth = _import_mlx_unsloth()
|
||||
|
||||
with pytest.raises(NotImplementedError, match = "callbacks"):
|
||||
unsloth.UnslothTrainer(
|
||||
model = _DummyModel(),
|
||||
tokenizer = None,
|
||||
train_dataset = [],
|
||||
callbacks = [object()],
|
||||
)
|
||||
with pytest.raises(NotImplementedError, match = "compute_metrics"):
|
||||
unsloth.UnslothTrainer(
|
||||
model = _DummyModel(),
|
||||
|
|
@ -725,6 +725,46 @@ def test_mlx_trainer_rejects_metrics_and_callbacks():
|
|||
)
|
||||
|
||||
|
||||
def test_mlx_trainer_accepts_callbacks():
|
||||
"""Callbacks are routed to MLXTrainer when the zoo backend supports them."""
|
||||
unsloth = _import_mlx_unsloth()
|
||||
from transformers import TrainerCallback
|
||||
|
||||
if not unsloth._mlx_trainer_supports_kwarg("callbacks"):
|
||||
pytest.skip("requires unsloth-zoo MLXTrainer callback support")
|
||||
|
||||
class Callback(TrainerCallback):
|
||||
pass
|
||||
|
||||
trainer = unsloth.UnslothTrainer(
|
||||
model = _DummyModel(),
|
||||
tokenizer = None,
|
||||
train_dataset = [],
|
||||
callbacks = [Callback()],
|
||||
)
|
||||
assert any(isinstance(cb, Callback) for cb in trainer.callback_handler.callbacks)
|
||||
|
||||
|
||||
def test_mlx_trainer_rejects_callbacks_with_old_zoo(monkeypatch):
|
||||
"""Older unsloth-zoo builds should fail clearly instead of TypeError."""
|
||||
unsloth = _import_mlx_unsloth()
|
||||
from transformers import TrainerCallback
|
||||
|
||||
monkeypatch.setattr(
|
||||
unsloth,
|
||||
"_mlx_trainer_supports_kwarg",
|
||||
lambda name: name != "callbacks",
|
||||
)
|
||||
|
||||
with pytest.raises(NotImplementedError, match = "callbacks require"):
|
||||
unsloth.UnslothTrainer(
|
||||
model = _DummyModel(),
|
||||
tokenizer = None,
|
||||
train_dataset = [],
|
||||
callbacks = [TrainerCallback()],
|
||||
)
|
||||
|
||||
|
||||
def test_mlx_trainer_rejects_custom_data_collator():
|
||||
"""MLXTrainer owns batching; custom SFT data collators must not be ignored."""
|
||||
unsloth = _import_mlx_unsloth()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue