Transformers 文档

回调

Hugging Face's logo
加入 Hugging Face 社区

并获得增强的文档体验

开始使用

回调函数

回调(Callbacks)是能够自定义 PyTorch Trainer 中训练循环行为的对象。它们可以检查训练循环的状态(用于进度报告、TensorBoard 或其他机器学习平台上的日志记录等),并做出决定(如提前停止)。

除了它们返回的 TrainerControl 对象外,回调是“只读”的代码片段,它们不能更改训练循环中的任何内容。对于需要更改训练循环的自定义需求,你应该继承 Trainer 并覆盖所需的方法(示例请参阅 trainer)。

默认情况下,TrainingArguments.report_to 设置为 "none"

实现回调的主要类是 TrainerCallback。它获取用于实例化 TrainerTrainingArguments,可以通过 TrainerState 访问该 Trainer 的内部状态,并可以通过 TrainerControl 对训练循环采取一些操作。

可用回调

以下是库中可用的 TrainerCallback 列表

class transformers.integrations.CometCallback

< >

( )

一个将日志发送到 Comet MLTrainerCallback

setup

< >

( args state model )

设置可选的 Comet 集成。

环境

  • COMET_MODE (str, 可选, 默认为 get_or_create):控制是创建并记录到新的 Comet 实验,还是附加到现有的实验。它接受以下值:
    • get_or_create:根据是否设置了 COMET_EXPERIMENT_KEY 以及使用该 key 的实验是否已存在,自动决定。
    • create:总是创建一个新的 Comet 实验。
    • get:总是尝试附加到一个现有的 Comet 实验。需要设置 COMET_EXPERIMENT_KEY
  • COMET_START_ONLINE (bool, 可选):是否创建在线或离线实验。
  • COMET_PROJECT_NAME (str, 可选):Comet 实验的项目名称。
  • COMET_LOG_ASSETS (str, 可选, 默认为 TRUE):是否将训练资产(检查点等)记录到 Comet。可以是 TRUEFALSE

有关环境中可配置项的列表,请参阅 此处

class transformers.DefaultFlowCallback

< >

( )

一个处理日志、评估和检查点训练循环默认流程的 TrainerCallback

class transformers.PrinterCallback

< >

( )

一个仅打印日志的基础 TrainerCallback

class transformers.ProgressCallback

< >

( max_str_len: int = 100 )

一个显示训练或评估进度的 TrainerCallback。你可以修改 max_str_len 来控制记录日志时字符串被截断的长度。

class transformers.EarlyStoppingCallback

< >

( early_stopping_patience: int = 1 early_stopping_threshold: float | None = 0.0 )

参数

  • early_stopping_patience (int) — 与 metric_for_best_model 一起使用,当指定指标在 early_stopping_patience 次评估调用后没有改善时,停止训练。
  • early_stopping_threshold (float, 可选) — 与 TrainingArguments 的 metric_for_best_modelearly_stopping_patience 一起使用,表示指定指标需要改善多少才能满足提前停止条件。

一个处理提前停止的 TrainerCallback

此回调依赖于 TrainingArguments 参数 load_best_model_at_end 功能来设置 TrainerState 中的 best_metric。注意,如果 TrainingArguments 参数 save_stepseval_steps 不同,提前停止将在下一次保存步骤之前不会发生。

class transformers.integrations.TensorBoardCallback

< >

( tb_writer = None )

参数

  • tb_writer (SummaryWriter, 可选) — 要使用的 writer。如果未设置,将实例化一个。

一个将日志发送到 TensorBoardTrainerCallback

环境

  • TENSORBOARD_LOGGING_DIR (str, 可选, 默认为 None):记录结果的日志目录。默认值为 os.path.join(args.output_dir, default_logdir())

class transformers.integrations.TrackioCallback

< >

( )

一个将指标记录到 Trackio 的 TrainerCallback

setup

< >

