nemo_automodel.recipes.llm.train_seq_cls

View as Markdown

Module Contents

Classes

NameDescription
TrainFinetuneRecipeForSequenceClassificationRecipe for fine-tuning a model for sequence classification.

Functions

NameDescription
mainRun the sequence-classification fine-tuning recipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification(
cfg
)

Bases: BaseRecipe

Recipe for fine-tuning a model for sequence classification.

cfg
nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification._run_train_optim_step(
batches
)
nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification._validate_one_epoch(
dataloader
)
nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification.log_train_metrics(
log_data
)

Log metrics to wandb and other loggers.

Parameters:

log_data

MetricsSample object, containing: step: int, the current step. epoch: int, the current epoch. metrics: Dict[str, float], containing: “loss”: Training loss. “accuracy”: Training accuracy. “grad_norm”: Gradient norm from the training step. “lr”: Learning rate. “mem”: Memory allocated. “tps”: Tokens per second (throughput). “tps_per_gpu”: Tokens per second per GPU.

nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification.log_val_metrics(
log_data
)

Log metrics to wandb and other loggers Args: log_data: MetricsSample object, containing: step: int, the current step. epoch: int, the current epoch. metrics: Dict[str, float], containing: “val_loss”: Validation loss. “lr”: Learning rate. “num_label_tokens”: Number of label tokens. “mem”: Memory allocated.

nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification.run_train_validation_loop()
nemo_automodel.recipes.llm.train_seq_cls.TrainFinetuneRecipeForSequenceClassification.setup()
nemo_automodel.recipes.llm.train_seq_cls.main(
config_path: str | None = None
)

Run the sequence-classification fine-tuning recipe.

nemo_automodel.recipes.llm.train_seq_cls.logger = logging.getLogger(__name__)