You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Keras/TensorFlow批量计算mean_iou时内存溢出致SIGKILL问题求助

解决Keras训练时内存泄漏导致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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.19 07:33:50