HuggingFace Trainer × WandB:实验追踪配置与指标扩展
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_to 由 none 改为 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/loss、train/learning_rate、train/epoch |
| 验证损失 | eval/loss、eval/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 一致。