一、开篇:为什么要关注这五个隐形坑

很多做AI模型训练的人,都会同时用到PyTorch Lightning(以下简称PL)和Weights & Biases(以下简称W&B)这两个工具。PL能帮你把重复的训练代码封装好,不用每次都写初始化、循环、保存模型这些重复活;W&B能帮你实时看训练的指标、对比不同模型的效果,不用自己写日志系统。但把这俩凑一起用的时候,很多人都会遇到“看着代码没毛病,跑起来结果不对”的情况,这就是遇到了隐形坑。这些坑不是代码报错,而是结果不符合预期,很难排查。接下来就从五个最常见的隐形坑说起,每个坑都给你说清楚怎么踩的、怎么避。

二、坑一:回调顺序搞反,日志全乱套

2.1 什么是回调顺序

PL的回调(Callback)是用来在训练的各个节点(比如每个epoch结束、每个batch结束)执行特定操作的,比如保存模型、打日志。W&B的日志其实也是通过一个回调来实现的,叫WandbLogger。很多人不知道,回调的执行顺序是按你在trainer里传的顺序来的,先传的先执行,后传的后执行。

2.2 踩坑示例

举个例子,你想在每个epoch结束后,先算模型的测试指标,再把指标传到W&B里。结果你传回调的时候,把WandbLogger放在了算测试指标的回调前面,那W&B就会记录上一个epoch的指标,而不是当前epoch的。

这里给你一个完整的踩坑代码,技术栈是Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15:

# 技术栈:Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.callbacks import Callback
from pytorch_lightning.loggers import WandbLogger
import wandb

# 自定义回调:每个epoch结束后算测试指标
class TestMetricCallback(Callback):
    def __init__(self, test_loader):
        self.test_loader = test_loader
        self.test_accuracy = 0.0

    def on_epoch_end(self, trainer, pl_module):
        # 切换模型到评估模式
        pl_module.eval()
        correct = 0
        total = 0
        with torch.no_grad():
            for x, y in self.test_loader:
                x = x.to(pl_module.device)
                y = y.to(pl_module.device)
                pred = pl_module(x)
                correct += (pred.argmax(1) == y).sum().item()
                total += y.size(0)
        # 计算准确率
        self.test_accuracy = correct / total
        # 把指标存到trainer里,方便后续用
        trainer.callback_metrics["test_acc"] = self.test_accuracy

# 定义简单的模型
class MyModel(LightningModule):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(10, 2)
        self.loss_fn = nn.CrossEntropyLoss()

    def forward(self, x):
        return self.layer(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self(x)
        loss = self.loss_fn(pred, y)
        self.log("train_loss", loss)
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 生成模拟数据
x_train = torch.randn(1000, 10)
y_train = torch.randint(0, 2, (1000,))
x_test = torch.randn(200, 10)
y_test = torch.randint(0, 2, (200,))
train_loader = DataLoader(TensorDataset(x_train, y_train), batch_size=32)
test_loader = DataLoader(TensorDataset(x_test, y_test), batch_size=32)

# 初始化W&B logger
wandb_logger = WandbLogger(project="test-project")

# 错误写法:回调顺序反了,先传WandbLogger,再传TestMetricCallback
trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger,
    callbacks=[WandbLogger, TestMetricCallback(test_loader)]
)

model = MyModel()
trainer.fit(model, train_loader)
wandb.finish()

这个代码跑起来后,W&B里的test_acc会比实际的晚一个epoch,因为WandbLogger先执行,记录的是上一次的指标,TestMetricCallback后执行,算的是当前epoch的指标,但已经晚了。

2.3 避坑方法

只要把回调的顺序反过来,先传TestMetricCallback,再传WandbLogger就行。改一下trainer的定义:

# 正确写法:先传自定义的TestMetricCallback,再传WandbLogger
trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger,
    callbacks=[TestMetricCallback(test_loader), WandbLogger]
)

这样TestMetricCallback先执行,算出当前epoch的指标,存到trainer里,然后WandbLogger执行,把当前的指标传到W&B里,就对了。

三、坑二:自动日志和手动日志混了,重复记录

3.1 什么是自动日志和手动日志

