You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何分别获取CNN模型每个训练epoch的权重?

获取CNN各训练Epoch的权重方案

核心问题说明

你直接调用model.Cnn1.weight得到的是该层权重的张量引用,训练过程中这个张量会持续更新,所以如果直接存储这个引用,最终所有epoch对应的权重都会指向训练结束后的最终值,无法保留各epoch的权重状态。


PyTorch 解决方案

通过自定义训练循环,在每个epoch结束后克隆并脱离计算图保存权重:

import torch
import torch.nn as nn
import torch.optim as optim

# 示例CNN模型
class MyCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.Cnn1 = nn.Conv2d(3, 16, kernel_size=3)
        # 其他层定义...

model = MyCNN()
optimizer = optim.SGD(model.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss()

# 初始化存储结构,按epoch存储权重
epoch_weights = {}

# 100个epoch的训练循环
for epoch in range(1, 101):
    model.train()
    # 训练步骤:前向传播、损失计算、反向传播、参数更新
    # 替换为你的实际训练代码:
    # outputs = model(input_batch)
    # loss = criterion(outputs, label_batch)
    # optimizer.zero_grad()
    # loss.backward()
    # optimizer.step()

    # 保存当前epoch的Cnn1层权重
    # clone()创建独立副本,detach()脱离计算图避免后续梯度影响
    cnn1_weight_copy = model.Cnn1.weight.clone().detach()
    epoch_weights[f'epoch {epoch}'] = cnn1_weight_copy

# 查看指定epoch的权重,示例:epoch 5
print(epoch_weights['epoch 5'])

TensorFlow/Keras 解决方案

使用自定义回调函数,在每个epoch结束时触发权重保存:

import tensorflow as tf
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D

# 示例CNN模型
model = Sequential([
    Conv2D(16, 3, name='Cnn1', input_shape=(28, 28, 3)),
    # 其他层定义...
])

model.compile(optimizer='sgd', loss='sparse_categorical_crossentropy')

# 初始化存储结构
epoch_weights = {}

# 自定义回调类
class SaveEpochWeights(tf.keras.callbacks.Callback):
    def on_epoch_end(self, epoch, logs=None):
        # 获取Cnn1层的权重(返回值是[权重张量, 偏置张量],取索引0为权重)
        cnn1_weights = self.model.get_layer('Cnn1').get_weights()[0]
        epoch_weights[f'epoch {epoch+1}'] = cnn1_weights

# 启动训练,传入回调函数
model.fit(x_train, y_train, epochs=100, callbacks=[SaveEpochWeights()])

# 查看指定epoch的权重,示例:epoch 5
print(epoch_weights['epoch 5'])

关键注意点

  • 必须保存权重的独立副本:PyTorch用clone().detach(),Keras中get_weights()返回的本身就是副本,无需额外处理。
  • 存储结构可按需选择:字典(按epoch键值对存储)或列表(按索引对应epoch)都可以。

内容的提问来源于stack exchange,提问作者Muhammadyusuf

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 16:10:33