diff --git a/src/setfit/trainer.py b/src/setfit/trainer.py index ba035d45..c0acaf7f 100644 --- a/src/setfit/trainer.py +++ b/src/setfit/trainer.py @@ -49,7 +49,11 @@ def __init__( self.logs_prefix = "embedding" super().__init__( model=setfit_model.model_body, - args=SentenceTransformerTrainingArguments(output_dir=setfit_args.output_dir), + # `report_to` must be set before the superclass initializes the reporting callbacks, otherwise + # integrations that the user excluded are still loaded and executed. + args=SentenceTransformerTrainingArguments( + output_dir=setfit_args.output_dir, report_to=setfit_args.report_to + ), **kwargs, ) self._apply_training_arguments(setfit_args) diff --git a/tests/test_trainer.py b/tests/test_trainer.py index 923d8527..e4c5cd11 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -8,11 +8,11 @@ import torch from datasets import Dataset, load_dataset from sentence_transformers import losses -from transformers import TrainerCallback +from transformers import TrainerCallback, integrations from transformers.testing_utils import require_optuna from transformers.utils.hp_naming import TrialShortNamer -from setfit import logging +from setfit import logging, training_args from setfit.losses import SupConLoss from setfit.modeling import SetFitModel from setfit.trainer import Trainer @@ -557,6 +557,25 @@ class TestCallback(TrainerCallback): assert callback not in trainer.st_trainer.callback_handler.callbacks +def test_trainer_report_to(model: SetFitModel, monkeypatch: pytest.MonkeyPatch): + # Issue #621: integrations that were not requested via `report_to` must not be loaded + class DummyReportCallback(TrainerCallback): + pass + + monkeypatch.setattr(training_args, "get_available_reporting_integrations", lambda: ["dummy"]) + monkeypatch.setitem(integrations.INTEGRATION_TO_CALLBACK, "dummy", DummyReportCallback) + + def reports_to_dummy(args: TrainingArguments) -> bool: + trainer = Trainer(model=model, args=args) + return any( + isinstance(callback, DummyReportCallback) for callback in trainer.st_trainer.callback_handler.callbacks + ) + + assert reports_to_dummy(TrainingArguments(report_to="all")) + assert reports_to_dummy(TrainingArguments(report_to="dummy")) + assert not reports_to_dummy(TrainingArguments(report_to="none")) + + def test_trainer_warn_freeze(model: SetFitModel): trainer = Trainer(model) with pytest.warns(