( args state model **kwargs )

设置可选的 Trackio 集成。

要自定义设置,你也可以在 TrainingArguments 中设置 projecttrackio_space_idtrackio_bucket_idtrackio_static_space_idhub_private_repo

class transformers.integrations.WandbCallback

< >

( )

一个将指标、媒体、模型检查点记录到 Weights and BiasesTrainerCallback

setup

< >

( args state model **kwargs )

设置可选的 Weights & Biases (wandb) 集成。

如果需要,可以继承并覆盖此方法以自定义设置。更多信息请参阅 这里。你还可以覆盖以下环境变量:

环境

  • WANDB_LOG_MODEL (str, 可选, 默认为 "false"):是否在训练期间记录模型和检查点。可以是 "end""checkpoint""false"。如果设置为 "end",模型将在训练结束时上传。如果设置为 "checkpoint",检查点将每 args.save_steps 次上传一次。如果设置为 "false",模型将不会上传。与 load_best_model_at_end() 一起使用以上传最佳模型。
  • WANDB_WATCH (str, 可选, 默认为 "false"):可以是 "gradients""all""parameters""false"。设置为 "all" 以记录梯度和参数。
  • WANDB_PROJECT (str, 可选, 默认为 "huggingface"):设置为自定义字符串以将结果存储在不同的项目中。

class transformers.integrations.MLflowCallback

< >

( )

一个将日志发送到 MLflowTrainerCallback。可以通过设置环境变量 DISABLE_MLFLOW_INTEGRATION = TRUE 来禁用。

setup

< >

( args state model )

设置可选的 MLflow 集成。

环境

  • HF_MLFLOW_LOG_ARTIFACTS (str, 可选):是否使用 MLflow .log_artifact() 工具来记录制品。这仅在记录到远程服务器(例如 s3 或 GCS)时才有意义。如果设置为 True1,将在 TrainingArgumentsoutput_dir 中每次保存检查点时,将每个已保存的检查点复制到本地或远程制品存储中。在没有远程存储的情况下使用它,只会将文件复制到你的制品位置。
  • MLFLOW_TRACKING_URI (str, 可选):是否将运行存储在特定路径或远程服务器上。默认未设置,这将完全跳过设置跟踪 URI。
  • MLFLOW_EXPERIMENT_NAME (str, 可选, 默认为 None):是否使用 MLflow experiment_name 来启动运行。默认值为 None,这将指向 MLflow 中的 Default 实验。否则,它是要激活的实验的大小写敏感名称。如果具有此名称的实验不存在,则创建一个具有此名称的新实验。
  • MLFLOW_TAGS (str, 可选):一个键/值对字典的字符串转储,作为标签添加到 MLflow 运行中。示例: os.environ['MLFLOW_TAGS']='{"release.candidate": "RC1", "release.version": "2.2.0"}'
  • MLFLOW_NESTED_RUN (str, 可选):是否使用 MLflow 嵌套运行。如果设置为 True1,将在当前运行内创建一个嵌套运行。
  • MLFLOW_RUN_ID (str, 可选):允许重新附加到现有的运行,这在从检查点恢复训练时很有用。当设置了 MLFLOW_RUN_ID 环境变量时,start_run 会尝试恢复具有指定运行 ID 的运行,其他参数将被忽略。
  • MLFLOW_FLATTEN_PARAMS (str, 可选, 默认为 False):是否在记录之前展平参数字典。
  • MLFLOW_MAX_LOG_PARAMS (int, 可选):设置在运行中记录的最大参数数量。

class transformers.integrations.AzureMLCallback

< >

( azureml_run = None )

一个将日志发送到 AzureMLTrainerCallback

class transformers.integrations.CodeCarbonCallback

< >

( )

一个跟踪训练过程中二氧化碳排放的 TrainerCallback

class transformers.integrations.ClearMLCallback

< >

( )

一个将日志发送到 ClearMLTrainerCallback

