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

Keras获取每层每epoch权重并保存,触发IndexError如何解决

报错原因
  • 核心问题是你默认所有层的get_weights()返回结果都至少有2个元素(权重、偏置),但实际上很多层没有可训练参数,或者参数结构不符合这个假设:
    • 你模型中的Dropout层没有任何可训练参数,get_weights()会返回空列表,此时访问[0]直接触发索引越界
    • 你的自定义Attention层也可能存在参数数量不等于2的情况(比如没有偏置、包含多个权重参数),访问[1]也会触发越界
修复后的实现方案

以下是修改后的Callback代码,兼容不同结构的层,同时支持权重写入本地文件:

import numpy as np
import keras
from matplotlib import pyplot as plt
import seaborn as sb

class GetWeights(keras.callbacks.Callback):
  def __init__(self, save_path="epoch_weights.npz"):
    super(GetWeights, self).__init__()
    self.weight_dict = {}
    self.save_path = save_path # 权重保存路径

  def on_epoch_end(self, epoch, logs=None):
    for layer_i, layer in enumerate(self.model.layers):
      weights = layer.get_weights()
      # 跳过无参数的层
      if len(weights) == 0:
        continue
      
      print(f"=== 第{layer_i}层:{layer.name} ===")
      # 遍历层的所有参数(不一定只有权重和偏置)
      for param_idx, param in enumerate(weights):
        param_key = f"layer_{layer_i}_param_{param_idx}"
        print(f"参数{param_idx}形状:{param.shape}")

        # 可选:如果是2维权重可以画热力图
        if len(param.shape) == 2:
          plt.figure()
          sb.heatmap(param)
          plt.title(f"Epoch {epoch} Layer {layer_i} Param {param_idx}")
          plt.show()

        # 存到字典里,新增epoch维度
        if epoch == 0:
          # 初始新增维度,shape变为(原形状..., epoch数)
          self.weight_dict[param_key] = np.expand_dims(param, axis=-1)
        else:
          # 沿epoch维度拼接
          self.weight_dict[param_key] = np.concatenate(
              [self.weight_dict[param_key], np.expand_dims(param, axis=-1)],
              axis=-1
          )
    
    # 每个epoch结束后自动存到文件,避免训练中断丢失数据
    np.savez_compressed(self.save_path, **self.weight_dict)
    print(f"第{epoch}轮权重已保存到{self.save_path}")
调用方法

你原来的模型定义不需要修改,只需要在训练时传入修改后的回调即可:

gw = GetWeights(save_path="my_model_weights.npz")
# 你的模型定义、编译代码不变
model.fit(trainX, trainy, epochs=epochs, batch_size=batch_size_masukan, verbose=verbose, callbacks=[gw])
权重文件读取方法

后续需要读取权重时,直接用numpy加载即可:

loaded_weights = np.load("my_model_weights.npz")
# 查看所有保存的参数键
print(loaded_weights.files)
# 获取指定层指定参数的所有epoch数据,比如第0层第0个参数
w = loaded_weights["layer_0_param_0"]
# w的最后一维是epoch维度,比如w[...,0]是第0轮的参数,w[...,1]是第1轮的参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 18:15:04