SensorFM 批量处理健康数据的内存占用降低步骤

文章导读
处理 SensorFM 批量健康数据时,进程内存持续上涨甚至被杀掉,往往不是单一原因,而是数据加载、长序列推理、中间张量累积共同作用。要降低内存占用,第一步是先定位峰值出现在哪个阶段,再用惰性加载、分段推理、显式回收来压住 RSS。以下步骤都可以用 psutil 和 tracemalloc 直接验证,不需要依赖官方内存数值。
📋 目录
  1. 用 tracemalloc 和 psutil 记录内存峰值位置
  2. 把数据加载改为惰性生成器
  3. 对长序列做分段推理并拼接结果
  4. 及时清理中间张量并显式回收内存
  5. 用恒定内存模式跑通一整天数据
A A

处理 SensorFM 批量健康数据时,进程内存持续上涨甚至被杀掉,往往不是单一原因,而是数据加载、长序列推理、中间张量累积共同作用。要降低内存占用,第一步是先定位峰值出现在哪个阶段,再用惰性加载、分段推理、显式回收来压住 RSS。以下步骤都可以用 psutil 和 tracemalloc 直接验证,不需要依赖官方内存数值。

内存上涨通常由一次性读入数据、长序列单次推理、中间张量未释放三处叠加造成。处理方向是先定位峰值阶段,再改惰性生成器、分段推理、及时回收。每一步以 RSS 回落为验证标准。优化效果与具体模型版本、输入长度相关,需要结合运行环境确认。

用 tracemalloc 和 psutil 记录内存峰值位置

先别急着改代码,用内存分析工具记录各阶段 RSS,确定是数据加载还是模型推理导致峰值。在数据加载、归一化、推理、后处理前后分别打印 RSS,同时用 tracemalloc 抓取占用前 10 的对象类型。

import psutil
import tracemalloc

def log_mem(tag):
    rss = psutil.Process().memory_info().rss / 1024**2
    print(f"{tag}: {rss:.1f} MB")

tracemalloc.start()
log_mem("start")

# 把下面各步骤替换成实际调用
data = load_all_data()   # 原数据加载
log_mem("after load")

data = normalize(data)   # 归一化
log_mem("after normalize")

outputs = model(data)    # 推理
log_mem("after inference")

results = postprocess(outputs)
log_mem("after postprocess")

current, peak = tracemalloc.get_traced_memory()
print(f"tracemalloc peak: {peak/1024**2:.1f} MB")
for stat in tracemalloc.take_snapshot().statistics('lineno')[:10]:
    print(stat)

运行后比较各阶段 RSS 增量:若加载后涨幅大,问题在数据读取;若推理后涨幅大,则需针对长序列做分段推理。同时注意 tracemalloc 本身会拖慢运行,仅在排查阶段开启。

把数据加载改为惰性生成器

如果内存峰值出现在数据加载阶段,原因是整个批次一次性读入。改成按时间窗口 yield 样本的生成器,让同一时间只有当前窗口驻留内存。下面是一个通用骨架,可替换成实际读取逻辑:

def load_windows(path, window_size, step):
    # 每次只读一个窗口,不保留全量数据
    for start in range(0, total_len, step):
        yield read_sensor_segment(path, start, start+window_size)

# 使用示例
for win in load_windows("recording.bin", window_size=300, step=300):
    pass  # 这里做归一化、推理

为了快速观察 RSS 是否稳定,可以用下面这段命令模拟生成器循环。它每次只构造一个窗口,然后打印 RSS,正常情况下 RSS 不会随迭代次数线性上涨:

SensorFM 批量处理健康数据的内存占用降低步骤
python -c "
import psutil, time

def windows():
    for i in range(10):
        data = bytearray(1024*1024)  # 假窗口
        yield data
        rss = psutil.Process().memory_info().rss / 1024**2
        print(f'step {i}: {rss:.1f} MB')
        time.sleep(0.1)

for w in windows():
    pass
"