环境

  • CLEARML_PROJECT (str, 可选, 默认为 HuggingFace Transformers): ClearML 项目名称。
  • CLEARML_TASK (str, 可选, 默认为 Trainer): ClearML 任务名称。
  • CLEARML_LOG_MODEL (bool, 可选, 默认为 False): 是否在训练过程中将模型记录为工件(artifacts)。

class transformers.integrations.DagsHubCallback

< >

( )

一个向 DagsHub 记录日志的 TrainerCallback。继承自 MLflowCallback

setup

< >

( *args **kwargs )

设置 DagsHub 的日志记录集成。

环境

  • HF_DAGSHUB_LOG_ARTIFACTS (str, 可选): 是否为实验保存数据和模型工件。默认为 False

class transformers.integrations.FlyteCallback

< >

( save_log_history: bool = True sync_checkpoints: bool = True )

参数

  • save_log_history (bool, 可选, 默认为 True) — 如果设置为 True,训练日志将保存为 Flyte Deck。
  • sync_checkpoints (bool, 可选, 默认为 True) — 如果设置为 True,检查点将与 Flyte 同步,并可在中断的情况下用于恢复训练。

一个向 Flyte 发送日志的 TrainerCallback。注意:此回调仅在 Flyte 任务中有效。

示例

# Note: This example skips over some setup steps for brevity.
from flytekit import current_context, task


@task
def train_hf_transformer():
    cp = current_context().checkpoint
    trainer = Trainer(..., callbacks=[FlyteCallback()])
    output = trainer.train(resume_from_checkpoint=cp.restore())

class transformers.integrations.KubeflowCallback

< >

( )

一个向 Kubeflow Trainer 报告训练进度的 TrainerCallback

当在启用 TrainJobRuntimeStatus 功能门的 Kubeflow TrainJob 内进行训练时,该回调会自动注册。Kubeflow 控制器会将所需的环境变量注入到训练 pod 中。

环境变量(由控制器注入)

  • KUBEFLOW_TRAINER_SERVER_URL:用于状态更新的 HTTPS 端点
  • KUBEFLOW_TRAINER_SERVER_CA_CERT:TLS 验证的 CA 证书路径
  • KUBEFLOW_TRAINER_SERVER_TOKEN:用于身份验证的服务账户令牌路径

报告信息

  • 进度百分比(0-100%)
  • 预计剩余时间(秒)
  • 训练指标(loss、learning_rate 等)

功能

  • 自动节流(每 5 秒最多 1 次更新),以避免压垮控制器
  • 令牌缓存(5 分钟),以最小化文件 I/O
  • 在分布式训练中,仅 rank 0 报告进度
  • 静默失败 - 网络问题不会中断训练

可以通过设置环境变量 DISABLE_KUBEFLOW_INTEGRATION=TRUE 来禁用。

class transformers.integrations.DVCLiveCallback

< >

( live: typing.Optional[typing.Any] = None log_model: typing.Union[typing.Literal['all'], bool, NoneType] = None **kwargs )

参数

  • live (dvclive.Live, 可选, 默认为 None) — 可选的 Live 实例。如果为 None,则会使用 **kwargs 创建一个新实例。
  • log_model (Union[Literal[“all”], bool], 可选, 默认为 None) — 是否使用 dvclive.Live.log_artifact() 来记录由 Trainer 创建的检查点。如果设置为 True,最终检查点将在训练结束时记录。如果设置为 "all"TrainingArguments 的整个 output_dir 将在每个检查点处被记录。

一个向 DVCLive 发送日志的 TrainerCallback

setup 中使用以下环境变量来配置集成。若要自定义此回调(超出这些环境变量的范围),请参阅 此处

setup

< >

( args state model )

设置可选的 DVCLive 集成。若要自定义此回调(超出以下环境变量的范围),请参阅 此处

