From 23da52455450e21018ee59ed3ac6c654f7bb95f8 Mon Sep 17 00:00:00 2001 From: Darsh Date: Sat, 4 Jul 2026 14:32:52 +0530 Subject: [PATCH] fix: reuse shared vocabulary expansion in vision/audio SFT (#291) --- src/soup_cli/trainer/sft.py | 12 +++ tests/test_trainer_init.py | 190 ++++++++++++++++++++++++++++++++++++ 2 files changed, 202 insertions(+) diff --git a/src/soup_cli/trainer/sft.py b/src/soup_cli/trainer/sft.py index bae5a7e..27c17eb 100644 --- a/src/soup_cli/trainer/sft.py +++ b/src/soup_cli/trainer/sft.py @@ -937,7 +937,13 @@ class SFTTrainerWrapper: model_kwargs["quantization_config"] = quant_config_obj self.model = AutoModelForVision2Seq.from_pretrained(cfg.base, **model_kwargs) + from soup_cli.utils.data_pipeline import apply_vocab_expansion + apply_vocab_expansion( + self.processor.tokenizer, + self.model, + cfg.data, + ) if tcfg.quantization in ("4bit", "8bit", "mxfp4"): self.model = prepare_model_for_kbit_training(self.model) @@ -1042,7 +1048,13 @@ class SFTTrainerWrapper: # Use AutoModel for audio models — AutoModelForCausalLM doesn't handle # audio-language architectures (Qwen2-Audio, Whisper, etc.) self.model = AutoModel.from_pretrained(cfg.base, **model_kwargs) + from soup_cli.utils.data_pipeline import apply_vocab_expansion + apply_vocab_expansion( + self.processor.tokenizer, + self.model, + cfg.data, + ) if tcfg.quantization in ("4bit", "8bit", "mxfp4"): self.model = prepare_model_for_kbit_training(self.model) diff --git a/tests/test_trainer_init.py b/tests/test_trainer_init.py index 4567ce3..b5d3d14 100644 --- a/tests/test_trainer_init.py +++ b/tests/test_trainer_init.py @@ -163,7 +163,197 @@ class TestSFTTrainerInit: assert calls["add_tokens"] == ["", ""] assert calls["add_special_tokens"] == [""] assert calls["resize"] == len(tokenizer) + def test_vision_vocab_expansion_adds_tokens_and_resizes(self, monkeypatch): + """Vision SFT should apply configured vocabulary expansion.""" + import sys + import types + from types import SimpleNamespace + + from soup_cli.trainer.sft import SFTTrainerWrapper + + cfg = _make_config( + modality="vision", + data={ + "train": "./data.jsonl", + "format": "llava", + "add_new_tokens": ["", ""], + "new_special_tokens": [""], + "resize_vocab": True, + }, + training={"quantization": "none"}, + ) + + calls = { + "add_tokens": None, + "add_special_tokens": None, + "resize": None, + } + + class _Tokenizer: + def __init__(self): + self.vocab = {"": 0} + + def get_vocab(self): + return self.vocab + + def add_tokens(self, tokens): + calls["add_tokens"] = list(tokens) + for t in tokens: + self.vocab.setdefault(t, len(self.vocab)) + return len(tokens) + + def add_special_tokens(self, data): + tokens = list(data["additional_special_tokens"]) + calls["add_special_tokens"] = tokens + for t in tokens: + self.vocab.setdefault(t, len(self.vocab)) + return len(tokens) + + def __len__(self): + return len(self.vocab) + + class _Processor: + def __init__(self): + self.tokenizer = _Tokenizer() + + class _Model: + config = SimpleNamespace() + + def resize_token_embeddings(self, size): + calls["resize"] = size + + def parameters(self): + return [] + + processor = _Processor() + model = _Model() + + fake_transformers = types.SimpleNamespace( + AutoProcessor=types.SimpleNamespace(from_pretrained=lambda *a, **k: processor), + AutoModelForVision2Seq=types.SimpleNamespace(from_pretrained=lambda *a, **k: model), + ) + + fake_peft = types.SimpleNamespace( + LoraConfig=lambda **kwargs: SimpleNamespace(**kwargs), + get_peft_model=lambda m, cfg: m, + prepare_model_for_kbit_training=lambda m: m, + ) + + monkeypatch.setitem(sys.modules, "transformers", fake_transformers) + monkeypatch.setitem(sys.modules, "peft", fake_peft) + + monkeypatch.setattr( + "soup_cli.utils.quant_menu.build_quantization_config_for_loader", + lambda **kwargs: None, + ) + + wrapper = object.__new__(SFTTrainerWrapper) + wrapper.config = cfg + wrapper.device = "cpu" + wrapper._trust_remote_code = False + + wrapper._setup_vision_transformers(cfg, cfg.training) + + assert calls["add_tokens"] == [""] + assert calls["add_special_tokens"] == [""] + assert calls["resize"] == len(processor.tokenizer) + + def test_audio_vocab_expansion_adds_tokens_and_resizes(self, monkeypatch): + """Audio SFT should apply configured vocabulary expansion.""" + + import sys + import types + from types import SimpleNamespace + + from soup_cli.trainer.sft import SFTTrainerWrapper + + cfg = _make_config( + modality="audio", + data={ + "train": "./data.jsonl", + "format": "audio", + "add_new_tokens": ["", ""], + "new_special_tokens": [""], + "resize_vocab": True, + }, + training={"quantization": "none"}, + ) + + calls = { + "add_tokens": None, + "add_special_tokens": None, + "resize": None, + } + + class _Tokenizer: + def __init__(self): + self.vocab = {"": 0} + + def get_vocab(self): + return self.vocab + + def add_tokens(self, tokens): + calls["add_tokens"] = list(tokens) + for t in tokens: + self.vocab.setdefault(t, len(self.vocab)) + return len(tokens) + + def add_special_tokens(self, data): + tokens = list(data["additional_special_tokens"]) + calls["add_special_tokens"] = tokens + for t in tokens: + self.vocab.setdefault(t, len(self.vocab)) + return len(tokens) + + def __len__(self): + return len(self.vocab) + + class _Processor: + def __init__(self): + self.tokenizer = _Tokenizer() + + class _Model: + config = SimpleNamespace() + + def resize_token_embeddings(self, size): + calls["resize"] = size + + def parameters(self): + return [] + + processor = _Processor() + model = _Model() + + fake_transformers = types.SimpleNamespace( + AutoProcessor=types.SimpleNamespace(from_pretrained=lambda *a, **k: processor), + AutoModel=types.SimpleNamespace(from_pretrained=lambda *a, **k: model), + ) + + fake_peft = types.SimpleNamespace( + LoraConfig=lambda **kwargs: SimpleNamespace(**kwargs), + get_peft_model=lambda m, cfg: m, + prepare_model_for_kbit_training=lambda m: m, + ) + + monkeypatch.setitem(sys.modules, "transformers", fake_transformers) + monkeypatch.setitem(sys.modules, "peft", fake_peft) + + monkeypatch.setattr( + "soup_cli.utils.quant_menu.build_quantization_config_for_loader", + lambda **kwargs: None, + ) + + wrapper = object.__new__(SFTTrainerWrapper) + wrapper.config = cfg + wrapper.device = "cpu" + wrapper._trust_remote_code = False + + wrapper._setup_audio_transformers(cfg, cfg.training) + + assert calls["add_tokens"] == [""] + assert calls["add_special_tokens"] == [""] + assert calls["resize"] == len(processor.tokenizer) class TestDPOTrainerInit: """Test DPOTrainerWrapper constructor."""