最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

PyTorch?Lightning?Callback使用指南

 更新時間:2025年12月24日 15:00:31   作者:其美杰布-富貴-李  
文章主要介紹了Callback在深度學習訓練過程中的重要性,包括其核心價值、核心概念與架構(gòu)、內(nèi)置Callback的詳細解釋以及自定義Callback的開發(fā)方法

1. 背景與動機

1.1 為什么需要 Callback?

在深度學習訓練過程中,我們經(jīng)常需要在特定時刻執(zhí)行特定操作:

訓練過程中的常見需求

  • ? 每個 epoch 結(jié)束后保存最佳模型
  • ? 當驗證損失不再下降時提前停止訓練
  • ? 記錄學習率變化曲線
  • ? 在訓練開始前初始化某些參數(shù)
  • ? 定期驗證模型在特定數(shù)據(jù)集上的表現(xiàn)
  • ? 動態(tài)調(diào)整訓練策略(如梯度累積)

傳統(tǒng)做法的問題

# ? 不使用 Callback 的代碼(耦合度高、難以維護)
for epoch in range(max_epochs):
    # 訓練邏輯
    train_loss = train_epoch(model, train_loader)
    
    # 驗證邏輯
    val_loss = validate(model, val_loader)
    
    # 手動保存最佳模型
    if val_loss < best_loss:
        best_loss = val_loss
        torch.save(model.state_dict(), 'best_model.pth')
    
    # 手動早停邏輯
    if val_loss > best_loss:
        patience_counter += 1
        if patience_counter >= patience:
            print("Early stopping!")
            break
    
    # 手動記錄日志
    log_metrics(epoch, train_loss, val_loss)
    
    # ... 更多邏輯混雜在一起

使用 Callback 的優(yōu)勢

# ? 使用 Callback 的代碼(清晰、模塊化、可復(fù)用)
trainer = pl.Trainer(
    max_epochs=100,
    callbacks=[
        ModelCheckpoint(monitor='val_loss', mode='min'),
        EarlyStopping(monitor='val_loss', patience=10),
        LearningRateMonitor(logging_interval='epoch'),
    ]
)
trainer.fit(model, train_loader, val_loader)

1.2 Callback 的核心價值

優(yōu)勢說明
解耦合訓練邏輯與輔助功能分離
模塊化每個 Callback 專注單一職責
可復(fù)用同一個 Callback 可用于多個項目
可組合多個 Callback 自由組合
易測試獨立的 Callback 易于單元測試
可擴展輕松添加自定義功能

2. 核心概念與架構(gòu)

2.1 什么是 Callback?

定義:Callback 是一個可以在訓練循環(huán)的特定階段被調(diào)用的對象,用于執(zhí)行自定義操作。

核心特點

  1. 繼承自 pytorch_lightning.callbacks.Callback 基類
  2. 通過重寫鉤子方法(hook methods)來插入自定義邏輯
  3. Trainer 的特定時刻自動被調(diào)用

2.2 Callback 的工作原理

訓練流程                    Callback 鉤子觸發(fā)時機
│
├─ Trainer.fit()
│  │
│  ├─ on_fit_start()        ← 訓練開始前
│  │
│  ├─ Epoch Loop
│  │  │
│  │  ├─ on_train_epoch_start()  ← 每個訓練 epoch 開始
│  │  │
│  │  ├─ Training Batch Loop
│  │  │  ├─ on_train_batch_start()   ← 每個訓練 batch 前
│  │  │  ├─ training_step()
│  │  │  └─ on_train_batch_end()     ← 每個訓練 batch 后
│  │  │
│  │  ├─ on_train_epoch_end()    ← 每個訓練 epoch 結(jié)束
│  │  │
│  │  ├─ Validation Loop
│  │  │  ├─ on_validation_epoch_start()
│  │  │  ├─ validation_step()
│  │  │  └─ on_validation_epoch_end()
│  │  │
│  │  └─ on_epoch_end()          ← 每個完整 epoch 結(jié)束
│  │
│  └─ on_fit_end()          ← 訓練完全結(jié)束
│
└─ Trainer.test()
   ├─ on_test_start()
   ├─ test_step()
   └─ on_test_end()

2.3 Callback 的分類

PyTorch Lightning 的 Callback 可以分為以下幾類:

類別典型 Callback用途
模型管理ModelCheckpoint保存/加載模型
訓練控制EarlyStopping, GradientAccumulationScheduler控制訓練流程
監(jiān)控與日志LearningRateMonitor, DeviceStatsMonitor記錄訓練指標
用戶界面RichProgressBar, TQDMProgressBar顯示訓練進度
優(yōu)化策略StochasticWeightAveraging高級優(yōu)化技巧
調(diào)試工具ModelSummary, Timer輔助調(diào)試
自定義用戶自定義 Callback特定需求

3. 內(nèi)置 Callback 詳解

3.1 ModelCheckpoint - 模型檢查點

作用:在訓練過程中自動保存模型,支持保存最佳模型或多個檢查點。

基礎(chǔ)用法

from pytorch_lightning.callbacks import ModelCheckpoint

# 示例1:保存驗證損失最低的模型
checkpoint = ModelCheckpoint(
    monitor='val_loss',           # 監(jiān)控的指標
    dirpath='checkpoints/',       # 保存目錄
    filename='best-{epoch:02d}-{val_loss:.4f}',  # 文件名模板
    save_top_k=1,                 # 保存最好的 1 個模型
    mode='min',                   # 'min' 表示越小越好,'max' 表示越大越好
    save_last=True,               # 額外保存最后一個 epoch 的模型
    verbose=True,                 # 打印日志
)

trainer = pl.Trainer(callbacks=[checkpoint])

完整參數(shù)說明

ModelCheckpoint(
    # 核心參數(shù)
    monitor='val_loss',              # 監(jiān)控的指標名稱(必須在 self.log() 中記錄)
    mode='min',                      # 'min'/'max'/'auto'
    
    # 保存策略
    save_top_k=3,                    # 保存最好的 k 個模型(-1 表示全部保存)
    save_last=True,                  # 是否額外保存最后一個模型(last.ckpt)
    save_weights_only=False,         # True: 僅保存權(quán)重,F(xiàn)alse: 保存完整狀態(tài)
    
    # 文件命名
    dirpath='checkpoints/',          # 保存目錄
    filename='epoch={epoch:02d}-val_loss={val_loss:.4f}',  # 文件名模板
    auto_insert_metric_name=True,    # 自動在文件名中插入 monitor 名稱
    
    # 觸發(fā)條件
    every_n_epochs=1,                # 每 n 個 epoch 檢查一次
    every_n_train_steps=None,        # 每 n 個訓練步檢查一次
    train_time_interval=None,        # 按時間間隔檢查(如 timedelta(minutes=30))
    
    # 其他
    verbose=True,                    # 是否打印保存信息
    save_on_train_epoch_end=None,    # 在訓練 epoch 結(jié)束時保存(默認驗證后)
)

高級用法

1. 同時保存多個指標的最佳模型