环境

  • HF_DVCLIVE_LOG_MODEL (str, 可选): 是否使用 dvclive.Live.log_artifact() 来记录由 Trainer 创建的检查点。如果设置为 True1,最终检查点将在训练结束时记录。如果设置为 allTrainingArguments 的整个 output_dir 将在每个检查点处被记录。

class transformers.integrations.SwanLabCallback

< >

( )

一个将指标、媒体和模型检查点记录到 SwanLabTrainerCallback

setup

< >

( args state model **kwargs )

设置可选的 SwanLab (swanlab) 集成。

如果需要自定义设置,可以继承并重写此方法。更多信息请见 此处

您还可以重写以下环境变量。有关环境变量的更多信息,请见 此处

环境

  • SWANLAB_API_KEY (str, 可选, 默认为 None): 云端 API Key。在登录过程中,系统首先会检查此环境变量。如果不存在,系统会检查用户是否已登录。如果未登录,则启动登录流程。

    • 如果将字符串传递给登录接口,则忽略此环境变量。
    • 如果用户已登录,此环境变量的优先级高于本地存储的登录信息。
  • SWANLAB_PROJECT (str, 可选, 默认为 None): 设置此项为自定义字符串,以便将结果存储在不同的项目中。如果未指定,将使用当前运行目录的名称。

  • SWANLAB_LOG_DIR (str, 可选, 默认为 swanlog): 此环境变量指定在本地模式下运行时的日志文件存储路径。默认情况下,日志保存在工作目录下的 swanlog 文件夹中。

  • SWANLAB_MODE (Literal["local", "cloud", "disabled"], 可选, 默认为 cloud): SwanLab 的解析模式,涉及操作员注册的回调。目前有三种模式:local、cloud 和 disabled。注意:区分大小写。更多信息请见 此处

  • SWANLAB_LOG_MODEL (str, 可选, 默认为 None): SwanLab 目前不支持保存模式功能。此功能将在未来版本中推出。

  • SWANLAB_WEB_HOST (str, 可选, 默认为 None): 用于私有版本(免费)的 SwanLab 云环境的 Web 地址。

  • SWANLAB_API_HOST (str, 可选, 默认为 None): 用于私有版本(免费)的 SwanLab 云环境的 API 地址。

  • SWANLAB_RUN_ID (str, 可选, 默认为 None): 用于恢复之前运行的实验 ID。配合 SWANLAB_RESUME 使用以继续现有的实验。

  • SWANLAB_RESUME (str, 可选, 默认为 None): 恢复模式 ("must", "allow", "never")。当使用 resume_from_checkpoint 时,默认为 "allow"

TrainerCallback

class transformers.TrainerCallback

< >

( )

参数

  • args (TrainingArguments) — 用于实例化 Trainer 的训练参数。
  • state (TrainerState) — Trainer 的当前状态。
  • control (TrainerControl) — 返回给 Trainer 并可用于做出决策的对象。
  • model (PreTrainedModeltorch.nn.Module) — 正在训练的模型。
  • processing_class ([PreTrainedTokenizerBaseImageProcessorProcessorMixinFeatureExtractionMixin]) — 用于编码数据的处理类。可以是分词器(tokenizer)、处理器(processor)、图像处理器(image processor)或特征提取器(feature extractor)。
  • optimizer (torch.optim.Optimizer) — 用于训练步骤的优化器。
  • lr_scheduler (torch.optim.lr_scheduler.LambdaLR) — 用于设置学习率的调度器。
  • train_dataloader (torch.utils.data.DataLoader, 可选) — 当前用于训练的数据加载器 (dataloader)。
  • eval_dataloader (torch.utils.data.DataLoader, 可选) — 当前用于评估的数据加载器。
  • metrics (dict[str, float]) — 上一次评估阶段计算出的指标。

    这些指标仅可在 on_evaluate 事件中访问。

  • logs (dict[str, float]) — 要记录的值。

    这些值仅可在 on_log 事件中访问。

