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
相关产品推荐
相关产品推荐

