Keras/TensorFlow批量计算mean_iou时内存溢出致SIGKILL问题求助
结合你的场景——2048x1024的大分辨率图像、每batch更新混淆矩阵计算mean_iou,30步后触发SIGKILL,大概率是计算过程中TensorFlow的张量/计算图节点没有被正确释放,导致显存持续累积占用,最终触发OOM(内存不足)终止。生成器单独跑正常,说明问题出在训练时模型计算与指标更新的交互上。下面是针对性的解决方案:
1. 用Keras自定义Metric封装mean_iou计算(最关键)
不要在训练循环里手动更新混淆矩阵,而是把指标计算逻辑封装成Keras的自定义Metric类。Keras会自动管理这些变量的生命周期,避免每步都创建新的张量节点导致内存泄漏。
示例代码:
from keras import backend as K from keras.metrics import Metric class CustomMeanIoU(Metric): def __init__(self, num_classes, name='mean_iou', **kwargs): super(CustomMeanIoU, self).__init__(name=name, **kwargs) self.num_classes = num_classes # 初始化混淆矩阵权重,仅创建一次 self.total_cm = self.add_weight(name='total_cm', shape=(num_classes, num_classes), initializer='zeros') def update_state(self, y_true, y_pred, sample_weight=None): # 转换为类别索引(假设你的输出是one-hot编码) y_true = K.argmax(y_true, axis=-1) y_pred = K.argmax(y_pred, axis=-1) # 计算当前batch的混淆矩阵并累加 cm = K.confusion_matrix(y_true, y_pred, num_classes=self.num_classes) self.total_cm.assign_add(cm) def result(self): # 计算各类别的IoU并取均值 sum_over_row = K.sum(self.total_cm, axis=0) sum_over_col = K.sum(self.total_cm, axis=1) true_positives = K.diag(self.total_cm) denominator = sum_over_row + sum_over_col - true_positives iou = true_positives / (denominator + K.epsilon()) # 避免除零 return K.mean(iou) def reset_states(self): # 每个epoch开始时重置混淆矩阵 K.set_value(self.total_cm, K.zeros_like(self.total_cm))
使用时,在模型编译阶段直接传入这个metric:
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=[CustomMeanIoU(num_classes=你的类别数)])
这样训练时Keras会自动在每个batch更新混淆矩阵,无需手动干预,从根源避免内存泄漏。
2. 开启TensorFlow显存增长模式
GTX1080的显存是8G,大分辨率图像batch_size=2已经很吃显存。开启显存增长模式可以让TensorFlow按需分配显存,避免一开始就占满全部显存,给后续计算留余量。
在代码最开头添加:
import tensorflow as tf from keras import backend as K config = tf.ConfigProto() config.gpu_options.allow_growth = True # 按需分配显存 sess = tf.Session(config=config) K.set_session(sess)
3. 优化生成器的内存回收
虽然生成器单独跑正常,但训练时生成器与模型计算并行,可能临时图像数据没有被及时回收。可以在生成器里显式清理临时变量:
def train_generator(...): while True: # 加载并预处理图像 img = load_image(...) img = preprocess_image(img) label = load_label(...) label = preprocess_label(label) yield (img, label) # 显式删除临时变量并触发垃圾回收 del img, label import gc gc.collect() # 可以每几个batch调用一次,避免频繁调用影响速度
4. 避免手动计算时保留Tensor引用
如果你必须手动在训练循环里计算混淆矩阵,不要直接累加Tensor对象,而是计算出数值后再更新Python变量:
# ❌ 错误:每次batch都创建新Tensor,导致内存累积 confusion_matrix += tf.confusion_matrix(y_true, y_pred, num_classes=num_classes) # ✅ 正确:计算出数值后更新Python变量 cm_val = sess.run(tf.confusion_matrix(y_true, y_pred, num_classes=num_classes)) confusion_matrix += cm_val
5. 考虑升级依赖版本(可选)
Keras2.1.2和TensorFlow1.4.1都是比较老的版本,存在一些已知的内存管理bug。如果项目允许,可以升级到Keras2.2.x(对应TensorFlow1.10.x左右),新版本在自定义指标和内存管理上更稳定。
内容的提问来源于stack exchange,提问作者Tobias Scheck