PL的LightningModule里有个self.log()方法,用来记录指标。WandbLogger会自动把self.log()里的指标传到W&B里,这叫自动日志。很多人还会手动写wandb.log()来传指标,这叫手动日志。很多人不知道,这俩是分开的,要是你同一个指标既用self.log()传,又用wandb.log()传,就会重复记录,导致W&B里的曲线有两个一样的指标,甚至数值不一样。

3.2 踩坑示例

比如你想记录训练损失,既在training_step里用self.log()传,又手动用wandb.log()传,就会重复。看下面的代码:

# 技术栈:Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.loggers import WandbLogger
import wandb

class MyModel(LightningModule):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(10, 2)
        self.loss_fn = nn.CrossEntropyLoss()

    def forward(self, x):
        return self.layer(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self(x)
        loss = self.loss_fn(pred, y)
        # 自动日志:用self.log传损失
        self.log("train_loss", loss)
        # 手动日志:用wandb.log传损失,重复了
        wandb.log({"train_loss_manual": loss})
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 生成模拟数据
x_train = torch.randn(1000, 10)
y_train = torch.randint(0, 2, (1000,))
train_loader = DataLoader(TensorDataset(x_train, y_train), batch_size=32)

# 初始化W&B logger
wandb_logger = WandbLogger(project="test-project")

trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger
)

model = MyModel()
trainer.fit(model, train_loader)
wandb.finish()

这个代码跑起来后,W&B里会有两个损失指标:train_loss和train_loss_manual,数值一样,但曲线重复,要是你没注意,就会以为有两个不同的指标,甚至手动改了其中一个,导致结果混乱。

3.3 避坑方法

要么只用自动日志,要么只用手动日志,不要混着用。要是你需要手动传一些特殊的指标(比如模型的权重分布),那手动日志的指标名字要和自动日志的不一样,避免重复。比如上面的例子,要是你确实需要手动传,就把名字改成train_loss_manual,和自动的train_loss区分开,不要用同一个名字。

四、坑三:日志级别没设,不该记录的也传了

4.1 什么是日志级别

PL的self.log()方法有个参数叫prog_bar,还有个参数叫on_step和on_epoch,用来控制什么时候记录指标。很多人不知道,WandbLogger会默认把所有的指标都传到W&B里,不管你是在step里记录的还是在epoch里记录的,也不管你是用来显示在进度条里的还是用来训练的。要是你在每个step里都记录指标,就会导致W&B里的曲线非常密,甚至因为数据太多,加载慢,还会占用W&B的存储空间。

4.2 踩坑示例

比如你想在每个step里记录训练损失,用来显示在进度条里,就会把on_step设为True,on_epoch设为False,然后用self.log()传。结果WandbLogger把每个step的损失都传到了W&B里,导致曲线有1000个点(假设一个epoch有1000个step),非常乱。看下面的代码:

# 技术栈:Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.loggers import WandbLogger
import wandb

class MyModel(LightningModule):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(10, 2)
        self.loss_fn = nn.CrossEntropyLoss()

    def forward(self, x):
        return self.layer(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self(x)
        loss = self.loss_fn(pred, y)
        # 错误写法:on_step设为True,on_epoch设为False,每个step都记录
        self.log("train_loss", loss, on_step=True, on_epoch=False, prog_bar=True)
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 生成模拟数据
x_train = torch.randn(1000, 10)
y_train = torch.randint(0, 2, (1000,))
train_loader = DataLoader(TensorDataset(x_train, y_train), batch_size=32)

# 初始化W&B logger
wandb_logger = WandbLogger(project="test-project")

trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger
)

model = MyModel()
trainer.fit(model, train_loader)
wandb.finish()

这个代码跑起来后,W&B里的train_loss曲线会有300个点(3个epoch,每个epoch100个step),非常密,要是你想看趋势,根本看不清,还会占用W&B的存储空间。

4.3 避坑方法

在self.log()里,把on_step设为False,on_epoch设为True,这样就只会在每个epoch结束后记录一次指标,曲线就会很清晰。要是你确实需要在step里记录指标,那就在WandbLogger里加个参数,过滤掉step级别的日志。比如:

# 正确写法:过滤掉step级别的日志,只保留epoch级别的
wandb_logger = WandbLogger(
    project="test-project",
    log_step_metrics=False # 关闭step级别的日志
)

这样WandbLogger就只会记录epoch级别的指标,不会记录step级别的了。

五、坑四:分布式训练下,日志重复上传

5.1 什么是分布式训练

分布式训练是指用多个GPU或者多台机器一起训练模型,这样能加快训练速度。PL的Trainer里有个参数叫devices,设成大于1的数,就是用多个GPU训练。很多人不知道,在分布式训练下,每个GPU都会跑一遍训练代码,要是你没处理好,每个GPU都会把日志传到W&B里,导致同一个指标重复上传多次,曲线就会有很多重复的点。

5.2 踩坑示例

比如你用2个GPU训练,每个GPU都会跑training_step,然后用self.log()传指标,WandbLogger就会把两个GPU的指标都传到W&B里,导致同一个step有两个一样的指标。看下面的代码:

# 技术栈:Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.loggers import WandbLogger
import wandb

class MyModel(LightningModule):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(10, 2)
        self.loss_fn = nn.CrossEntropyLoss()

    def forward(self, x):
        return self.layer(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self(x)
        loss = self.loss_fn(pred, y)
        self.log("train_loss", loss)
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 生成模拟数据
x_train = torch.randn(1000, 10)
y_train = torch.randint(0, 2, (1000,))
train_loader = DataLoader(TensorDataset(x_train, y_train), batch_size=32)

# 初始化W&B logger
wandb_logger = WandbLogger(project="test-project")

# 错误写法:用2个GPU训练,没处理分布式日志
trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger,
    devices=2, # 用2个GPU
    accelerator="gpu"
)

model = MyModel()
trainer.fit(model, train_loader)
wandb.finish()

这个代码跑起来后,W&B里的train_loss曲线会有每个step两个点,因为两个GPU都传了日志,导致曲线重复,甚至数值不一样,因为每个GPU的batch不一样。

5.3 避坑方法

在分布式训练下,WandbLogger会自动处理,但是你要确保只在主进程里传日志。PL的LightningModule里有个属性叫self.trainer.is_global_zero,用来判断是不是主进程。要是你用手动日志的话,就加个判断,只在主进程里传。比如:

# 正确写法:手动日志只在主进程里传
def training_step(self, batch, batch_idx):
    x, y = batch
    pred = self(x)
    loss = self.loss_fn(pred, y)
    self.log("train_loss", loss)
    # 只在主进程里传手动日志
    if self.trainer.is_global_zero:
        wandb.log({"train_loss_manual": loss})
    return loss

要是你用自动日志的话,WandbLogger会自动把所有进程的指标合并,只传一次,不用你处理,但是你要确保WandbLogger是在主进程里初始化的。

六、坑五:自定义指标没同步,W&B里看不到

6.1 什么是自定义指标

自定义指标是指你自己定义的、不是PL自动生成的指标,比如模型的准确率、召回率、F1值等。很多人不知道,在分布式训练下,自定义指标需要同步到主进程,不然W&B里就看不到,或者数值不对。

6.2 踩坑示例

比如你在每个epoch结束后,自己算模型的准确率,然后用self.log()传,结果在分布式训练下,每个GPU算的准确率不一样,W&B里的准确率是错的,甚至看不到。看下面的代码:

# 技术栈:Python 3.10 + PyTorch Lightning 2.0 + Weights & Biases 0.15
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.callbacks import Callback
from pytorch_lightning.loggers import WandbLogger
import wandb

# 自定义回调:每个epoch结束后算测试指标
class TestMetricCallback(Callback):
    def __init__(self, test_loader):
        self.test_loader = test_loader
        self.test_accuracy = 0.0

    def on_epoch_end(self, trainer, pl_module):
        # 切换模型到评估模式
        pl_module.eval()
        correct = 0
        total = 0
        with torch.no_grad():
            for x, y in self.test_loader:
                x = x.to(pl_module.device)
                y = y.to(pl_module.device)
                pred = pl_module(x)
                correct += (pred.argmax(1) == y).sum().item()
                total += y.size(0)
        # 计算准确率,错误写法:没同步到主进程
        self.test_accuracy = correct / total
        trainer.callback_metrics["test_acc"] = self.test_accuracy

