PyTorch?Lightning?Callback使用指南
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 的代碼(清晰、模塊化、可復用)
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 專注單一職責 |
| 可復用 | 同一個 Callback 可用于多個項目 |
| 可組合 | 多個 Callback 自由組合 |
| 易測試 | 獨立的 Callback 易于單元測試 |
| 可擴展 | 輕松添加自定義功能 |
2. 核心概念與架構(gòu)
2.1 什么是 Callback?
定義:Callback 是一個可以在訓練循環(huán)的特定階段被調(diào)用的對象,用于執(zhí)行自定義操作。
核心特點:
- 繼承自
pytorch_lightning.callbacks.Callback基類 - 通過重寫鉤子方法(hook methods)來插入自定義邏輯
- 在
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 - 模型檢查點
作用:在訓練過程中自動保存模型,支持保存最佳模型或多個檢查點。
基礎用法
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)不再改善時,自動停止訓練,防止過擬合。
基礎用法
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. 基礎早停(驗證損失不下降)
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)度策略。
基礎用法
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 訓練。
基礎用法
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)重進行平均,提升泛化性能。
基礎用法
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)控訓練耗時,可設置最大訓練時間。
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 - 設備狀態(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)絡
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ù)預處理 |
| 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 后 |
預測流程鉤子
| 鉤子方法 | 觸發(fā)時機 |
|---|---|
| on_predict_start(trainer, pl_module) | 預測開始前 |
| on_predict_end(trainer, pl_module) | 預測結(jié)束后 |
| on_predict_epoch_start(trainer, pl_module) | 預測 epoch 開始前 |
| on_predict_epoch_end(trainer, pl_module) | 預測 epoch 結(jié)束后 |
| on_predict_batch_start(...) | 每個預測 batch 前 |
| on_predict_batch_end(...) | 每個預測 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) | 加載檢查點時 | 恢復自定義狀態(tài) |
| on_before_backward(trainer, pl_module, loss) | 反向傳播前 | 梯度預處理 |
| 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 基礎模板
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:保存驗證集預測結(jié)果
class SaveValidationPredictionsCallback(Callback):
"""保存每個 epoch 的驗證集預測結(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):
# 收集預測結(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"驗證集預測已保存: {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 基礎訓練配置
場景:標準的分類/回歸任務
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. 設備狀態(tài)監(jiān)控
DeviceStatsMonitor(),
# 5. 時間限制(如云服務器按時計費)
Timer(duration=timedelta(hours=10)),
# 6. 自定義訓練報告
TrainingReportCallback(),
]
6.3 研究實驗配置
場景:科研項目,需要詳細記錄和復現(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. 保存驗證集預測(用于后續(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. 高級應用與最佳實踐
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í)行(避免重復)
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)但未恢復
# ? 錯誤示例
def on_validation_start(self, trainer, pl_module):
pl_module.train() # 錯誤地切換到訓練模式
解決:
# ? 正確示例
def on_validation_start(self, trainer, pl_module):
# Lightning 會自動處理模式切換,無需手動干預
pass
錯誤3:在錯誤的鉤子中執(zhí)行耗時操作
# ? 錯誤示例(會嚴重拖慢訓練)
def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
# 每個 batch 都執(zhí)行復雜計算
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ù)預處理
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 的設計
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中注冊使用
搭配使用策略:
- 基礎訓練:
ModelCheckpoint+EarlyStopping+LearningRateMonitor - 研究實驗:增加梯度監(jiān)控、預測保存等
- 生產(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)境:
- 啟用
ModelCheckpoint和EarlyStopping - 配置合理的
patience和save_top_k - 添加異常處理和通知機制
附錄:快速參考
Callback 常用參數(shù)速查
| Callback | 關(guān)鍵參數(shù) | 說明 |
|---|---|---|
| ModelCheckpoint | monitor, mode, save_top_k | 保存最佳模型 |
| EarlyStopping | monitor, patience, mode | 防止過擬合 |
| LearningRateMonitor | logging_interval | 記錄學習率 |
| GradientAccumulationScheduler | scheduling | 動態(tài)調(diào)整累積步數(shù) |
| StochasticWeightAveraging | swa_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)控 |
常用代碼片段
基礎配置:
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)文章
Python光學仿真數(shù)值分析求解波動方程繪制波包變化圖
這篇文章主要為大家介紹了Python光學仿真通過數(shù)值分析求解波動方程并繪制波包變化圖的示例詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助2021-10-10

