Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/setfit/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
23 changes: 21 additions & 2 deletions tests/test_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down