# 保存 val_loss 最低的模型
checkpoint_loss = ModelCheckpoint(
    monitor='val_loss',
    dirpath='checkpoints/loss/',
    filename='best-loss-{epoch:02d}-{val_loss:.4f}',
    mode='min',
    save_top_k=1,
)

# 保存 val_acc 最高的模型
checkpoint_acc = ModelCheckpoint(
    monitor='val_acc',
    dirpath='checkpoints/acc/',
    filename='best-acc-{epoch:02d}-{val_acc:.4f}',
    mode='max',
    save_top_k=1,
)

trainer = pl.Trainer(callbacks=[checkpoint_loss, checkpoint_acc])

2. 定期保存檢查點(無論性能如何)

# 每 5 個 epoch 保存一次
checkpoint_periodic = ModelCheckpoint(
    dirpath='checkpoints/periodic/',
    filename='epoch={epoch:02d}',
    every_n_epochs=5,
    save_top_k=-1,  # 保存所有
)

3. 按訓練步數(shù)保存

checkpoint_steps = ModelCheckpoint(
    dirpath='checkpoints/steps/',
    filename='step={step}',
    every_n_train_steps=1000,  # 每 1000 步保存
    save_top_k=-1,
)

4. 按時間間隔保存

from datetime import timedelta

checkpoint_time = ModelCheckpoint(
    dirpath='checkpoints/timed/',
    train_time_interval=timedelta(minutes=30),  # 每 30 分鐘保存
    save_top_k=-1,
)

訪問最佳模型路徑

trainer.fit(model, train_loader, val_loader)

# 獲取最佳模型路徑
best_model_path = checkpoint.best_model_path
print(f"Best model: {best_model_path}")

# 獲取最佳分數(shù)
best_score = checkpoint.best_model_score
print(f"Best score: {best_score}")

# 加載最佳模型
best_model = MyModel.load_from_checkpoint(best_model_path)

3.2 EarlyStopping - 早停

作用:當監(jiān)控指標在一定時間內(nèi)不再改善時,自動停止訓練,防止過擬合。

基礎(chǔ)用法

from pytorch_lightning.callbacks import EarlyStopping

early_stop = EarlyStopping(
    monitor='val_loss',     # 監(jiān)控的指標
    patience=10,            # 容忍多少個 epoch 不改善
    mode='min',             # 'min' 或 'max'
    verbose=True,           # 打印停止信息
    min_delta=0.001,        # 最小改善量(小于此值不算改善)
)

trainer = pl.Trainer(callbacks=[early_stop])

完整參數(shù)說明

EarlyStopping(
    # 核心參數(shù)
    monitor='val_loss',              # 監(jiān)控的指標
    mode='min',                      # 'min'/'max'/'auto'
    patience=3,                      # 容忍的 epoch 數(shù)
    
    # 判斷標準
    min_delta=0.0,                   # 最小改善閾值(絕對值)
    strict=True,                     # 是否嚴格要求改善(False 允許相等)
    
    # 停止行為
    stopping_threshold=None,         # 達到此值立即停止(如 val_loss < 0.01)
    divergence_threshold=None,       # 超過此值立即停止(如 val_loss > 10.0)
    check_finite=True,               # 檢查指標是否為有限值
    check_on_train_epoch_end=None,   # 在訓練 epoch 結(jié)束時檢查(默認驗證后)
    
    # 日志
    verbose=True,
    log_rank_zero_only=False,        # 僅在主進程打印
)

實用場景

1. 基礎(chǔ)早停(驗證損失不下降)

early_stop = EarlyStopping(
    monitor='val_loss',
    patience=15,
    mode='min',
    verbose=True,
)

2. 準確率不提升時停止

early_stop = EarlyStopping(
    monitor='val_acc',
    patience=10,
    mode='max',
    min_delta=0.005,  # 提升小于 0.5% 不算改善
)

3. 達到目標后立即停止

early_stop = EarlyStopping(
    monitor='val_acc',
    stopping_threshold=0.95,  # 準確率達到 95% 立即停止
    mode='max',
)

4. 檢測發(fā)散(loss 爆炸)

early_stop = EarlyStopping(
    monitor='train_loss',
    divergence_threshold=10.0,  # 訓練損失超過 10 立即停止
    mode='min',
)

3.3 LearningRateMonitor - 學習率監(jiān)控

作用:自動記錄學習率變化,用于可視化學習率調(diào)度策略。

基礎(chǔ)用法

from pytorch_lightning.callbacks import LearningRateMonitor

lr_monitor = LearningRateMonitor(
    logging_interval='epoch',  # 'step' 或 'epoch'
    log_momentum=False,        # 是否記錄 momentum(SGD 優(yōu)化器)
)

trainer = pl.Trainer(callbacks=[lr_monitor])

使用場景

1. 監(jiān)控學習率調(diào)度器

class MyModel(pl.LightningModule):
    def configure_optimizers(self):
        optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)
        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
            optimizer, mode='min', factor=0.5, patience=5
        )
        return {
            'optimizer': optimizer,
            'lr_scheduler': {
                'scheduler': scheduler,
                'monitor': 'val_loss',
            }
        }

# 在 TensorBoard 中自動記錄 lr 曲線
trainer = pl.Trainer(
    callbacks=[LearningRateMonitor(logging_interval='epoch')],
    logger=TensorBoardLogger('logs/')
)

2. 每步記錄(用于 OneCycleLR 等)

lr_monitor = LearningRateMonitor(logging_interval='step')

3.4 RichProgressBar / TQDMProgressBar - 進度條

作用:顯示訓練進度和實時指標。

RichProgressBar(推薦)

from pytorch_lightning.callbacks import RichProgressBar

# 默認配置
progress_bar = RichProgressBar()

# 自定義配置
progress_bar = RichProgressBar(
    refresh_rate=1,              # 刷新頻率(步數(shù))
    leave=True,                  # 訓練結(jié)束后保留進度條
    theme=RichProgressBarTheme(  # 自定義主題
        description="green_yellow",
        progress_bar="green1",
        progress_bar_finished="green1",
        batch_progress="green_yellow",
        time="grey82",
        processing_speed="grey82",
        metrics="grey82",
    ),
)

自定義進度條顯示

class CustomProgressBar(RichProgressBar):
    def get_metrics(self, trainer, model):
        # 獲取父類的指標
        items = super().get_metrics(trainer, model)
        
        # 自定義顯示格式(如顯示更多小數(shù)位)
        items = {
            k: f"{v:.6f}" if isinstance(v, (int, float)) else v
            for k, v in items.items()
        }
        return items

3.5 GradientAccumulationScheduler - 梯度累積調(diào)度

作用:動態(tài)調(diào)整梯度累積步數(shù),實現(xiàn)變 batch size 訓練。

基礎(chǔ)用法

from pytorch_lightning.callbacks import GradientAccumulationScheduler

# 在不同 epoch 使用不同的累積步數(shù)
accumulator = GradientAccumulationScheduler(
    scheduling={
        0: 8,   # epoch 0-4: 累積 8 步
        5: 4,   # epoch 5-9: 累積 4 步
        10: 2,  # epoch 10+: 累積 2 步
    }
)

