如何在TensorFlow循环中不使用Python集合实现append操作适配Keras自定义层
问题根因
TensorFlow构建静态计算图时,不支持使用Python原生list存储动态生成的张量。你使用tf.range驱动的循环属于图内执行逻辑,Python侧的append操作无法被计算图追踪,会导致类型不兼容、梯度传递失败等错误。
解决方案1:最小改动(替换原生list为tf.TensorArray)
tf.TensorArray是TensorFlow官方提供的计算图兼容的动态张量数组,仅需替换原list相关操作即可修复问题,同时修正了原代码中compute_output_shape缩进错误、输出shape不符合要求的问题:
import tensorflow as tf from tensorflow.keras.layers import Layer import numpy as np class WeightedLayer(Layer): def __init__(self, n_input, n_memb, **kwargs): super(WeightedLayer, self).__init__(**kwargs) self.n = n_input # 16 features self.m = n_memb # 3 def build(self, batch_input_shape): super(WeightedLayer, self).build(batch_input_shape) def call(self, input_): self.batch_size = tf.shape(input_)[0] # 替换原生list为TensorArray CP = tf.TensorArray(dtype=input_.dtype, size=self.batch_size) for batch in tf.range(self.batch_size): xd_shape = [self.m] c_shape = [1] cp = input_[batch, 0, :] for d in range(1, self.n): c_shape.insert(0, self.m) xd_shape.insert(0, 1) xd = tf.reshape(input_[batch, d, :], xd_shape) c = tf.reshape(cp, c_shape) cp = tf.matmul(c, xd) flat_cp = tf.reshape(cp, (1, self.m**self.n)) # 替换append为write方法 CP = CP.write(batch, flat_cp) c = CP.stack() # 调整输出shape为要求的 (batch_size, 1, 3^16) return tf.reshape(c, (self.batch_size, 1, self.m**self.n)) def compute_output_shape(self, batch_input_shape): return tf.TensorShape([batch_input_shape[0], 1, self.m ** self.n]) # 测试代码 X_train = np.random.uniform(0, 1, (200, 16, 3)) X_test = np.random.uniform(0, 1, (200, 16, 3)) y_train = np.random.uniform(0, 1, (200,)) y_test = np.random.uniform(0, 1, (200,)) Model = tf.keras.models.Sequential() Model.add(WeightedLayer(n_input=16, n_memb=3, input_shape=(16, 3))) Model.compile(loss='mean_squared_error', optimizer='adam') Model.fit(X_train, y_train, epochs=20, batch_size=10, validation_data=(X_test, y_test))
解决方案2:高性能优化版(移除batch维度循环)
逐样本循环在batch较大时会严重拖慢训练速度,你可以利用TensorFlow的广播机制,直接对整个batch做并行计算,无需遍历每个样本:
class OptimizedWeightedLayer(Layer): def __init__(self, n_input, n_memb, **kwargs): super(OptimizedWeightedLayer, self).__init__(**kwargs) self.n = n_input # 16 features self.m = n_memb # 3 def build(self, batch_input_shape): super(OptimizedWeightedLayer, self).build(batch_input_shape) def call(self, input_): batch_size = tf.shape(input_)[0] # 初始cp形状为 (batch, 3) cp = input_[:, 0, :] for d in range(1, self.n): # 保留batch维度,添加广播所需的轴 xd = tf.reshape(input_[:, d, :], (batch_size, *([1]*d), self.m)) c = tf.reshape(cp, (batch_size, *([self.m]*d), 1)) cp = tf.matmul(c, xd) # 直接reshape为要求的输出形状 return tf.reshape(cp, (batch_size, 1, self.m**self.n)) def compute_output_shape(self, batch_input_shape): return tf.TensorShape([batch_input_shape[0], 1, self.m ** self.n])
内容的提问来源于stack exchange,提问作者zahra moradi
相关产品推荐
相关产品推荐

