自定义TensorFlow Layer中向量广播拼接时batch_size使用报错求助
问题场景
需要将单个向量广播至一批向量并完成拼接操作,但在自定义TensorFlow Layer中处理batch_size时遇到问题:构建阶段第一个输入的batch_size(第一维度)为None,调用阶段则为整数。
原实现代码
import sys import numpy as np import tensorflow as tf from tensorflow.keras.layers import Layer, Input, Concatenate, UpSampling2D from tensorflow.keras import Model print(f"Python Version: {sys.version}") print(f"Numpy Version: {np.version.version}") print(f"Tensorflow Version: {tf.version.VERSION}") class MyLayer(Layer): def __init__(self, **kwargs): super(MyLayer, self).__init__(**kwargs) self.concat = Concatenate(axis = 1) def call(self, inputs): x, y = inputs batchSize = tf.shape(x)[0] batchSize = 1 if batchSize is None else batchSize yShape = tf.shape(y) y = tf.reshape(y, (1, 1, yShape[1], 1)) y = UpSampling2D(size = (batchSize, 1))(y) y = tf.reshape(y, (batchSize, yShape[1])) return self.concat([x, y]) inputsX = Input(shape = (4, ), dtype = tf.float32) inputsY = Input(shape = (5, ), dtype = tf.float32) outputs = MyLayer()([inputsX, inputsY]) model = Model(inputs = [inputsX, inputsY], outputs = outputs) model.build(input_shape = ((None, 4), (1, 5))) model.summary() pointsX = tf.constant(np.reshape(range(12), [3, 4]), dtype = tf.float32) pointsY = tf.constant(np.reshape(range(5), [1, 5]), dtype = tf.float32) print(model((pointsX, pointsY)))
运行环境版本
Python Version: 3.8.10 (default, Jun 22 2022, 20:18:18) [GCC 9.4.0] Numpy Version: 1.23.4 Tensorflow Version: 2.11.0
预期与错误
预期得到一个3×9的张量,每行左侧4列为pointsX的值,右侧5列为pointsY的值,但运行时报错:
ValueError: The
sizeargument must be a tuple of 2 integers. Received: (<tf.Tensor 'my_layer/strided_slice:0' shape=() dtype=int32>, 1)including element Tensor("my_layer/strided_slice:0", shape=(), dtype=int32) of type <class 'tensorflow.python.framework.ops.Tensor'>Call arguments received by layer "my_layer" (type MyLayer):
• inputs=['tf.Tensor(shape=(None, 4), dtype=float32)', 'tf.Tensor(shape=(None, 5), dtype=float32)']
解决方案
错误原因
UpSampling2D的size参数要求传入静态整数,而原代码中batchSize = tf.shape(x)[0]得到的是动态Tensor,在模型构建阶段无法确定具体数值,因此触发类型错误。
正确实现
利用TensorFlow的动态广播机制,无需依赖UpSampling2D,直接将单个向量扩展至批量维度:
import sys import numpy as np import tensorflow as tf from tensorflow.keras.layers import Layer, Input, Concatenate from tensorflow.keras import Model print(f"Python Version: {sys.version}") print(f"Numpy Version: {np.version.version}") print(f"Tensorflow Version: {tf.version.VERSION}") class MyLayer(Layer): def __init__(self, **kwargs): super(MyLayer, self).__init__(**kwargs) self.concat = Concatenate(axis=1) def call(self, inputs): x, y = inputs batch_size = tf.shape(x)[0] # 将y从(1,5)广播至(batch_size,5) y_broadcasted = tf.tile(y, [batch_size, 1]) return self.concat([x, y_broadcasted]) inputsX = Input(shape=(4,), dtype=tf.float32) inputsY = Input(shape=(5,), dtype=tf.float32) outputs = MyLayer()([inputsX, inputsY]) model = Model(inputs=[inputsX, inputsY], outputs=outputs) model.build(input_shape=((None, 4), (1, 5))) model.summary() pointsX = tf.constant(np.reshape(range(12), [3, 4]), dtype=tf.float32) pointsY = tf.constant(np.reshape(range(5), [1, 5]), dtype=tf.float32) print(model((pointsX, pointsY)))
代码解释
- 动态获取批量大小:
tf.shape(x)[0]在调用阶段会自动获取实际的batch_size数值,TensorFlow会动态处理构建阶段的None情况,无需额外判断。 - 广播向量:
tf.tile(y, [batch_size, 1])将形状为(1,5)的y在第0维度重复batch_size次,得到形状为(batch_size,5)的张量,与x的批量维度完全匹配。 - 拼接操作:直接拼接
x和广播后的y,得到预期的(batch_size, 9)形状的张量。
运行结果
执行后会输出符合预期的3×9张量:
tf.Tensor( [[ 0. 1. 2. 3. 0. 1. 2. 3. 4.] [ 4. 5. 6. 7. 0. 1. 2. 3. 4.] [ 8. 9. 10. 11. 0. 1. 2. 3. 4.]], shape=(3, 9), dtype=float32)
内容的提问来源于stack exchange,提问作者louis