trainer = pl.Trainer(callbacks=[accumulator])

實用場景

場景:GPU 顯存有限,初期用小 batch,后期逐步增大。

# 等效 batch size 變化:
# epoch 0-4:  batch_size=16 × accumulate=8 = 128
# epoch 5-9:  batch_size=16 × accumulate=4 = 64
# epoch 10+:  batch_size=16 × accumulate=2 = 32

accumulator = GradientAccumulationScheduler(
    scheduling={0: 8, 5: 4, 10: 2}
)

3.6 StochasticWeightAveraging (SWA) - 隨機權(quán)重平均

作用:對訓練后期的模型權(quán)重進行平均,提升泛化性能。

基礎(chǔ)用法

from pytorch_lightning.callbacks import StochasticWeightAveraging

swa = StochasticWeightAveraging(
    swa_lrs=1e-2,              # SWA 階段的學習率
    swa_epoch_start=0.8,       # 從 80% epoch 開始 SWA(0.8 × max_epochs)
    annealing_epochs=10,       # 退火 epoch 數(shù)
    annealing_strategy='cos',  # 'cos' 或 'linear'
)

trainer = pl.Trainer(
    max_epochs=100,
    callbacks=[swa]
)

原理與效果

正常訓練: 模型權(quán)重在最優(yōu)點附近震蕩
SWA:      對后期權(quán)重求平均,得到更平滑的模型

訓練曲線:
        ╱╲  ╱╲  ╱╲
Loss   ╱  ╲╱  ╲╱  ╲  ← 正常訓練
      ╱____________╲  ← SWA 平均后(更穩(wěn)定)
      ↑
    SWA Start

3.7 ModelSummary - 模型摘要

作用:在訓練開始前打印模型結(jié)構(gòu)和參數(shù)統(tǒng)計。

from pytorch_lightning.callbacks import ModelSummary

summary = ModelSummary(
    max_depth=2,  # 顯示的最大層級深度(-1 表示全部)
)

trainer = pl.Trainer(callbacks=[summary])

輸出示例

  | Name      | Type   | Params
------------------------------------
0 | layer1    | Linear | 320
1 | layer2    | Linear | 640
2 | layer3    | Linear | 10
------------------------------------
970       Trainable params
0         Non-trainable params
970       Total params

3.8 Timer - 訓練時間監(jiān)控

作用:監(jiān)控訓練耗時,可設(shè)置最大訓練時間。

from pytorch_lightning.callbacks import Timer
from datetime import timedelta

timer = Timer(
    duration=timedelta(hours=2),  # 最大訓練時間 2 小時
    interval='epoch',             # 檢查間隔('step' 或 'epoch')
    verbose=True,
)

trainer = pl.Trainer(callbacks=[timer])

3.9 DeviceStatsMonitor - 設(shè)備狀態(tài)監(jiān)控

作用:監(jiān)控 GPU/CPU 使用情況。

from pytorch_lightning.callbacks import DeviceStatsMonitor

device_stats = DeviceStatsMonitor()

trainer = pl.Trainer(callbacks=[device_stats])

記錄的指標

  • GPU 利用率
  • GPU 內(nèi)存使用
  • CPU 內(nèi)存使用

3.10 BaseFinetuning - 微調(diào)輔助

作用:輔助實現(xiàn)凍結(jié)-解凍訓練策略。

from pytorch_lightning.callbacks import BaseFinetuning

class FeatureExtractorFreezeUnfreeze(BaseFinetuning):
    def __init__(self, unfreeze_at_epoch=10):
        super().__init__()
        self._unfreeze_at_epoch = unfreeze_at_epoch

    def freeze_before_training(self, pl_module):
        # 初始凍結(jié)骨干網(wǎng)絡(luò)
        self.freeze(pl_module.feature_extractor)

    def finetune_function(self, pl_module, current_epoch, optimizer):
        # 在指定 epoch 解凍
        if current_epoch == self._unfreeze_at_epoch:
            self.unfreeze_and_add_param_group(
                modules=pl_module.feature_extractor,
                optimizer=optimizer,
                lr=1e-5,  # 使用更小的學習率
            )

trainer = pl.Trainer(callbacks=[FeatureExtractorFreezeUnfreeze(unfreeze_at_epoch=10)])

4. Callback 生命周期鉤子方法

4.1 完整鉤子方法列表

PyTorch Lightning 提供了豐富的鉤子方法,覆蓋訓練的各個階段:

訓練流程鉤子

鉤子方法觸發(fā)時機常用場景
on_fit_start(trainer, pl_module)fit() 開始前初始化全局狀態(tài)
on_fit_end(trainer, pl_module)fit() 結(jié)束后清理資源、保存最終結(jié)果
on_train_start(trainer, pl_module)訓練開始前打印訓練配置
on_train_end(trainer, pl_module)訓練結(jié)束后生成訓練報告
on_train_epoch_start(trainer, pl_module)每個訓練 epoch 開始前重置 epoch 級別的統(tǒng)計
on_train_epoch_end(trainer, pl_module)每個訓練 epoch 結(jié)束后計算 epoch 級別的指標
on_train_batch_start(trainer, pl_module, batch, batch_idx)每個訓練 batch 前數(shù)據(jù)預(yù)處理
on_train_batch_end(trainer, pl_module, outputs, batch, batch_idx)每個訓練 batch 后記錄 batch 級別的指標

驗證流程鉤子

鉤子方法觸發(fā)時機常用場景
on_validation_start(trainer, pl_module)驗證開始前切換到評估模式
on_validation_end(trainer, pl_module)驗證結(jié)束后計算驗證集總體指標
on_validation_epoch_start(trainer, pl_module)驗證 epoch 開始前重置驗證統(tǒng)計
on_validation_epoch_end(trainer, pl_module)驗證 epoch 結(jié)束后計算混淆矩陣等
on_validation_batch_start(trainer, pl_module, batch, batch_idx, dataloader_idx)每個驗證 batch 前-
on_validation_batch_end(trainer, pl_module, outputs, batch, batch_idx, dataloader_idx)每個驗證 batch 后-

測試流程鉤子

鉤子方法觸發(fā)時機
on_test_start(trainer, pl_module)測試開始前
on_test_end(trainer, pl_module)測試結(jié)束后
on_test_epoch_start(trainer, pl_module)測試 epoch 開始前
on_test_epoch_end(trainer, pl_module)測試 epoch 結(jié)束后
on_test_batch_start(...)每個測試 batch 前
on_test_batch_end(...)每個測試 batch 后

預(yù)測流程鉤子

鉤子方法觸發(fā)時機
on_predict_start(trainer, pl_module)預(yù)測開始前
on_predict_end(trainer, pl_module)預(yù)測結(jié)束后
on_predict_epoch_start(trainer, pl_module)預(yù)測 epoch 開始前
on_predict_epoch_end(trainer, pl_module)預(yù)測 epoch 結(jié)束后
on_predict_batch_start(...)每個預(yù)測 batch 前
on_predict_batch_end(...)每個預(yù)測 batch 后

