如何正确拼接维度尺寸不同的张量,解决tf.concat维度不匹配报错问题
错误原因
你使用tf.concat([t1,t2], 0)时报错,是因为指定了沿**第0维(batch维度)**做拼接,tf.concat要求除了拼接维度外,其余所有维度的尺寸必须完全相等。你两个张量的第1维(序列长度维度)分别是20和5,沿第0维拼接时要求该维度尺寸匹配,因此触发报错。
你的两个张量维度规则为(batch_size, 序列长度, 嵌入维度):
- 第0维:动态batch大小,对应你看到的
None - 第1维:序列长度,分别为20、5
- 第2维:嵌入维度,均为32
你需要得到(None, 25, 32)的输出,本质是要沿第1维(序列长度维度)拼接两个张量。
正确实现
直接将tf.concat的axis参数修改为1即可,也可以用负索引axis=-2(倒数第2个维度),写法鲁棒性更高:
# 两种写法均可 concat_tensor = tf.concat([t1, t2], axis=1) # concat_tensor = tf.concat([t1, t2], axis=-2)
效果验证
你可以用以下测试代码确认输出符合预期:
import tensorflow as tf # 模拟两个输入张量,batch大小固定为32 t1 = tf.random.normal((32, 20, 32)) t2 = tf.random.normal((32, 5, 32)) result = tf.concat([t1, t2], axis=1) print(result.shape) # 输出:(32, 25, 32)
注意事项
拼接前请确保两个张量的第0维(batch大小)完全一致,该维度不匹配时仍会触发报错。
内容的提问来源于stack exchange,提问作者user3284804
相关产品推荐
相关产品推荐

