TensorFlow如何实现类似PyTorch nn.ModuleList的可索引层列表
TensorFlow实现可索引层列表(对标PyTorch nn.ModuleList)方案
报错根因
你遇到的解包错误和层列表能力无关,完全是代码写法问题:
- 层列表
layers定义在了类作用域,没有赋值给实例属性self.layers,调用self.layers[i]时实际访问的是Keras Layer基类自带的空层列表,根本不是你实例化的GTLayer对象 - 构建层时错误使用
tf.range做循环,返回的是张量值而非Python原生整数,既既不能直接做i==0的布尔判断,也不能作为列表索引使用 - 无意义创建了Sequential实例但没有加入任何层,属于无效代码
- Keras自定义层前向传播需要重写
call方法,自定义forward方法不会触发Keras内置的权重跟踪、自动求导逻辑
正确实现方式
TensorFlow/Keras不需要特殊封装就能实现和nn.ModuleList完全一致的效果:
- 最简单的写法:在
__init__方法中创建普通Python列表,把实例化的自定义层逐个append进去,再把列表赋值给self的属性,Keras会自动跟踪列表内所有层的权重、支持梯度更新,完全兼容self.layers[i](*args)的索引调用写法 - 更稳妥的写法:直接使用
tf.keras.layers.LayerList,这个API就是官方对标PyTorch nn.ModuleList实现的,行为完全对齐
修正后的GTN完整实现代码
import tensorflow as tf from tensorflow import keras # 自定义GTLayer实现,继承keras.layers.Layer class GTLayer(keras.layers.Layer): def __init__(self, num_edge, num_channels, first=False, **kwargs): super().__init__(**kwargs) self.num_edge = num_edge self.num_channels = num_channels self.first = first # 此处补充你自己的层权重、子层初始化逻辑 def call(self, A, H=None): # 此处补充你自己的前向计算逻辑 # first=True时仅输入A,返回(H, W);其余情况输入A、H,返回(H, W) if self.first: # 首层计算逻辑 return H, W # 非首层计算逻辑 return H, W class GTN(keras.layers.Layer): def __init__(self, num_edge, num_channels, num_layers, **kwargs): super().__init__(**kwargs) self.num_layers = num_layers # 写法1:普通Python列表,足够应对绝大多数场景 self.layers = [] # 写法2:官方LayerList,行为和PyTorch nn.ModuleList完全一致 # self.layers = keras.layers.LayerList() # 必须用Python原生range循环,不能用tf.range for i in range(num_layers): if i == 0: self.layers.append(GTLayer(num_edge, num_channels, first=True)) else: self.layers.append(GTLayer(num_edge, num_channels, first=False)) def normalization(self, H): # 补充你自己的H归一化逻辑 return normalized_H # 自定义层前向传播必须重写call方法 def call(self, A, X=None, target_x=None, target=None): # 对应PyTorch的A.unsqueeze(0).permute(0,3,1,2) A = tf.transpose(tf.expand_dims(A, 0), perm=[0, 3, 1, 2]) Ws = [] H = None for i in range(self.num_layers): if i == 0: H, W = self.layers[i](A) else: H = self.normalization(H) H, W = self.layers[i](A, H) Ws.append(W) # 补充后续损失计算、预测输出等逻辑 return Ws, H
注意事项
- 所有子层必须在
__init__方法中实例化并存入挂载到self的列表,不要在call方法中创建层,不要把层列表定义为类属性,否则Keras无法跟踪权重,每次前向传播都会重新初始化参数 - 不要用
tf.keras.Sequential处理这类非串行结构:Sequential仅支持单输入单输出、层输出直接作为下一层输入的简单场景,GTN这种首层输入特殊、中间插入归一化操作、需要收集每层输出的结构完全不适用 - 层遍历、索引必须使用Python原生整数,不要用张量值作为列表索引
内容的提问来源于stack exchange,提问作者Anshuman Sinha
相关产品推荐
相关产品推荐