其他重要鉤子

鉤子方法觸發(fā)時機常用場景
on_epoch_start(trainer, pl_module)每個完整 epoch 開始前(訓練+驗證)-
on_epoch_end(trainer, pl_module)每個完整 epoch 結(jié)束后保存中間結(jié)果
on_save_checkpoint(trainer, pl_module, checkpoint)保存檢查點時添加自定義數(shù)據(jù)到檢查點
on_load_checkpoint(trainer, pl_module, checkpoint)加載檢查點時恢復(fù)自定義狀態(tài)
on_before_backward(trainer, pl_module, loss)反向傳播前梯度預(yù)處理
on_after_backward(trainer, pl_module)反向傳播后梯度裁剪、檢查
on_before_optimizer_step(trainer, pl_module, optimizer)優(yōu)化器更新前-
on_before_zero_grad(trainer, pl_module, optimizer)梯度清零前-

4.2 鉤子方法參數(shù)說明

通用參數(shù)

  • trainer: pl.Trainer 實例,可訪問訓練器的狀態(tài)
  • pl_module: pl.LightningModule 實例,即你的模型
  • batch: 當前批次的數(shù)據(jù)
  • batch_idx: 批次索引
  • dataloader_idx: 數(shù)據(jù)加載器索引(多數(shù)據(jù)集時)
  • outputs: 模型輸出(如 training_step 的返回值)

訪問訓練狀態(tài)

def on_train_epoch_end(self, trainer, pl_module):
    # 訪問當前 epoch
    current_epoch = trainer.current_epoch
    
    # 訪問全局步數(shù)
    global_step = trainer.global_step
    
    # 訪問日志記錄的指標
    logged_metrics = trainer.logged_metrics
    
    # 訪問回調(diào)指標(用于 ModelCheckpoint 等)
    callback_metrics = trainer.callback_metrics
    
    # 訪問模型參數(shù)
    for name, param in pl_module.named_parameters():
        print(f"{name}: {param.shape}")

4.3 鉤子方法調(diào)用順序示例

# 完整訓練流程的鉤子調(diào)用順序

trainer.fit(model, train_loader, val_loader)
│
├─ on_fit_start()
│  ├─ on_train_start()
│  │
│  ├─ Epoch 0
│  │  ├─ on_epoch_start()
│  │  ├─ on_train_epoch_start()
│  │  │
│  │  ├─ Training Batches
│  │  │  ├─ on_train_batch_start(batch_idx=0)
│  │  │  ├─ on_before_backward()
│  │  │  ├─ on_after_backward()
│  │  │  ├─ on_before_optimizer_step()
│  │  │  ├─ on_before_zero_grad()
│  │  │  ├─ on_train_batch_end(batch_idx=0)
│  │  │  │
│  │  │  ├─ on_train_batch_start(batch_idx=1)
│  │  │  └─ ...
│  │  │
│  │  ├─ on_train_epoch_end()
│  │  │
│  │  ├─ Validation (如果啟用)
│  │  │  ├─ on_validation_epoch_start()
│  │  │  ├─ on_validation_batch_start(batch_idx=0)
│  │  │  ├─ on_validation_batch_end(batch_idx=0)
│  │  │  └─ on_validation_epoch_end()
│  │  │
│  │  └─ on_epoch_end()
│  │
│  ├─ Epoch 1
│  │  └─ ... (同上)
│  │
│  └─ on_train_end()
│
└─ on_fit_end()

5. 自定義 Callback 開發(fā)

5.1 基礎(chǔ)模板

import pytorch_lightning as pl
from pytorch_lightning.callbacks import Callback

class MyCustomCallback(Callback):
    """自定義 Callback 模板"""
    
    def __init__(self, custom_param):
        super().__init__()
        self.custom_param = custom_param
        # 初始化自定義狀態(tài)
        self.state = {}
    
    def on_train_start(self, trainer, pl_module):
        """訓練開始時調(diào)用"""
        print(f"訓練開始,參數(shù): {self.custom_param}")
    
    def on_train_epoch_end(self, trainer, pl_module):
        """每個訓練 epoch 結(jié)束時調(diào)用"""
        # 訪問訓練指標
        metrics = trainer.callback_metrics
        print(f"Epoch {trainer.current_epoch} 結(jié)束")
    
    def on_validation_epoch_end(self, trainer, pl_module):
        """每個驗證 epoch 結(jié)束時調(diào)用"""
        pass

5.2 實用自定義 Callback 示例

示例1:打印訓練進度報告

class TrainingReportCallback(Callback):
    """每個 epoch 結(jié)束后打印詳細報告"""
    
    def on_train_epoch_end(self, trainer, pl_module):
        metrics = trainer.callback_metrics
        
        print("\n" + "="*60)
        print(f"Epoch {trainer.current_epoch} 訓練報告")
        print("="*60)
        
        for key, value in metrics.items():
            if isinstance(value, torch.Tensor):
                value = value.item()
            print(f"{key:30s}: {value:.6f}")
        
        print("="*60 + "\n")

示例2:保存驗證集預(yù)測結(jié)果

class SaveValidationPredictionsCallback(Callback):
    """保存每個 epoch 的驗證集預(yù)測結(jié)果"""
    
    def __init__(self, save_dir='predictions/'):
        super().__init__()
        self.save_dir = save_dir
        self.predictions = []
        self.targets = []
    
    def on_validation_epoch_start(self, trainer, pl_module):
        # 重置存儲
        self.predictions = []
        self.targets = []
    
    def on_validation_batch_end(self, trainer, pl_module, outputs, 
                                batch, batch_idx, dataloader_idx=0):
        # 收集預(yù)測結(jié)果
        if isinstance(outputs, dict) and 'preds' in outputs:
            self.predictions.append(outputs['preds'].cpu())
            self.targets.append(outputs['targets'].cpu())
    
    def on_validation_epoch_end(self, trainer, pl_module):
        # 合并并保存
        if self.predictions:
            all_preds = torch.cat(self.predictions)
            all_targets = torch.cat(self.targets)
            
            save_path = f"{self.save_dir}/epoch_{trainer.current_epoch}.pt"
            torch.save({
                'predictions': all_preds,
                'targets': all_targets,
                'epoch': trainer.current_epoch
            }, save_path)
            
            print(f"驗證集預(yù)測已保存: {save_path}")

示例3:動態(tài)學習率調(diào)整

class CustomLRScheduler(Callback):
    """基于驗證損失的自定義學習率調(diào)整"""
    
    def __init__(self, patience=5, factor=0.5, min_lr=1e-6):
        super().__init__()
        self.patience = patience
        self.factor = factor
        self.min_lr = min_lr
        self.best_loss = float('inf')
        self.wait = 0
    
    def on_validation_epoch_end(self, trainer, pl_module):
        # 獲取當前驗證損失
        val_loss = trainer.callback_metrics.get('val_loss')
        
        if val_loss is None:
            return
        
        val_loss = val_loss.item()
        
        # 檢查是否改善
        if val_loss < self.best_loss:
            self.best_loss = val_loss
            self.wait = 0
        else:
            self.wait += 1
            
            if self.wait >= self.patience:
                # 降低學習率
                for optimizer in trainer.optimizers:
                    for param_group in optimizer.param_groups:
                        old_lr = param_group['lr']
                        new_lr = max(old_lr * self.factor, self.min_lr)
                        param_group['lr'] = new_lr
                        
                        print(f"\n學習率調(diào)整: {old_lr:.6f} → {new_lr:.6f}")
                
                self.wait = 0