如果替换后 RSS 仍然持续上涨,再看下一步的中间张量清理。

对长序列做分段推理并拼接结果

SensorFM 处理整晚睡眠记录时,单条样本可能超过模型最大输入长度,导致推理时显存/内存压力倍增。可以把长序列切成带重叠的窗口,每个窗口单独推理,只保留中间不含重叠部分的输出,再拼接成完整序列。

def split_with_overlap(seq, chunk_size, overlap):
    step = chunk_size - overlap
    for start in range(0, len(seq), step):
        yield seq[start:start+chunk_size]

def segment_infer(model, seq, chunk_size=256, overlap=32):
    step = chunk_size - overlap
    out_parts = []
    for start, chunk in enumerate(split_with_overlap(seq, chunk_size, overlap)):
        # 不足 chunk_size 时补零,也可以直接处理最后一段
        chunk = np.pad(chunk, (0, max(0, chunk_size - len(chunk))))
        out = model(chunk)
        if start == 0:
            keep = out[: chunk_size - overlap//2]
        else:
            keep = out[overlap//2 : chunk_size - overlap//2]
        out_parts.append(keep)
    return np.concatenate(out_parts, axis=0)

full_seq = np.random.randn(10000, 8)
result = segment_infer(model, full_seq)
print("输入长度:", len(full_seq), "输出长度:", len(result))

拼接后的长度应等于原始序列长度(或按步长略有差异),如果长度对不上,检查重叠部分的保留逻辑。这样可以单条推理改用多窗口,能显著降低单次推理的内存峰值。

SensorFM 批量处理健康数据的内存占用降低步骤

及时清理中间张量并显式回收内存

推理过程中,模型输出、中间激活、DataLoader 批次变量都会累积在内存里。即使改成分段推理,如果不主动释放,峰值依然会升高。每批推理后调用 del 删除中间张量,再用 gc.collect() 回收循环引用。

import gc, psutil

def process_batches(model, data_loader):
    baseline = psutil.Process().memory_info().rss
    results = []
    for i, batch in enumerate(data_loader):
        out = model(batch)
        results.append(out.detach().cpu())

        del out, batch     # 删除中间张量
        if i % 10 == 0:
            gc.collect()   # 回收循环引用

        rss = psutil.Process().memory_info().rss
        print(f"batch {i}: RSS {rss/1024**2:.1f} MB, baseline {baseline/1024**2:.1f} MB")
    return results

执行后观察每批 RSS 是否回落。由于 Python 解释器可能保留空闲内存,RSS 不会完全回到 baseline,但只要不再逐批上涨,就说明清理起到作用。若仍上涨,检查是否有全局引用持有了中间结果。

用恒定内存模式跑通一整天数据

最后把上面的步骤串起来,用模拟一整天数据跑批处理,每处理一小时记录一次 RSS,绘制内存时间曲线。目标是确认峰值不随数据量线性上涨,而是保持在一个近似水平。

import psutil, time, matplotlib.pyplot as plt

def simulate_day(model, hours=24):
    rss_history = []
    timestamps = []
    for hour in range(hours):
        # 生成一小时模拟数据并用惰性加载、分段推理处理
        data = generate_one_hour_data(hour)  # 替换成真实数据的模拟
        process_hour(data, model)

        rss = psutil.Process().memory_info().rss / 1024**2
        rss_history.append(rss)
        timestamps.append(hour)
        print(f"hour {hour}: RSS {rss:.1f} MB")
        time.sleep(0.1)

    plt.plot(timestamps, rss_history)
    plt.xlabel("hour")
    plt.ylabel("RSS (MB)")
    plt.savefig("memory_curve.png")
    return max(rss_history)

simulate_day(model)

观察曲线:如果 RSS 在开始几小时后进入平台期,说明流程已实现恒定内存。如果曲线持续上升,说明仍有某些对象在积累,例如 results 列表无限增长或缓存未清理。此时回到前几步,逐个阶段检查 RSS 增量。