如何改写支持批量图像的Keras GradCam实现以移除tf.map_fn
移除GradCam批量处理中的tf.map_fn以提升性能
我参考Keras的GradCam示例实现了批量处理图像的功能,但目前使用了tf.map_fn,希望移除该函数来提升性能。当前代码中需要改写的片段是:
def single_image(index): return last_conv_layer_output[index] @ pooled_grads[index][tf.newaxis, ..., tf.newaxis] heatmaps = tf.map_fn(single_image, tf.range(tf.shape(grads)[0]), dtype=tf.float32)
解决方案
可以通过张量广播与批量运算直接实现,完全不需要循环或tf.map_fn,代码如下:
# 扩展pooled_grads的维度,匹配特征图的空间维度 pooled_grads_expanded = tf.expand_dims(tf.expand_dims(pooled_grads, axis=1), axis=1) # 批量计算特征图与权重的乘积,再沿通道维度求和得到热力图 heatmaps = tf.reduce_sum(last_conv_layer_output * pooled_grads_expanded, axis=-1)
逻辑说明
这段代码和原tf.map_fn的逻辑完全等价,但利用TensorFlow的批量运算特性,避免了逐样本循环,性能会明显提升:
pooled_grads的原始形状是(batch_size, channels),扩展维度后变为(batch_size, 1, 1, channels),可以和last_conv_layer_output(形状(batch_size, H, W, channels))直接做元素级乘法- 最后沿通道维度(
axis=-1)求和,得到每个样本的热力图(形状(batch_size, H, W))
内容的提问来源于stack exchange,提问作者Fabian N.
相关产品推荐
相关产品推荐