示例4:梯度監(jiān)控

class GradientLoggingCallback(Callback):
    """記錄梯度統(tǒng)計信息"""
    
    def __init__(self, log_every_n_steps=100):
        super().__init__()
        self.log_every_n_steps = log_every_n_steps
    
    def on_after_backward(self, trainer, pl_module):
        if trainer.global_step % self.log_every_n_steps != 0:
            return
        
        # 計算梯度統(tǒng)計
        grad_norms = []
        for name, param in pl_module.named_parameters():
            if param.grad is not None:
                grad_norm = param.grad.norm().item()
                grad_norms.append(grad_norm)
                
                # 記錄每層梯度
                pl_module.log(f'grad_norm/{name}', grad_norm)
        
        # 記錄平均梯度范數(shù)
        if grad_norms:
            avg_grad_norm = sum(grad_norms) / len(grad_norms)
            pl_module.log('grad_norm/average', avg_grad_norm)

示例5:檢查點管理(清理舊文件)

import os
import glob

class CheckpointCleanupCallback(Callback):
    """自動清理舊的檢查點文件,僅保留最新的 N 個"""
    
    def __init__(self, checkpoint_dir='checkpoints/', keep_last_n=3):
        super().__init__()
        self.checkpoint_dir = checkpoint_dir
        self.keep_last_n = keep_last_n
    
    def on_train_epoch_end(self, trainer, pl_module):
        # 獲取所有檢查點文件
        ckpt_files = glob.glob(f"{self.checkpoint_dir}/*.ckpt")
        
        # 按修改時間排序
        ckpt_files.sort(key=os.path.getmtime, reverse=True)
        
        # 刪除舊文件
        for ckpt_file in ckpt_files[self.keep_last_n:]:
            try:
                os.remove(ckpt_file)
                print(f"刪除舊檢查點: {ckpt_file}")
            except Exception as e:
                print(f"刪除失敗: {e}")

示例6:郵件通知

import smtplib
from email.mime.text import MIMEText

class EmailNotificationCallback(Callback):
    """訓練完成或異常時發(fā)送郵件通知"""
    
    def __init__(self, recipient_email, smtp_config):
        super().__init__()
        self.recipient_email = recipient_email
        self.smtp_config = smtp_config
    
    def send_email(self, subject, message):
        """發(fā)送郵件"""
        msg = MIMEText(message)
        msg['Subject'] = subject
        msg['From'] = self.smtp_config['from']
        msg['To'] = self.recipient_email
        
        try:
            with smtplib.SMTP(self.smtp_config['server'], 
                            self.smtp_config['port']) as server:
                server.login(self.smtp_config['username'], 
                           self.smtp_config['password'])
                server.send_message(msg)
        except Exception as e:
            print(f"郵件發(fā)送失敗: {e}")
    
    def on_train_end(self, trainer, pl_module):
        """訓練結(jié)束時發(fā)送通知"""
        metrics = trainer.callback_metrics
        
        message = f"""
        訓練已完成!
        
        最終指標:
        {metrics}
        
        總 Epoch: {trainer.current_epoch}
        總步數(shù): {trainer.global_step}
        """
        
        self.send_email("訓練完成通知", message)
    
    def on_exception(self, trainer, pl_module, exception):
        """發(fā)生異常時發(fā)送通知"""
        message = f"訓練發(fā)生異常: {exception}"
        self.send_email("訓練異常通知", message)

示例7:實時可視化(Matplotlib)

import matplotlib.pyplot as plt

class RealTimePlotCallback(Callback):
    """實時繪制訓練曲線"""
    
    def __init__(self):
        super().__init__()
        self.train_losses = []
        self.val_losses = []
        self.epochs = []
        
        # 創(chuàng)建圖形
        plt.ion()  # 交互模式
        self.fig, self.ax = plt.subplots()
    
    def on_train_epoch_end(self, trainer, pl_module):
        # 記錄數(shù)據(jù)
        metrics = trainer.callback_metrics
        self.epochs.append(trainer.current_epoch)
        
        if 'train_loss' in metrics:
            self.train_losses.append(metrics['train_loss'].item())
    
    def on_validation_epoch_end(self, trainer, pl_module):
        metrics = trainer.callback_metrics
        
        if 'val_loss' in metrics:
            self.val_losses.append(metrics['val_loss'].item())
            
            # 更新圖形
            self.ax.clear()
            self.ax.plot(self.epochs, self.train_losses, label='Train Loss')
            if len(self.val_losses) > 0:
                self.ax.plot(self.epochs, self.val_losses, label='Val Loss')
            self.ax.legend()
            self.ax.set_xlabel('Epoch')
            self.ax.set_ylabel('Loss')
            self.fig.canvas.draw()
            self.fig.canvas.flush_events()
    
    def on_train_end(self, trainer, pl_module):
        # 保存最終圖形
        plt.ioff()
        self.fig.savefig('training_curve.png')
        print("訓練曲線已保存: training_curve.png")

5.3 訪問模型和數(shù)據(jù)

在自定義 Callback 中,可以訪問:

class DataInspectionCallback(Callback):
    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        # 訪問模型
        model = pl_module
        
        # 訪問批次數(shù)據(jù)
        x, y = batch  # 根據(jù)實際數(shù)據(jù)結(jié)構(gòu)解包
        
        # 訪問模型輸出
        predictions = outputs['preds']  # 根據(jù) training_step 返回值
        
        # 訪問優(yōu)化器
        optimizer = trainer.optimizers[0]
        current_lr = optimizer.param_groups[0]['lr']
        
        # 訪問日志器
        logger = trainer.logger
        logger.log_metrics({'custom_metric': 1.0}, step=trainer.global_step)

6. Callback 搭配使用策略

6.1 基礎(chǔ)訓練配置

場景:標準的分類/回歸任務(wù)

callbacks = [
    # 保存最佳模型
    ModelCheckpoint(
        monitor='val_loss',
        mode='min',
        save_top_k=1,
        filename='best-{epoch:02d}-{val_loss:.4f}',
    ),
    
    # 早停
    EarlyStopping(
        monitor='val_loss',
        patience=15,
        mode='min',
    ),
    
    # 學習率監(jiān)控
    LearningRateMonitor(logging_interval='epoch'),
    
    # 進度條
    RichProgressBar(),
]

trainer = pl.Trainer(
    max_epochs=100,
    callbacks=callbacks,
    logger=TensorBoardLogger('logs/'),
)

6.2 高性能訓練配置

場景:大模型、長時間訓練,需要多重保護

