You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义TensorFlow Layer中向量广播拼接时batch_size使用报错求助

问题:自定义TensorFlow Layer中广播单个向量至批量向量并拼接的错误解决

问题场景

需要将单个向量广播至一批向量并完成拼接操作,但在自定义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 size argument 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)))

代码解释

  1. 动态获取批量大小:tf.shape(x)[0]在调用阶段会自动获取实际的batch_size数值,TensorFlow会动态处理构建阶段的None情况,无需额外判断。
  2. 广播向量:tf.tile(y, [batch_size, 1])将形状为(1,5)的y在第0维度重复batch_size次,得到形状为(batch_size,5)的张量,与x的批量维度完全匹配。
  3. 拼接操作:直接拼接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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 11:30:47