TensorFlow 0.12中如何反转flatten操作恢复张量原始形状?
恢复Flatten后的张量形状
嘿,这事儿其实挺简单的,用 TensorFlow 的 tf.reshape() 就能完美解决!
原理说明
tf.contrib.layers.flatten() 干的核心事儿就是把输入张量除第一个 batch 维度之外的所有维度,全部展平成一维——你这里就是把 [64, 32, 256, 2] 转成了 [64, 32*256*2],也就是 [64, 16384]。那反转操作其实就是把展平后的张量重新 reshape 回原始的多维形状就行。
具体代码实现
你可以直接指定完整的原始形状,针对性很强:
# 假设你展平后的张量叫做 flattened_output,形状是 [64, 16384] original_shape = [64, 32, 256, 2] restored_tensor = tf.reshape(flattened_output, original_shape)
如果想让代码更通用(比如后续 batch size 可能调整),可以用 -1 让 TensorFlow 自动推断 batch 维度的大小,这样不管 batch 是 64 还是其他数值,都能正确恢复:
restored_tensor = tf.reshape(flattened_output, [-1, 32, 256, 2])
注意事项
要确保展平后的元素总数和原始形状的元素总数完全一致——你这里 64*32*256*2 = 64*16384,刚好匹配,所以不会有问题。如果总数不匹配,TensorFlow 会直接抛出错误提示你。
另外,你的环境是 TensorFlow 0.12 + Python 2.7,tf.reshape() 在这个版本里完全支持这种用法,放心用就行~
内容的提问来源于stack exchange,提问作者Peter111
相关产品推荐
相关产品推荐

