如何在TensorFlow自定义算子/层后添加slim.batch_norm?代码是否正确?
你的代码存在几个需要修正的地方,我们一步步拆解来看:
1. 与需求不符的参数问题
你原本想要调用 My_ops(x, 128),但代码里写的是 my_ops(x, 256),这里的通道数参数和你的预期需求不一致,需要先改成128。
2. Slim arg_scope 的冗余写法
你当前的 arg_scope 使用方式有点多余,with slim.arg_scope([slim.batch_norm], **batch_norm_params) as bn: 里的 as bn 完全没必要——因为你并没有用到这个变量。直接在 arg_scope 块里返回 slim.batch_norm(o) 是可行的,但我们可以让代码更简洁清晰。
3. Is_training 参数的正确性
你通过函数参数传递 is_training 并在 batch_norm_params 里设置的方式是正确的,这样能保证训练阶段开启批量归一化的均值/方差更新,推理阶段固定这些统计值。不过要注意调用这个函数时,训练场景传入 is_training=True,推理场景传入 is_training=False,别搞混场景。
修正后的完整代码:
def _build_block(self, x, name, is_training=True): with tf.variable_scope(name) as scope: # 修正通道数为需求的128 o = my_ops(x, 128) batch_norm_params = { 'decay': 0.9997, 'epsilon': 1e-5, 'scale': True, 'updates_collections': tf.GraphKeys.UPDATE_OPS, 'fused': None, # 使用可用的融合批归一化实现 'is_training': is_training } # 简化arg_scope的使用,去掉不必要的变量赋值 with slim.arg_scope([slim.batch_norm], **batch_norm_params): return slim.batch_norm(o)
额外小提示:
如果你希望代码可读性更强,也可以在调用 slim.batch_norm 时显式指定 is_training 参数(虽然 arg_scope 已经统一设置了,但显式传递能让阅读代码的人更直观看到这个关键参数):
return slim.batch_norm(o, is_training=is_training)
内容的提问来源于stack exchange,提问作者Jame
相关产品推荐
相关产品推荐

