自定义TensorFlow MaxPooling层报错:tf.function无法提取张量
问题描述
尝试实现步长为2、池化尺寸为2、"same"填充的自定义MaxPooling2D层时,触发如下错误:
ValueError: Argument
initial_value(Tensor("custom_max_pooling2d/zeros:0", shape=(None, 64, 28, 28), dtype=float32)) could not be lifted out of atf.function. (Tried to create variable with name='None'). To avoid this error, when constructingtf.Variables inside oftf.functionyou can create theinitial_valuetensor in atf.init_scopeor pass a callableinitial_value(e.g.,tf.Variable(lambda : tf.truncated_normal([10, 40]))). Please file a feature request if this restriction inconveniences you.Call arguments received by layer "custom_max_pooling2d" (type CustomMaxPooling2D):
• inputs=tf.Tensor(shape=(None, 28, 28, 64), dtype=float32)
同时发现fit方法中指定的batch_size=128未生效。要求保留底层池化逻辑,不调用TensorFlow内置池化函数,仅修正CustomMaxPooling2D类的问题。
完整代码
import tensorflow as tf fashion_mnist = tf.keras.datasets.fashion_mnist (x_train_full, y_train_full), (x_test, y_test) = fashion_mnist.load_data() x_valid, x_train = x_train_full[:5000]/255.0, x_train_full[5000:]/255.0 y_valid, y_train = y_train_full[:5000], y_train_full[5000:] x_test = x_test/255.0 tensorboard_writer = tf.summary.create_file_writer('tensorboard/') class ValidationLossCallback(tf.keras.callbacks.Callback): def on_epoch_end(self, epoch, logs=None): with tensorboard_writer.as_default(): val_loss = logs['val_loss'] tf.summary.scalar('val_loss', val_loss, step=epoch) tensorboard_writer.flush() loss = tf.keras.losses.SparseCategoricalCrossentropy() opt = tf.keras.optimizers.Adam() early_stopping = tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=25, restore_best_weights=True) class CustomMaxPooling2D(tf.keras.layers.Layer): def __init__(self): super(CustomMaxPooling2D, self).__init__() self.padding_format = tf.constant([[0, 0], [1, 1], [1, 1], [0, 0]]) self.transpose_format = tf.constant([0, 3, 1, 2]) self.reverse_transpose_format = tf.constant([0, 2, 3, 1]) self.batch_size = tf.Variable(128, dtype=tf.int32) def build(self, input_shape): _, self.height, self.width, self.channels = input_shape self.output_width = tf.floor((self.height-1)/2)+1 def call(self, inputs): self.out= tf.Variable(tf.zeros(shape=(self.batch_size, self.channels, self.width, self.width))) if self.width%self.output_width!=0: inputs = tf.pad(tensor=inputs, paddings=self.padding_format, mode='constant') inputs = tf.transpose(inputs, perm=self.transpose_format) for b in range(self.batch_size): for c in range(self.channels): input_data = inputs[b][c] for w in range(self.output_width): for h in range(self.output_width): self.out[int(b),int(c),int(w),int(h)].assign(tf.reduce_max(input_data[int(w)*2:(int(w)+1)*2, int(h)*2:(int(h)+1)*2])) return tf.transpose(self.out, perm=self.reverse_transpose_format) model = tf.keras.models.Sequential([ tf.keras.layers.Conv2D(64,7,activation="relu",padding="same",input_shape=[28,28,1]), CustomMaxPooling2D(), tf.keras.layers.Conv2D(128,3,activation="relu",padding="same"), tf.keras.layers.Conv2D(128,3,activation="relu",padding="same"), CustomMaxPooling2D(), tf.keras.layers.Conv2D(256,3,activation="relu",padding="same"), tf.keras.layers.Conv2D(256,3,activation="relu",padding="same"), CustomMaxPooling2D(), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activation="relu"), tf.keras.layers.Dropout(0.5), tf.keras.layers.Dense(64, activation="relu"), tf.keras.layers.Dropout(0.5), tf.keras.layers.Dense(10, activation="softmax") ]) model.compile(optimizer=opt, loss=loss) model.fit(x_train, y_train, batch_size=128, epochs=10000,shuffle=True, validation_data=(x_valid, y_valid), callbacks=[early_stopping, ValidationLossCallback()]) predictions = model.predict(x_test) # Step 5: Evaluate the predictions accuracy = tf.keras.metrics.Accuracy() accuracy.update_state(tf.argmax(predictions, axis=1), y_test) accuracy_result = accuracy.result().numpy() print("Accuracy: {:.2f}%".format(accuracy_result * 100)) model.save_weights('models/maxpoolmodel.h5')
核心问题点
call方法中创建tf.Variable:call会被TensorFlow转为tf.function,其中不能动态创建变量,池化输出是临时计算结果,无需用变量存储。- 硬编码
batch_size:固定的batch_size变量无法适配fit传入的批次(含验证时的不同批次),需从输入张量动态获取。 - Python原生循环不兼容图模式:嵌套Python循环在图模式下效率极低,且强制类型转换(如
int(w))易报错。 - 填充逻辑错误:固定上下左右各填1不符合"same"填充的动态计算规则。
修正后的CustomMaxPooling2D类
class CustomMaxPooling2D(tf.keras.layers.Layer): def __init__(self, pool_size=2, strides=2): super(CustomMaxPooling2D, self).__init__() self.pool_size = pool_size self.strides = strides self.transpose_format = tf.constant([0, 3, 1, 2]) self.reverse_transpose_format = tf.constant([0, 2, 3, 1]) def build(self, input_shape): _, self.height, self.width, self.channels = input_shape # 计算same填充后的输出尺寸 self.output_height = tf.cast(tf.math.ceil(self.height / self.strides), tf.int32) self.output_width = tf.cast(tf.math.ceil(self.width / self.strides), tf.int32) # 动态计算填充量 pad_height = (self.output_height * self.strides - self.height) self.pad_top = pad_height // 2 self.pad_bottom = pad_height - self.pad_top pad_width = (self.output_width * self.strides - self.width) self.pad_left = pad_width // 2 self.pad_right = pad_width - self.pad_left def call(self, inputs): # 动态获取当前批次大小 batch_size = tf.shape(inputs)[0] # 应用动态计算的same填充 inputs = tf.pad(inputs, [[0, 0], [self.pad_top, self.pad_bottom], [self.pad_left, self.pad_right], [0, 0]], mode='constant') # 转置为(batch, channel, height, width)格式 inputs_transposed = tf.transpose(inputs, perm=self.transpose_format) # 定义单样本池化逻辑,用vectorized_map批量处理 def pool_single_sample(sample): output = tf.TensorArray(tf.float32, size=self.channels) for c in tf.range(self.channels): channel_data = sample[c] channel_output = tf.TensorArray(tf.float32, size=self.output_height) for h in tf.range(self.output_height): start_h = h * self.strides end_h = start_h + self.pool_size row_output = tf.TensorArray(tf.float32, size=self.output_width) for w in tf.range(self.output_width): start_w = w * self.strides end_w = start_w + self.pool_size patch = channel_data[start_h:end_h, start_w:end_w] row_output = row_output.write(w, tf.reduce_max(patch)) channel_output = channel_output.write(h, row_output.stack()) output = output.write(c, channel_output.stack()) return output.stack() # 批量处理所有样本 pooled = tf.vectorized_map(pool_single_sample, inputs_transposed) # 转置回原格式(batch, height, width, channels) return tf.transpose(pooled, perm=self.reverse_transpose_format)
修复说明
- 移除不必要的变量:用
tf.TensorArray临时存储计算结果,避免在call中创建tf.Variable。 - 动态适配批次:通过
tf.shape(inputs)[0]获取当前输入的批次大小,不再硬编码。 - 动态计算填充:根据输入和输出尺寸的关系,精准计算"same"填充的上下左右填充量。
- 兼容图模式的循环:使用
tf.vectorized_map和tf.TensorArray实现循环,既保留底层池化逻辑,又适配TensorFlow的图模式执行。 - 增强通用性:将
pool_size和strides设为初始化参数,方便调整池化配置。
内容的提问来源于stack exchange,提问作者Bryan Carty