# 定义简单的模型
class MyModel(LightningModule):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(10, 2)
        self.loss_fn = nn.CrossEntropyLoss()

    def forward(self, x):
        return self.layer(x)

    def training_step(self, batch, batch_idx):
        x, y = batch
        pred = self(x)
        loss = self.loss_fn(pred, y)
        self.log("train_loss", loss)
        return loss

    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=0.001)

# 生成模拟数据
x_train = torch.randn(1000, 10)
y_train = torch.randint(0, 2, (1000,))
x_test = torch.randn(200, 10)
y_test = torch.randint(0, 2, (200,))
train_loader = DataLoader(TensorDataset(x_train, y_train), batch_size=32)
test_loader = DataLoader(TensorDataset(x_test, y_test), batch_size=32)

# 初始化W&B logger
wandb_logger = WandbLogger(project="test-project")

# 错误写法:用2个GPU训练,自定义指标没同步
trainer = Trainer(
    max_epochs=3,
    logger=wandb_logger,
    callbacks=[TestMetricCallback(test_loader)],
    devices=2,
    accelerator="gpu"
)

model = MyModel()
trainer.fit(model, train_loader)
wandb.finish()

这个代码跑起来后,W&B里的test_acc会是错的,因为每个GPU算的准确率不一样,没同步到主进程,导致主进程的准确率是错的。

6.3 避坑方法

在自定义回调里,算完指标后,要把指标同步到主进程。PL的Callback里有个方法叫all_gather,用来把所有进程的指标同步到主进程。比如:

# 正确写法:自定义回调里同步指标到主进程
class TestMetricCallback(Callback):
    def __init__(self, test_loader):
        self.test_loader = test_loader
        self.test_accuracy = 0.0

    def on_epoch_end(self, trainer, pl_module):
        # 切换模型到评估模式
        pl_module.eval()
        correct = torch.tensor(0, device=pl_module.device)
        total = torch.tensor(0, device=pl_module.device)
        with torch.no_grad():
            for x, y in self.test_loader:
                x = x.to(pl_module.device)
                y = y.to(pl_module.device)
                pred = pl_module(x)
                correct += (pred.argmax(1) == y).sum()
                total += y.size(0)
        # 同步所有进程的correct和total到主进程
        correct = trainer.strategy.all_gather(correct).sum()
        total = trainer.strategy.all_gather(total).sum()
        # 计算准确率
        self.test_accuracy = correct.item() / total.item()
        trainer.callback_metrics["test_acc"] = self.test_accuracy

这样所有进程的correct和total都会同步到主进程,主进程算出的准确率就是对的,W&B里就能看到正确的指标了。

七、总结与应用场景

这五个隐形坑,本质上都是对PL和W&B的底层逻辑不熟悉导致的。回调顺序的坑是因为没搞清楚回调的执行机制,日志重复的坑是因为没搞清楚自动日志和手动日志的区别,日志级别的坑是因为没搞清楚日志的过滤机制,分布式日志重复的坑是因为没搞清楚分布式训练下的进程机制,自定义指标没同步的坑是因为没搞清楚分布式训练下的指标同步机制。

这些坑的应用场景主要是AI模型训练的过程中,尤其是需要实时监控训练指标、对比不同模型效果的场景。比如你在训练一个分类模型,需要实时看训练损失、测试准确率,对比不同学习率的效果,就会用到PL和W&B,要是遇到这些坑,就会导致结果不对,排查起来很麻烦。

避坑的核心方法是:搞清楚每个工具的底层逻辑,不要想当然的用,遇到问题多查官方文档,多做测试。比如回调顺序的坑,只要记住先传自定义回调,再传WandbLogger就行;日志重复的坑,只要记住要么只用自动日志,要么只用手动日志,不要混着用就行;日志级别的坑,只要记住在self.log()里把on_step设为False,on_epoch设为True,或者在WandbLogger里过滤step级别的日志就行;分布式日志重复的坑,只要记住手动日志只在主进程里传就行;自定义指标没同步的坑,只要记住在自定义回调里用all_gather同步指标就行。