用于检查训练循环在某些事件节点的状态并做出相应决策的类。在每个事件中,以下参数均可用。

control 对象是唯一可以被回调修改的对象,若需要修改,事件应返回该对象的修改版本。

参数 argsstatecontrol 是所有事件的位置参数,其他所有参数都包含在 kwargs 中。您可以在事件的签名中解包所需参数。例如,请参见 PrinterCallback 的简单代码实现。

示例

class PrinterCallback(TrainerCallback):
    def on_log(self, args, state, control, logs=None, **kwargs):
        _ = logs.pop("total_flos", None)
        if state.is_local_process_zero:
            print(logs)

on_epoch_begin

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

Epoch 开始时调用的事件。

on_epoch_end

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

Epoch 结束时调用的事件。

on_evaluate

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

评估阶段之后调用的事件。

on_init_end

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

Trainer 初始化结束时调用的事件。

on_log

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

记录最近的日志后调用的事件。

on_optimizer_step

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

优化器步骤之后、梯度清零之前调用的事件。可用于监控梯度。

on_pre_optimizer_step

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

优化器步骤之前、梯度裁剪之后调用的事件。可用于监控梯度。

on_predict

< >

( args: TrainingArguments state: TrainerState control: TrainerControl metrics **kwargs )

成功预测后调用的事件。

on_prediction_step

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

预测步骤之后调用的事件。

on_push_begin

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

Trainer.push_to_hubTrainer._push_from_checkpoint 开始时,将模型推送到 hub 之前调用的事件。

on_save

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

保存检查点后调用的事件。

on_step_begin

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

训练步骤开始时调用的事件。如果使用梯度累积,一个训练步骤可能包含多个输入。

on_step_end

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

训练步骤结束时调用的事件。如果使用梯度累积,一个训练步骤可能包含多个输入。

on_substep_end

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

梯度累积期间,子步骤结束时调用的事件。

on_train_begin

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

训练开始时调用的事件。

on_train_end

< >

( args: TrainingArguments state: TrainerState control: TrainerControl **kwargs )

训练结束时调用的事件。

以下是如何向 PyTorch Trainer 注册自定义回调的示例

class MyCallback(TrainerCallback):
    "A callback that prints a message at the beginning of training"

    def on_train_begin(self, args, state, control, **kwargs):
        print("Starting training")


trainer = Trainer(
    model,
    args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    callbacks=[MyCallback],  # We can either pass the callback class this way or an instance of it (MyCallback())
)

注册回调的另一种方法是如下调用 trainer.add_callback()

trainer = Trainer(...)
trainer.add_callback(MyCallback)
# Alternatively, we can pass an instance of the callback class
trainer.add_callback(MyCallback())

TrainerState

class transformers.TrainerState

< >

( epoch: float = 0 global_step: int = 0 max_steps: int = 0 logging_steps: int = 500 eval_steps: int = 500 save_steps: int = 500 train_batch_size: int | None = None num_train_epochs: int = 0 num_input_tokens_seen: int = 0 total_flos: float = 0 log_history: list = None best_metric: float | None = None best_global_step: int | None = None best_model_checkpoint: str | None = None is_local_process_zero: bool = True is_world_process_zero: bool = True is_hyper_param_search: bool = False trial_name: str | None = None trial_params: dict[str, str | float | int | bool] | None = None stateful_callbacks: list['TrainerCallback'] | None = None )

