tf.keras.layers.Flatten()未生效原因咨询
问题原因与解决方案
原因解析
Keras的Flatten层是为神经网络批量输入场景设计的,默认保留第一个维度作为批量维度,仅展平从第二个维度开始的所有维度。你的输入张量x形状为(14, 2),在Keras层的逻辑中会被识别为(批量大小=14, 特征维度=2)——不存在需要展平的后续维度,因此输出形状和输入完全一致。
解决方案
方法1:适配Keras层的批量维度逻辑
给输入张量增加一个批量维度,让Flatten层能正确展平特征部分,之后可按需移除批量维度:
import tensorflow as tf # 创建形状为(14, 2)的张量 x = tf.constant([[1, 2], [3, 4], [5, 6], [7, 8], [9, 10], [11, 12], [13, 14], [15, 16], [17, 18], [19, 20], [21, 22], [23, 24], [25, 26], [27, 28]]) # 增加批量维度,形状变为(1, 14, 2) x_with_batch = tf.expand_dims(x, axis=0) y = tf.keras.layers.Flatten()(x_with_batch) # 移除批量维度,得到形状(28,)的张量 y = tf.squeeze(y) print(x) print(y)
方法2:直接用tf.reshape展平
如果不需要适配Keras的批量输入逻辑,直接使用tf.reshape会更简单直接:
import tensorflow as tf # 创建形状为(14, 2)的张量 x = tf.constant([[1, 2], [3, 4], [5, 6], [7, 8], [9, 10], [11, 12], [13, 14], [15, 16], [17, 18], [19, 20], [21, 22], [23, 24], [25, 26], [27, 28]]) # 直接展平为一维张量,-1表示自动计算维度大小 y = tf.reshape(x, (-1,)) print(x) print(y)
内容的提问来源于stack exchange,提问作者Rocky
相关产品推荐
相关产品推荐

