Skip to content
Merged
Changes from 2 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: 1 addition & 5 deletions src/setfit/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -346,8 +346,6 @@ def train(
distance_metric=self.distance_metric,
margin=self.margin,
)

train_steps = len(train_dataloader) * self.num_epochs
else:
train_examples = []

Expand All @@ -363,19 +361,17 @@ def train(

train_dataloader = DataLoader(train_examples, shuffle=True, batch_size=batch_size)
train_loss = self.loss_class(self.model.model_body)
train_steps = len(train_dataloader) * num_epochs

logger.info("***** Running training *****")
logger.info(f" Num examples = {len(train_examples)}")
logger.info(f" Num epochs = {num_epochs}")
logger.info(f" Total optimization steps = {train_steps}")
logger.info(f" Total optimization steps = {len(train_dataloader) * num_epochs}")
logger.info(f" Total train batch size = {batch_size}")

warmup_steps = math.ceil(train_steps * self.warmup_proportion)

@tomaarsen tomaarsen Jan 19, 2023

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
warmup_steps = math.ceil(train_steps * self.warmup_proportion)
warmup_steps = math.ceil(train_steps * self.warmup_proportion)

This still uses the train_steps variable that you removed.

self.model.model_body.fit(
train_objectives=[(train_dataloader, train_loss)],
epochs=num_epochs,
steps_per_epoch=train_steps,
optimizer_params={"lr": learning_rate},
warmup_steps=warmup_steps,
show_progress_bar=show_progress_bar,
Expand Down