AI 2 分钟阅读

HuggingFace Trainer × WandB:实验追踪配置与指标扩展

· 更新于 2026/6/25

WandB(Weights & Biases)是主流实验追踪平台,HuggingFace Trainer 对其有原生支持。训练过程的 loss 曲线、超参数、硬件指标均可自动记录,并支持跨实验对比。


一、安装与认证

pip install wandb

获取 API Key:访问 https://wandb.ai/authorize,然后在服务器上执行:

wandb login <your-api-key>

Key 存储在 ~/.netrc,后续训练自动读取,只需执行一次。


二、基础配置

2.1 开启上报

TrainingArguments 中将 report_tonone 改为 wandb

from transformers import TrainingArguments

args = TrainingArguments(
    output_dir="output/my_model",
    report_to="wandb",          # 开启 WandB 上报
    logging_steps=10,
    eval_strategy="steps",
    eval_steps=200,
)

Trainer 内部已集成 WandB Callback,无需手动调用 wandb.log()

2.2 初始化 Run

trainer.train() 之前调用 wandb.init()

import wandb

wandb.init(
    project="my-project",
    name=f"lora_r{lora_r}_lr{lr}_ep{epochs}",   # 命名包含关键超参,便于识别
    config={
        "base_model": base_model,
        "lora_r": lora_r,
        "lora_alpha": lora_alpha,
        "learning_rate": lr,
        "epochs": epochs,
        "batch_size": batch_size,
        "train_samples": len(train_dataset),
        "val_samples": len(val_dataset),
    },
    tags=[f"r{lora_r}", f"lr{lr}"],
)

name 建议包含超参标识,config 记录所有超参,tags 用于过滤筛选。

2.3 多卡训练:只在主进程初始化

import os
from transformers import TrainingArguments

is_main = int(os.environ.get("LOCAL_RANK", 0)) == 0

if is_main:
    wandb.init(
        project="my-project",
        name=run_name,
        config=config_dict,
    )
    os.environ["WANDB_PROJECT"] = "my-project"

args = TrainingArguments(
    report_to="wandb" if is_main else "none",
    ...
)

多卡训练时只有 rank 0 上报,避免重复记录。


三、自动追踪的内容

开启后,Trainer 自动上报以下指标:

类别 指标
训练损失 train/losstrain/learning_ratetrain/epoch
验证损失 eval/losseval/runtime
梯度 train/grad_norm
超参数 wandb.init(config=...) 中传入的所有字段
硬件 GPU 利用率、显存占用、系统内存(自动采集)

四、自定义指标

Trainer 内置指标之外,可通过 TrainerCallback 扩展。

4.1 train/eval gap(过拟合监控)

eval_loss - train_loss > 0.1 是过拟合的早期信号:

from transformers import TrainerCallback
import wandb

class GapMonitorCallback(TrainerCallback):
    def on_log(self, args, state, control, logs=None, **kwargs):
        if not (logs and state.is_world_process_zero):
            return

        metrics = {}

        trains = [l for l in state.log_history if "loss" in l and "eval_loss" not in l]
        evals  = [l for l in state.log_history if "eval_loss" in l]
        if trains and evals:
            gap = evals[-1]["eval_loss"] - trains[-1]["loss"]
            metrics["train_eval_gap"] = gap

        if "grad_norm" in logs:
            metrics["grad_norm"] = logs["grad_norm"]

        if metrics:
            wandb.log(metrics, step=state.global_step)

4.2 有效 token 比例

prompt masking 后,记录实际参与训练的 token 比例:

class TokenEfficiencyCallback(TrainerCallback):
    def __init__(self):
        self._total = 0
        self._effective = 0

    def record(self, total: int, effective: int):
        self._total += total
        self._effective += effective

    def on_step_end(self, args, state, control, **kwargs):
        if state.is_world_process_zero and self._total > 0:
            ratio = self._effective / self._total
            wandb.log({"effective_token_ratio": ratio}, step=state.global_step)
            self._total = self._effective = 0

在数据集的 __getitem__ 中调用 callback.record(total_len, label_len),其中 label_len 为非 -100 的 token 数。

4.3 注册 Callback

trainer = Trainer(
    model=model,
    args=args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    callbacks=[
        GapMonitorCallback(),
        TokenEfficiencyCallback(),
    ],
)
trainer.train()

五、训练结束后关闭

trainer.train()
if is_main:
    wandb.finish()

不调用 finish() 时 WandB 会在进程退出时自动关闭,但显式调用可确保所有数据上传完毕。


六、离线模式

服务器无法访问外网时,先本地保存,后同步:

# 训练时设置离线模式
WANDB_MODE=offline python train.py

# 网络恢复后同步
wandb sync wandb/offline-run-*/

七、国内替代:SwanLab

SwanLab 是功能对标 WandB 的国产平台,服务器在国内,访问速度更快。HuggingFace Trainer 通过 Callback 集成:

from swanlab.integration.huggingface import SwanLabCallback

trainer = Trainer(
    ...
    callbacks=[SwanLabCallback(
        project="my-project",
        experiment_name=run_name,
        config=config_dict,
    )],
)

主要差异:SwanLab 不支持 report_to="swanlab",必须通过 Callback 接入;功能覆盖基本与 WandB 一致。


References