callbacks = [
    # 1. 多重模型保存策略
    ModelCheckpoint(
        monitor='val_loss',
        mode='min',
        save_top_k=3,
        filename='best-loss-{epoch:02d}-{val_loss:.4f}',
    ),
    ModelCheckpoint(
        monitor='val_acc',
        mode='max',
        save_top_k=1,
        filename='best-acc-{epoch:02d}-{val_acc:.4f}',
    ),
    ModelCheckpoint(
        every_n_epochs=10,
        filename='periodic-{epoch:02d}',
        save_top_k=-1,  # 保存所有
    ),
    
    # 2. 早停 + 發(fā)散檢測
    EarlyStopping(
        monitor='val_loss',
        patience=20,
        mode='min',
        min_delta=0.001,
    ),
    EarlyStopping(
        monitor='train_loss',
        divergence_threshold=10.0,  # 檢測 loss 爆炸
        mode='min',
    ),
    
    # 3. 學習率監(jiān)控
    LearningRateMonitor(logging_interval='step'),
    
    # 4. 設(shè)備狀態(tài)監(jiān)控
    DeviceStatsMonitor(),
    
    # 5. 時間限制(如云服務(wù)器按時計費)
    Timer(duration=timedelta(hours=10)),
    
    # 6. 自定義訓練報告
    TrainingReportCallback(),
]

6.3 研究實驗配置

場景:科研項目,需要詳細記錄和復(fù)現(xiàn)

callbacks = [
    # 1. 模型保存
    ModelCheckpoint(
        monitor='val_loss',
        mode='min',
        save_top_k=5,
        save_last=True,
    ),
    
    # 2. 早停
    EarlyStopping(monitor='val_loss', patience=30),
    
    # 3. 學習率監(jiān)控
    LearningRateMonitor(logging_interval='step'),
    
    # 4. 梯度監(jiān)控(檢測梯度消失/爆炸)
    GradientLoggingCallback(log_every_n_steps=50),
    
    # 5. 保存驗證集預(yù)測(用于后續(xù)分析)
    SaveValidationPredictionsCallback(save_dir='predictions/'),
    
    # 6. 實時可視化
    RealTimePlotCallback(),
    
    # 7. 模型摘要
    ModelSummary(max_depth=3),
]

# 同時使用 TensorBoard 和 WandB
trainer = pl.Trainer(
    callbacks=callbacks,
    logger=[
        TensorBoardLogger('logs/tensorboard/'),
        WandbLogger(project='my_research', name='exp_001'),
    ],
)

6.4 生產(chǎn)部署配置

場景:模型訓練后需要部署到生產(chǎn)環(huán)境

callbacks = [
    # 1. 僅保存權(quán)重(減小文件體積)
    ModelCheckpoint(
        monitor='val_loss',
        mode='min',
        save_top_k=1,
        save_weights_only=True,  # 僅保存權(quán)重
        filename='production-best',
    ),
    
    # 2. 早停
    EarlyStopping(
        monitor='val_loss',
        patience=10,
        stopping_threshold=0.05,  # 達到目標即停止
    ),
    
    # 3. SWA 提升泛化性能
    StochasticWeightAveraging(swa_lrs=1e-2),
    
    # 4. 檢查點清理(節(jié)省存儲)
    CheckpointCleanupCallback(keep_last_n=2),
    
    # 5. 訓練完成通知
    EmailNotificationCallback(
        recipient_email='team@company.com',
        smtp_config={...}
    ),
]

6.5 調(diào)試配置

場景:快速調(diào)試代碼,檢測 Bug

# 使用 Trainer 的快速開發(fā)標志
trainer = pl.Trainer(
    max_epochs=2,  # 少量 epoch
    limit_train_batches=10,  # 僅訓練 10 個 batch
    limit_val_batches=5,     # 僅驗證 5 個 batch
    callbacks=[
        RichProgressBar(),
        ModelSummary(max_depth=-1),  # 查看完整模型結(jié)構(gòu)
    ],
    logger=False,  # 不記錄日志
    enable_checkpointing=False,  # 不保存檢查點
)

6.6 超參數(shù)搜索配置

場景:使用 Ray Tune / Optuna 進行超參數(shù)優(yōu)化

from ray import tune
from ray.tune.integration.pytorch_lightning import TuneReportCallback

def train_func(config):
    model = MyModel(
        lr=config['lr'],
        hidden_dim=config['hidden_dim'],
    )
    
    trainer = pl.Trainer(
        max_epochs=20,
        callbacks=[
            # Ray Tune 回調(diào)(報告指標)
            TuneReportCallback(
                metrics={'val_loss': 'val_loss'},
                on='validation_end',
            ),
            EarlyStopping(monitor='val_loss', patience=5),
        ],
        enable_progress_bar=False,  # 禁用進度條(避免輸出混亂)
        enable_model_summary=False,
    )
    
    trainer.fit(model, train_loader, val_loader)

# 啟動超參數(shù)搜索
analysis = tune.run(
    train_func,
    config={
        'lr': tune.loguniform(1e-4, 1e-1),
        'hidden_dim': tune.choice([64, 128, 256]),
    },
    num_samples=20,
)

7. 高級應(yīng)用與最佳實踐

7.1 Callback 之間的通信

場景:不同 Callback 需要共享狀態(tài)

class SharedStateCallback(Callback):
    """使用 Trainer 的自定義屬性共享狀態(tài)"""
    
    def on_train_start(self, trainer, pl_module):
        # 初始化共享狀態(tài)
        trainer.my_shared_state = {'counter': 0}
    
    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        # 更新共享狀態(tài)
        trainer.my_shared_state['counter'] += 1

class AnotherCallback(Callback):
    def on_validation_epoch_end(self, trainer, pl_module):
        # 讀取共享狀態(tài)
        counter = trainer.my_shared_state['counter']
        print(f"已訓練 {counter} 個 batch")

7.2 條件執(zhí)行 Callback

class ConditionalCallback(Callback):
    """僅在特定條件下執(zhí)行"""
    
    def __init__(self, execute_after_epoch=10):
        super().__init__()
        self.execute_after_epoch = execute_after_epoch
    
    def on_validation_epoch_end(self, trainer, pl_module):
        # 僅在第 10 個 epoch 后執(zhí)行
        if trainer.current_epoch >= self.execute_after_epoch:
            print("執(zhí)行特殊操作...")

7.3 Callback 優(yōu)先級

Callback 的執(zhí)行順序由添加順序決定

callbacks = [
    CallbackA(),  # 第一個執(zhí)行
    CallbackB(),  # 第二個執(zhí)行
    CallbackC(),  # 第三個執(zhí)行
]

# 注意:ModelCheckpoint 和 EarlyStopping 的順序很重要!
callbacks = [
    ModelCheckpoint(...),  # 先保存模型
    EarlyStopping(...),    # 再判斷是否停止
]

7.4 在 Callback 中使用日志器

