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

scikit-learn多线程场景下内存泄漏问题求助

scikit-learn多线程场景下内存泄漏问题求助

嘿,我太懂你这种头疼的感觉了——明明手动清了变量、调用了GC,内存还是蹭蹭往上涨,而且只在多线程模式下出问题,确实挺闹心的。结合我之前碰到过的类似情况,咱们来捋一捋:

为啥会出现这种情况?

你猜的方向完全对!这个内存泄漏就是n_jobs≠1时,sklearn启动的子线程搞的鬼。sklearn底层一般用joblib来实现多线程,子线程里生成的临时工作对象,主线程的gc.collect()根本管不到——因为子线程有自己的引用链,这些对象没被主线程的GC扫描到,就一直占着内存。

尤其是OPTICS这类需要在多线程里处理大量中间数据的算法,每次循环重新实例化模型时,都会创建新的线程池和相关对象,旧的线程池残留的内存没被彻底释放,循环次数一多就漏得越来越明显。

给你几个亲测有效的解决办法:

1. 把模型实例移到循环外面复用

别每次循环都重新创建OPTICS和TSNE实例!把模型初始化放在循环之前,这样就能避免重复创建线程池和相关内存对象,内存占用会稳定很多:

import gc
import numpy as np
from sklearn.manifold import TSNE
from sklearn.cluster import OPTICS
import psutil
process = psutil.Process()

def main():
    data = np.random.random((100, 100))
    # 提前初始化模型,循环里复用
    optics = OPTICS(n_jobs=2)
    tsne = TSNE()
    for _i in range(1, 50):
        points = tsne.fit_transform(data)
        prediction = optics.fit_predict(points)
        # 清理操作保持不变
        points = None
        prediction = None
        del prediction, points
        gc.collect()
        print(process.memory_info().rss)

main()

2. 用joblib的上下文管理器显式管理线程池

既然sklearn用joblib处理多线程,咱们就直接用joblib的上下文管理器来控制线程池的生命周期,用完就强制回收:

import gc
import numpy as np
from sklearn.manifold import TSNE
from sklearn.cluster import OPTICS
import psutil
from joblib import parallel_backend

process = psutil.Process()

def main():
    data = np.random.random((100, 100))
    # 用上下文管理器指定线程后端,自动回收资源
    with parallel_backend('threading', n_jobs=2):
        for _i in range(1, 50):
            points = TSNE().fit_transform(data)
            prediction = OPTICS(n_jobs=2).fit_predict(points)
            del prediction, points
            gc.collect()
            print(process.memory_info().rss)

main()

3. 换成多进程模式兜底

如果线程模式的泄漏实在搞不定,那就换多进程试试——进程结束后系统会直接回收所有内存,不会有残留。sklearn的部分算法支持通过backend参数切换:

# 初始化OPTICS时指定多进程后端
prediction = OPTICS(n_jobs=2, backend='multiprocessing').fit_predict(points)

不过要注意,多进程的启动开销比线程大一点,小数据量的话可能会慢一丢丢,但内存问题肯定能解决。

额外小工具:定位泄漏点

要是你想搞清楚到底是哪些对象在占内存,可以用tracemalloc来追踪:

import tracemalloc

tracemalloc.start()
# 运行你的循环代码
snapshot = tracemalloc.take_snapshot()
top_stats = snapshot.statistics('lineno')

print("Top 10内存泄漏点:")
for stat in top_stats[:10]:
    print(stat)

这样能精准看到哪行代码生成的对象没被回收,方便针对性处理。

备注:内容来源于stack exchange,提问作者nelolpp

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 16:44:35