关于「读取数据时内存溢出(OOM)」的说明与解决

由hxgre创建,最终由hxgre 被浏览 13 用户

一、现象

不少同学把端到端 Transformer 示例(Transformer_modelsave_predict)拿去提交后,运行时报内存溢出(Out Of Memory),进程被系统杀掉。观察日志会发现,崩溃几乎都发生在读取数据、构建数据集的阶段,而不是模型训练或推理本身。


二、为什么会溢出

问题出在数据读取的方式上。示例里的 build_dataset 是这样取数的:

sql = f"SELECT date, instrument, {', '.join(FEATURE_COLS)} FROM {table} ORDER BY instrument, date"
df = dai.query(sql, filters={"date": [buf, ed], "instrument": instruments}).df()

这一句会把整个区间、全体股票的 1 分钟原始 K 线一次性全部读进内存(.df() 会把结果整体 materialize 成一个 DataFrame)。

分钟级数据的体量很容易被低估:

  • 1 只股票 1 个交易日约有 240 根 1 分钟 bar;
  • 成分股约 1200 只,一段较长区间有 很多个交易日
  • 粗算就是 上千万行,再乘以十几列特征。

真正喂给模型的样本其实很小——每天只在收盘决策点取一个「回看窗口」,但在切出这些小样本之前,那一大坨原始分钟数据已经把内存占满了。这就是崩溃的根因。


三、解决思路:分块读取,用完即释放

核心只有一句话:

按股票分成小批查询,每批切完窗口就立刻释放原始行,让内存里始终只留下「一小批股票的原始数据」,而不是「全体」。

为什么按「股票」分批,而不是按「日期」分批?

因为回看窗口需要单只股票在时间上连续(连续的 N 根 bar 才能构成一条样本)。

  • 日期切:窗口会跨越切分边界而损坏,样本不完整;
  • 股票切:每只股票的历史保持完整,各批之间互不影响,安全。

三个配合的技巧

  1. 分块查询 + 及时释放:每批只查 CHUNK_SIZE 只股票,切完窗口后 del df 主动释放原始数据。
  2. 流式计算标准化统计:不要把所有窗口堆在一起再算 mean/std,而是用累加器(sum / sumsq / count)在线累加,结果与「全量算一次」在数值上等价,却不需要同时把全部数据留在内存。
  3. 推理时「读一批、预测一批、只留结果」:每批 查数 → 切窗口 → 标准化 → 立即预测,只累积 date / instrument / score 三列极小的结果,窗口和原始数据当批就释放。这样峰值内存与区间长度、股票总数无关,只由 CHUNK_SIZE 决定——这也是提交后最需要省内存的一步。

四、参考模板代码

我们提供了一份分块读取版的完整可运行示例,模型结构、保存格式、输出格式与原示例完全一致,唯一的区别就是取数方式。两版对照着看,就能直接明白改了哪里、为什么这么改:

关键片段:分块读取的骨架

CHUNK_SIZE = 50   # 每次只查这么多只股票的原始数据,用完即释放

def _chunks(seq, size):
    for i in range(0, len(seq), size):
        yield seq[i:i + size]

def _query_chunk(table, buf, ed, chunk):
    """查一小批股票的原始数据(只有这一步会把原始行放进内存)。"""
    sql = f"SELECT date, instrument, {', '.join(FEATURE_COLS)} FROM {table} ORDER BY instrument, date"
    df = dai.query(sql, filters={"date": [buf, ed], "instrument": chunk}).df()
    # ...(按需做 log1p 等预处理)
    return df

关键片段:推理时「读一批、预测一批、只留结果」

out = []
for chunk in _chunks(instruments, CHUNK_SIZE):
    df = _query_chunk(table, buf, ed, chunk)      # 只查这一小批
    wins, keys = build_windows(df)                # 切窗口
    del df                                        # ★原始行用完立即释放

    X = normalize(np.stack(wins), mean, std)
    scores = predict_in_batches(model, X)         # 立即预测

    chunk_df = pd.DataFrame(keys, columns=["date", "instrument"])
    chunk_df["score"] = scores
    out.append(chunk_df)                          # 只保留三列结果
    del wins, keys, X                             # 窗口/张量当批释放

result = pd.concat(out, ignore_index=True)

关键片段:流式计算标准化统计

class RunningStats:
    """按字段累加 sum / sumsq / count,最后一次算出 mean/std,
    与「全量堆在一起算一次」数值上等价,但无需同时留住全部数据。"""
    def __init__(self, n_feat):
        self.n = 0
        self.s  = np.zeros(n_feat, np.float64)
        self.ss = np.zeros(n_feat, np.float64)

    def update(self, x):                          # x: (m, n_feat)
        self.n  += x.shape[0]
        self.s  += x.sum(0)
        self.ss += (x.astype(np.float64) ** 2).sum(0)

    def finalize(self):
        mean = self.s / self.n
        var  = self.ss / self.n - mean ** 2
        std  = np.sqrt(np.clip(var, 0, None)) + 1e-6
        return mean.astype(np.float32), std.astype(np.float32)

五、几点建议

  • CHUNK_SIZE 是内存和速度的调节旋钮。模板里默认设成较保守的 50:值越小越省内存、耗时略增;如果运行环境内存有余量,可以适当调大以兼顾速度。遇到 OOM 就先把它调小。

  • 优先改造推理路径。提交后跑的是推理,这里最容易 OOM,收益也最大;训练同理适用分块,只是本地训练时资源通常更宽裕。

  • 不要用 SELECT * 全表扫描。只 SELECT 用得到的字段,并始终通过 filters 传入日期区间,既省内存也更快。

  • 善用 SDK 的流式读取dai 结果对象支持按 batch 流式读取,同样能避免一次性占满内存:

    reader = result.fetch_arrow_reader(batch_size=10000)
    for batch in reader:
        process(batch)
    

简言之:别一次性把全体分钟数据读进内存。按股票分小批读、用完即释放、只留最终结果,就能稳稳避开读数据阶段的内存溢出。

{link}