class CustomLoggingCallback(Callback):
    def on_train_epoch_end(self, trainer, pl_module):
        # 方式1:通過 pl_module 記錄
        pl_module.log('custom_metric', 1.0)
        
        # 方式2:直接使用 logger
        if trainer.logger:
            trainer.logger.log_metrics(
                {'another_metric': 2.0},
                step=trainer.global_step
            )
            
            # 如果是 TensorBoard
            if isinstance(trainer.logger, TensorBoardLogger):
                trainer.logger.experiment.add_scalar(
                    'special_metric', 3.0, trainer.global_step
                )

7.5 處理分布式訓練

class DistributedAwareCallback(Callback):
    """在分布式訓練中正確處理"""
    
    def on_validation_epoch_end(self, trainer, pl_module):
        # 僅在主進程執(zhí)行(避免重復(fù))
        if trainer.is_global_zero:
            print("這只在主進程打印一次")
        
        # 所有進程都執(zhí)行
        local_rank = trainer.local_rank
        print(f"進程 {local_rank} 執(zhí)行")

7.6 Callback 的測試

import unittest

class TestMyCallback(unittest.TestCase):
    def test_callback_logic(self):
        # 創(chuàng)建模擬的 trainer 和 model
        trainer = MockTrainer()
        model = MockModel()
        
        # 測試 callback
        callback = MyCustomCallback()
        callback.on_train_start(trainer, model)
        
        # 驗證行為
        self.assertEqual(callback.state['initialized'], True)

8. 常見問題與調(diào)試技巧

8.1 常見錯誤

錯誤1:在on_train_epoch_end中訪問不存在的指標

# ? 錯誤示例
def on_train_epoch_end(self, trainer, pl_module):
    val_loss = trainer.callback_metrics['val_loss']  # KeyError!

原因on_train_epoch_end 在驗證之前調(diào)用,此時 val_loss 還未計算。

解決

# ? 正確示例
def on_validation_epoch_end(self, trainer, pl_module):
    # 在驗證后訪問
    val_loss = trainer.callback_metrics.get('val_loss')
    if val_loss is not None:
        print(f"驗證損失: {val_loss}")

錯誤2:Callback 修改了模型狀態(tài)但未恢復(fù)

# ? 錯誤示例
def on_validation_start(self, trainer, pl_module):
    pl_module.train()  # 錯誤地切換到訓練模式

解決

# ? 正確示例
def on_validation_start(self, trainer, pl_module):
    # Lightning 會自動處理模式切換,無需手動干預(yù)
    pass

錯誤3:在錯誤的鉤子中執(zhí)行耗時操作

# ? 錯誤示例(會嚴重拖慢訓練)
def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
    # 每個 batch 都執(zhí)行復(fù)雜計算
    expensive_operation()

解決

# ? 正確示例
def on_train_epoch_end(self, trainer, pl_module):
    # 每個 epoch 執(zhí)行一次
    expensive_operation()

8.2 調(diào)試技巧

技巧1:打印所有可用指標

class DebugCallback(Callback):
    def on_validation_epoch_end(self, trainer, pl_module):
        print("\n可用指標:")
        for key, value in trainer.callback_metrics.items():
            print(f"  {key}: {value}")

技巧2:檢查 Callback 是否被調(diào)用

class TestCallback(Callback):
    def __init__(self):
        super().__init__()
        self.call_count = {}
    
    def _log_call(self, method_name):
        self.call_count[method_name] = self.call_count.get(method_name, 0) + 1
        print(f"[{method_name}] 被調(diào)用 {self.call_count[method_name]} 次")
    
    def on_train_start(self, trainer, pl_module):
        self._log_call('on_train_start')
    
    def on_train_epoch_end(self, trainer, pl_module):
        self._log_call('on_train_epoch_end')

技巧3:使用斷點調(diào)試

class DebugCallback(Callback):
    def on_validation_epoch_end(self, trainer, pl_module):
        # 在特定條件下觸發(fā)斷點
        if trainer.current_epoch == 5:
            import pdb; pdb.set_trace()

8.3 性能優(yōu)化

優(yōu)化1:避免頻繁的 I/O 操作

# ? 低效
class BadCallback(Callback):
    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        # 每個 batch 都寫文件
        with open('log.txt', 'a') as f:
            f.write(f"Batch {batch_idx} done\n")

# ? 高效
class GoodCallback(Callback):
    def __init__(self):
        super().__init__()
        self.buffer = []
    
    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        self.buffer.append(f"Batch {batch_idx} done\n")
    
    def on_train_epoch_end(self, trainer, pl_module):
        # 每個 epoch 寫一次
        with open('log.txt', 'a') as f:
            f.writelines(self.buffer)
        self.buffer = []

優(yōu)化2:使用條件判斷減少計算

class OptimizedCallback(Callback):
    def __init__(self, log_every_n_epochs=5):
        super().__init__()
        self.log_every_n_epochs = log_every_n_epochs
    
    def on_validation_epoch_end(self, trainer, pl_module):
        # 僅每 5 個 epoch 執(zhí)行一次
        if trainer.current_epoch % self.log_every_n_epochs == 0:
            expensive_visualization()

9. 擴展閱讀與進階方向

9.1 官方文檔

PyTorch Lightning Callbacks 文檔:https://lightning.ai/docs/pytorch/stable/extensions/callbacks.html

內(nèi)置 Callback API 參考:https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.callbacks.html

9.2 高級主題

9.2.1 與其他框架集成

  • Ray Tune 集成:分布式超參數(shù)優(yōu)化
  • Optuna 集成:貝葉斯超參數(shù)優(yōu)化
  • MLflow 集成:實驗追蹤與模型管理

9.2.2 自定義訓練循環(huán)

class CustomTrainLoop(Callback):
    """完全自定義訓練循環(huán)"""
    
    def on_train_batch_start(self, trainer, pl_module, batch, batch_idx):
        # 自定義數(shù)據(jù)預(yù)處理
        pass
    
    def on_before_backward(self, trainer, pl_module, loss):
        # 自定義損失縮放
        pass
    
    def on_before_optimizer_step(self, trainer, pl_module, optimizer):
        # 自定義梯度處理
        pass

9.2.3 高級模型管理

  • 模型版本控制:使用 DVC 或 Git LFS
  • A/B 測試:保存多個候選模型進行對比
  • 模型蒸餾:在 Callback 中實現(xiàn)教師-學生訓練

9.3 實戰(zhàn)案例學習

推薦閱讀以下開源項目的 Callback 實現(xiàn):

Transformers (Hugging Face)

查看 transformers.TrainerCallback 的設(shè)計

Lightning-Hydra-Template

完整的 PyTorch Lightning 項目模板

PyTorch Lightning Bolts

高級 Callback 示例集合

9.4 社區(qū)資源

PyTorch Lightning GitHub Discussions:https://github.com/Lightning-AI/lightning/discussions

PyTorch Lightning Slack:加入社區(qū)討論

總結(jié)

核心要點回顧

Callback 是什么

  • 在訓練循環(huán)特定階段執(zhí)行的可插拔模塊
  • 通過鉤子方法(hook)實現(xiàn)自定義邏輯

