TensorFlow 1.4中batch_normalization CPU报错及非降级解决方法咨询
问题分析与解决方案
潜在原因
你的问题核心出在两个关键点上:
axis参数设置不匹配格式:在NHWC数据格式下,通道维度是最后一维(索引为-1或3),但你指定了axis=1,这会让批量归一化操作针对高度维度(而非通道维度)计算。而TensorFlow 1.4中CPU版本的FusedBatchNorm(默认启用)仅支持对NHWC格式的最后一维(通道)做归一化,因此触发了格式不兼容的错误。- 版本间的API行为变化:TensorFlow 1.0中
tf.contrib.layers.batch_normalization默认没有启用fused模式,非融合的批量归一化实现对axis参数的限制更宽松,所以你的代码能正常运行。
无需回退版本的解决办法
这里有两个可行的修复方案,你可以根据实际需求选择:
方案1:修正axis参数(推荐)
既然你确认数据是NHWC格式,直接将axis改为通道对应的维度即可,这也是批量归一化的标准用法:
bn1 = tf.contrib.layers.batch_normalization(inputs=conv1, axis=-1, training=is_training) # 也可以显式写axis=3,效果完全一致
方案2:关闭fused模式
如果你因为特殊业务需求必须保留axis=1(比如需要对空间维度做归一化),可以通过设置fused=False禁用融合实现,绕开CPU FusedBatchNorm的格式限制:
bn1 = tf.contrib.layers.batch_normalization(inputs=conv1, axis=1, training=is_training, fused=False)
注意:禁用融合模式会略微降低性能,但能直接解决当前的兼容性问题。
另外你提到尝试给tf.layers.batch_normalization加fused和data_format参数报错,这是因为tf.layers.batch_normalization和tf.contrib.layers.batch_normalization是两个独立的API,前者确实不支持这两个参数,所以正确的调整对象是tf.contrib.layers.batch_normalization的参数,而非切换API版本。
内容的提问来源于stack exchange,提问作者Vignesh Venugopal
相关产品推荐
相关产品推荐

