---
title: "HuggingFace Trainer × WandB：实验追踪配置与指标扩展"
description: "HuggingFace Trainer 原生支持 WandB 实验追踪，配置步骤、自定义指标 Callback、多卡训练注意事项与国内替代方案完整说明。"
pubDate: 2026-06-25
tags: ["WandB","HuggingFace Trainer","实验追踪","PyTorch","训练监控","LoRA"]
category: "AI"
lang: "zh"
math: false
---
WandB（Weights & Biases）是主流实验追踪平台，HuggingFace Trainer 对其有原生支持。训练过程的 loss 曲线、超参数、硬件指标均可自动记录，并支持跨实验对比。

---

## 一、安装与认证

```bash
pip install wandb
```

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

```bash
wandb login <your-api-key>
```

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

---

## 二、基础配置

### 2.1 开启上报

在 `TrainingArguments` 中将 `report_to` 由 `none` 改为 `wandb`：

```python
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()`：

```python
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 多卡训练：只在主进程初始化

```python
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` 是过拟合的早期信号：

```python
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 比例：

```python
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

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

---

## 五、训练结束后关闭

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

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

---

## 六、离线模式

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

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

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

---

## 七、国内替代：SwanLab

[SwanLab](https://swanlab.cn/) 是功能对标 WandB 的国产平台，服务器在国内，访问速度更快。HuggingFace Trainer 通过 Callback 集成：

```python
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

- [WandB 官方文档](https://docs.wandb.ai/)
- [HuggingFace Trainer WandB 集成](https://huggingface.co/docs/transformers/main/en/main_classes/trainer#transformers.Trainer)
- [SwanLab HuggingFace 集成文档](https://docs.swanlab.cn/)