关于「读取数据时内存溢出(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 才能构成一条样本)。
- 按日期切:窗口会跨越切分边界而损坏,样本不完整;
- 按股票切:每只股票的历史保持完整,各批之间互不影响,安全。
三个配合的技巧
- 分块查询 + 及时释放:每批只查
CHUNK_SIZE只股票,切完窗口后del df主动释放原始数据。 - 流式计算标准化统计:不要把所有窗口堆在一起再算
mean/std,而是用累加器(sum / sumsq / count)在线累加,结果与「全量算一次」在数值上等价,却不需要同时把全部数据留在内存。 - 推理时「读一批、预测一批、只留结果」:每批
查数 → 切窗口 → 标准化 → 立即预测,只累积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)
简言之:别一次性把全体分钟数据读进内存。按股票分小批读、用完即释放、只留最终结果,就能稳稳避开读数据阶段的内存溢出。