常用內(nèi)置 Callback

  • ModelCheckpoint:保存模型
  • EarlyStopping:早停
  • LearningRateMonitor:學習率監(jiān)控
  • RichProgressBar:進度條
  • StochasticWeightAveraging:SWA 優(yōu)化

生命周期鉤子

  • 訓練階段:on_train_start, on_train_epoch_end, on_train_batch_end
  • 驗證階段:on_validation_epoch_end
  • 其他:on_save_checkpoint, on_load_checkpoint

自定義 Callback

  • 繼承 Callback 基類
  • 重寫所需的鉤子方法
  • Trainer 中注冊使用

搭配使用策略

  • 基礎(chǔ)訓練:ModelCheckpoint + EarlyStopping + LearningRateMonitor
  • 研究實驗:增加梯度監(jiān)控、預(yù)測保存等
  • 生產(chǎn)部署:增加 SWA、檢查點清理等

最佳實踐建議

  • ? 模塊化:每個 Callback 專注單一職責
  • ? 可配置:通過參數(shù)控制行為
  • ? 高效:避免在高頻鉤子中執(zhí)行耗時操作
  • ? 魯棒:處理邊界情況(如指標不存在)
  • ? 可測試:編寫單元測試驗證邏輯
  • ? 文檔化:為自定義 Callback 添加詳細注釋

Callback 使用清單

訓練前檢查

  • 確認監(jiān)控的指標在 self.log() 中記錄
  • 檢查 mode 參數(shù)(‘min’ 或 ‘max’)
  • 驗證文件保存路徑存在且有寫權(quán)限

調(diào)試階段

  • 使用 verbose=True 查看詳細日志
  • 添加 DebugCallback 檢查調(diào)用順序
  • 使用小數(shù)據(jù)集快速驗證

生產(chǎn)環(huán)境

  • 啟用 ModelCheckpointEarlyStopping
  • 配置合理的 patiencesave_top_k
  • 添加異常處理和通知機制

附錄:快速參考

Callback 常用參數(shù)速查

Callback關(guān)鍵參數(shù)說明
ModelCheckpointmonitor, mode, save_top_k保存最佳模型
EarlyStoppingmonitor, patience, mode防止過擬合
LearningRateMonitorlogging_interval記錄學習率
GradientAccumulationSchedulerscheduling動態(tài)調(diào)整累積步數(shù)
StochasticWeightAveragingswa_lrs, swa_epoch_start權(quán)重平均優(yōu)化

鉤子方法速查

鉤子觸發(fā)時機常用場景
on_train_start訓練開始前初始化狀態(tài)
on_train_epoch_end訓練 epoch 結(jié)束計算 epoch 指標
on_validation_epoch_end驗證 epoch 結(jié)束保存驗證結(jié)果
on_save_checkpoint保存檢查點時添加自定義數(shù)據(jù)
on_after_backward反向傳播后梯度監(jiān)控

常用代碼片段

基礎(chǔ)配置

callbacks = [
    ModelCheckpoint(monitor='val_loss', mode='min', save_top_k=1),
    EarlyStopping(monitor='val_loss', patience=10),
    LearningRateMonitor(),
]

自定義 Callback 模板

class MyCallback(Callback):
    def on_train_epoch_end(self, trainer, pl_module):
        metrics = trainer.callback_metrics
        # 自定義邏輯

以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • 16中Python機器學習類別特征處理方法總結(jié)

    16中Python機器學習類別特征處理方法總結(jié)

    類別型特征(categorical?feature)主要是指職業(yè),血型等在有限類別內(nèi)取值的特征。在這篇文章中,小編將給大家分享一下16種類別特征處理方法,需要的可以參考一下
    2022-09-09
  • python中的for循環(huán)

    python中的for循環(huán)

    Python for循環(huán)可以遍歷任何序列的項目,如一個列表或者一個字符串。這篇文章主要介紹了python的for循環(huán),需要的朋友可以參考下
    2018-09-09
  • Python騷操作之動態(tài)定義函數(shù)

    Python騷操作之動態(tài)定義函數(shù)

    這篇文章主要介紹了Python騷操作之動態(tài)定義函數(shù),文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友們下面隨著小編來一起學習學習吧
    2019-03-03
  • pytorch教程之Tensor的值及操作使用學習

    pytorch教程之Tensor的值及操作使用學習

    這篇文章主要為大家介紹了pytorch教程中關(guān)于Tensor的操作使用,有需要的朋友可以借鑒參考下,希望可以有所幫助,祝大家升職加薪,共同進步
    2021-09-09
  • django框架防止XSS注入的方法分析

    django框架防止XSS注入的方法分析

    這篇文章主要介紹了django框架防止XSS注入的方法,結(jié)合實例形式分析了XSS攻擊的原理及Django框架防止XSS攻擊的相關(guān)操作技巧,需要的朋友可以參考下
    2019-06-06
  • Python的圖像處理庫Pillow安裝與使用教程

    Python的圖像處理庫Pillow安裝與使用教程

    Pillow庫是Python中用于圖像處理的開源庫,提供了豐富的圖像處理功能,如圖像讀取、保存、裁剪、調(diào)整大小、旋轉(zhuǎn)、添加文字等,這篇文章主要給大家介紹了關(guān)于Python的圖像處理庫Pillow安裝與使用的相關(guān)資料,需要的朋友可以參考下
    2024-04-04
  • Python光學仿真數(shù)值分析求解波動方程繪制波包變化圖

    Python光學仿真數(shù)值分析求解波動方程繪制波包變化圖

    這篇文章主要為大家介紹了Python光學仿真通過數(shù)值分析求解波動方程并繪制波包變化圖的示例詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助
    2021-10-10
  • python實現(xiàn)根據(jù)文件格式分類

    python實現(xiàn)根據(jù)文件格式分類

    這篇文章主要為大家詳細介紹了python實現(xiàn)根據(jù)文件格式分類,文中示例代碼介紹的非常詳細,具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2019-10-10
  • Python列表生成式與生成器操作示例

    Python列表生成式與生成器操作示例

    這篇文章主要介紹了Python列表生成式與生成器操作,結(jié)合實例形式分析了Python列表生成式與生成器的功能、使用方法及相關(guān)操作技巧,需要的朋友可以參考下
    2018-08-08
  • python socket實現(xiàn)聊天室

    python socket實現(xiàn)聊天室

    這篇文章主要為大家詳細介紹了python socket實現(xiàn)聊天室,文中示例代碼介紹的非常詳細,具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2021-07-07

最新評論

江津市| 新和县| 军事| 莲花县| 卫辉市| 福海县| 冕宁县| 天祝| 孝感市| 赣州市| 安国市| 左贡县| 南平市| 沧州市| 金沙县| 泉州市| 八宿县| 岫岩| 临颍县| 汉沽区| 株洲县| 连江县| 无棣县| 苗栗市| 囊谦县| 天祝| 辛集市| 黔东| 苗栗县| 申扎县| 玛曲县| 锦屏县| 青铜峡市| 昌宁县| 石屏县| 江达县| 徐闻县| 淅川县| 定州市| 海晏县| 巴彦淖尔市|