diff --git a/src/setfit/trainer.py b/src/setfit/trainer.py index ba035d45..7994a424 100644 --- a/src/setfit/trainer.py +++ b/src/setfit/trainer.py @@ -579,7 +579,19 @@ def train_embeddings( SupConLoss, ): self.st_trainer.args.batch_sampler = BatchSamplers.GROUP_BY_LABEL - self.st_trainer.train() + + # `args.max_length` is honored in the classifier phase (via `SetFitModel.fit`), + # but was previously ignored while finetuning the embedding body. Apply it to the + # body's truncation length for this phase too, clamping to the model's maximum as + # `_prepare_dataloader` does, and restore the original afterwards so that encoding + # at inference time is unaffected. + original_max_seq_length = self.model.model_body.max_seq_length + if args.max_length is not None: + self.model.model_body.max_seq_length = min(args.max_length, self.model.model_body.get_max_seq_length()) + try: + self.st_trainer.train() + finally: + self.model.model_body.max_seq_length = original_max_seq_length def get_dataset( self, x: List[str], y: Union[List[int], List[List[int]]], args: TrainingArguments, max_pairs: int = -1 diff --git a/tests/test_trainer.py b/tests/test_trainer.py index 923d8527..66125570 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -603,3 +603,28 @@ def test_trainer_wrong_args(model: SetFitModel) -> None: expected = "`args` must be a `TrainingArguments` instance imported from `setfit`." with pytest.raises(ValueError, match=expected): Trainer(model, dataset) + + +def test_trainer_max_length_applied_to_embedding_phase(model: SetFitModel) -> None: + # Regression test for #561: `TrainingArguments.max_length` was honored when fitting the + # classifier head but silently ignored when finetuning the embedding body. + dataset = Dataset.from_dict({"text": ["a", "b", "c"], "label": [0, 1, 2]}) + max_length = 32 + assert max_length < model.model_body.get_max_seq_length() + original_max_seq_length = model.model_body.max_seq_length + + args = TrainingArguments(num_iterations=1, max_length=max_length) + trainer = Trainer(model, args=args, train_dataset=dataset) + + observed = {} + + def fake_train() -> None: + observed["max_seq_length"] = model.model_body.max_seq_length + + trainer.st_trainer.train = fake_train + trainer.train_embeddings(dataset["text"], dataset["label"], args=args) + + # The body is truncated to `max_length` while the embeddings are trained, and the + # original truncation length is restored afterwards so inference is unaffected. + assert observed["max_seq_length"] == max_length + assert model.model_body.max_seq_length == original_max_seq_length