参数

  • epoch (float, 可选) — 仅在训练期间设置,表示训练所处的 epoch(小数部分表示当前 epoch 已完成的百分比)。
  • global_step (int, 可选,默认为 0) — 训练期间,表示已完成的更新步数。
  • max_steps (int, 可选,默认为 0) — 当前训练中要执行的更新步数。
  • logging_steps (int, 可选,默认为 500) — 每隔 X 个更新步数进行一次日志记录。
  • eval_steps (int, 可选) — 每隔 X 步进行一次评估。
  • save_steps (int, 可选,默认为 500) — 每隔 X 个更新步数保存一次检查点。
  • train_batch_size (int, 可选) — 训练数据加载器的批次大小。仅在使用了 auto_find_batch_size 时需要。
  • num_input_tokens_seen (int, 可选,默认为 0) — 当跟踪输入 token 时,训练期间看到的 token 总数(指的是输入 token 的数量,而非预测 token 的数量)。
  • total_flos (float, 可选,默认为 0) — 从训练开始以来模型执行的总浮点运算数(以浮点数存储以避免溢出)。
  • log_history (list[dict[str, float]], 可选) — 自训练开始以来执行的日志列表。
  • best_metric (float, 可选) — 在跟踪最佳模型时,迄今为止遇到的最佳指标值。
  • best_global_step (int, 可选) — 在跟踪最佳模型时,遇到最佳指标时的步数。用于设置 best_model_checkpoint
  • best_model_checkpoint (str, 可选) — 在跟踪最佳模型时,迄今为止遇到的最佳模型检查点的名称值。
  • is_local_process_zero (bool, 可选,默认为 True) — 该进程是否为本地主进程(例如,如果在多台机器上进行分布式训练,则是指单台机器上的主进程)。
  • is_world_process_zero (bool, 可选,默认为 True) — 该进程是否为全局主进程(当在多台机器上进行分布式训练时,只有一个进程为 True)。
  • is_hyper_param_search (bool, 可选,默认为 False) — 是否正在使用 Trainer.hyperparameter_search 进行超参数搜索。这将影响数据记录到 TensorBoard 的方式。
  • stateful_callbacks (list[StatefulTrainerCallback], 可选) — 附加到 Trainer 的回调,其状态需要保存或恢复。相关回调应实现 statefrom_state 函数。

一个包含 Trainer 内部状态的类,在创建检查点时会随模型和优化器一起保存,并传递给 TrainerCallback

在本类中,一步(step)被理解为一次更新步(update step)。当使用梯度累积时,一次更新步可能需要多次前向和后向传播:如果您使用 gradient_accumulation_steps=n,则一次更新步需要经过 n 个批次。

compute_steps

< >

( args max_steps )

根据是否为比例值,计算并存储用于日志记录、评估和保存步骤的绝对值。

init_training_references

< >

( trainer max_steps num_train_epochs trial )

self 中所需的初始训练参考信息存储起来。

load_from_json

< >

( json_path: str )

json_path 的内容创建一个实例。

save_to_json

< >

( json_path: str )

将此实例的内容以 JSON 格式保存到 json_path 中。

TrainerControl

class transformers.TrainerControl

< >

( should_training_stop: bool = False should_epoch_stop: bool = False should_save: bool = False should_evaluate: bool = False should_log: bool = False )

参数

  • should_training_stop (bool, 可选,默认为 False) — 是否应该中断训练。

    如果为 True,此变量将不会被改回 False。训练将直接停止。

  • should_epoch_stop (bool, 可选,默认为 False) — 是否应该中断当前 epoch。

    如果为 True,此变量将在下一个 epoch 开始时被改回 False

  • should_save (bool, 可选,默认为 False) — 是否应该在当前步保存模型。

    如果为 True,此变量将在下一步开始时被改回 False

  • should_evaluate (bool, 可选,默认为 False) — 是否应该在当前步对模型进行评估。

    如果为 True,此变量将在下一步开始时被改回 False

  • should_log (bool, 可选,默认为 False) — 是否应该在当前步报告日志。

    如果为 True,此变量将在下一步开始时被改回 False

一个处理 Trainer 控制流的类。该类被 TrainerCallback 用于激活训练循环中的某些开关。

在 GitHub 上更新

© . This site is unofficial and not affiliated with Hugging Face, Inc.