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

TensorFlow Dataset转多列DataFrame时内存占用过高问题咨询

大列数TFDS数据处理的内存优化与落地方案

1. 优化流程降低内存消耗的核心方法

(1)摒弃预分配全空DataFrame

你当前代码中预分配162078行×8500列的空DataFrame,仅每列的np.nan(float64类型)就会占用≈10GB内存,这是内存耗尽的核心原因。改为逐行/逐块收集有效数据,只存储非空值:

import tensorflow as tf
import pandas as pd
import numpy as np
from tqdm import tqdm

datapoint_indices = [x[0] for x in filtered_ranking_table]
column_names = ["class"] + [f'datapoint_{i}' for i in datapoint_indices]

# 定义稀疏数据类型,仅存储非空值
dtype_map = {'class': pd.SparseDtype(np.int32, np.nan)}
for col in column_names[1:]:
    dtype_map[col] = pd.SparseDtype(np.float64, np.nan)  # 根据实际数据类型调整

data_rows = []
batch_size = 10000  # 每批处理行数,避免内存堆积

for datapoint_n, clusters in tqdm(dataset.take(114003), total=114003):
    dp_n = datapoint_n.numpy()
    if dp_n not in datapoint_indices:
        continue
    dp_col = f'datapoint_{dp_n}'
    
    for i, cluster in enumerate(clusters):
        # 用numpy操作过滤0值,比列表推导高效
        cluster_arr = cluster.numpy()
        cluster_arr = cluster_arr[cluster_arr != 0]
        
        for val in cluster_arr:
            data_rows.append({'class': i, dp_col: val})
        
        # 达到批次大小就写入临时DataFrame并清空列表
        if len(data_rows) >= batch_size:
            temp_df = pd.DataFrame(data_rows).astype(dtype_map)
            df = pd.concat([df, temp_df]) if 'df' in locals() else temp_df
            data_rows = []

# 处理剩余未批次的数据
if data_rows:
    temp_df = pd.DataFrame(data_rows).astype(dtype_map)
    df = pd.concat([df, temp_df]) if 'df' in locals() else temp_df

df = df.dropna(how='all')

(2)使用稀疏数据结构

pandas的稀疏类型(SparseDtype)仅存储非空值,对于你的高稀疏度数据(大部分为NaN),内存占用可降低90%以上。上述代码已集成稀疏类型的使用。

(3)优化循环中的数据操作

  • 用numpy数组替代列表推导过滤0值,减少内存复制;
  • 避免df.loc切片赋值(会产生大量临时对象),改用字典逐行收集数据后批量转换。

2. 直接写入文件而非加载全量数据到内存

完全可以,这是处理大列数/大行数数据的最优方案,以下是几种主流格式的实现:

(1)Parquet(推荐,列式存储+高压缩)

Parquet适合大列数数据,支持高效压缩和列式查询,用pyarrow实现逐块写入:

import pyarrow as pa
import pyarrow.parquet as pq
from tqdm import tqdm

datapoint_indices = [x[0] for x in filtered_ranking_table]
# 定义Parquet schema
schema_fields = [pa.field('class', pa.int32())]
for dp in datapoint_indices:
    schema_fields.append(pa.field(f'datapoint_{dp}', pa.float64()))
schema = pa.schema(schema_fields)

# 初始化写入器,使用snappy压缩
writer = pq.ParquetWriter('output.parquet', schema, compression='snappy')
batch_size = 10000
data_rows = []

for datapoint_n, clusters in tqdm(dataset.take(114003), total=114003):
    dp_n = datapoint_n.numpy()
    if dp_n not in datapoint_indices:
        continue
    dp_col = f'datapoint_{dp_n}'
    
    for i, cluster in enumerate(clusters):
        cluster_arr = cluster.numpy()
        cluster_arr = cluster_arr[cluster_arr != 0]
        for val in cluster_arr:
            data_rows.append({'class': i, dp_col: val})
        
        if len(data_rows) >= batch_size:
            table = pa.Table.from_pylist(data_rows, schema=schema)
            writer.write_table(table)
            data_rows = []

# 写入剩余数据
if data_rows:
    table = pa.Table.from_pylist(data_rows, schema=schema)
    writer.write_table(table)

writer.close()

(2)HDF5(支持分块存储与快速查询)

用pandas的HDFStore实现append模式写入:

import pandas as pd
from tqdm import tqdm

with pd.HDFStore('output.h5', mode='w') as store:
    batch_size = 10000
    data_rows = []
    for datapoint_n, clusters in tqdm(dataset.take(114003), total=114003):
        dp_n = datapoint_n.numpy()
        if dp_n not in datapoint_indices:
            continue
        dp_col = f'datapoint_{dp_n}'
        
        for i, cluster in enumerate(clusters):
            cluster_arr = cluster.numpy()
            cluster_arr = cluster_arr[cluster_arr != 0]
            for val in cluster_arr:
                data_rows.append({'class': i, dp_col: val})
            
            if len(data_rows) >= batch_size:
                temp_df = pd.DataFrame(data_rows)
                store.append('data', temp_df, format='table', data_columns=True)
                data_rows = []
    
    if data_rows:
        temp_df = pd.DataFrame(data_rows)
        store.append('data', temp_df, format='table', data_columns=True)

(3)CSV(兼容性好但效率较低)

适合需要通用格式的场景,逐行写入避免内存占用:

import csv
from tqdm import tqdm

datapoint_indices = [x[0] for x in filtered_ranking_table]
column_names = ["class"] + [f'datapoint_{i}' for i in datapoint_indices]

with open('output.csv', 'w', newline='') as f:
    writer = csv.DictWriter(f, fieldnames=column_names)
    writer.writeheader()
    
    for datapoint_n, clusters in tqdm(dataset.take(114003), total=114003):
        dp_n = datapoint_n.numpy()
        if dp_n not in datapoint_indices:
            continue
        dp_col = f'datapoint_{dp_n}'
        
        for i, cluster in enumerate(clusters):
            cluster_arr = cluster.numpy()
            cluster_arr = cluster_arr[cluster_arr != 0]
            for val in cluster_arr:
                row = {col: '' for col in column_names}
                row['class'] = str(i)
                row[dp_col] = str(val)
                writer.writerow(row)

3. 针对大列数数据的其他可行思路

(1)用Dask替代pandas处理超内存数据

Dask支持分块并行处理,可直接处理超出内存的数据集,无需全量加载:

import dask.bag as db
import dask.dataframe as dd
from dask.diagnostics import ProgressBar

datapoint_indices = [x[0] for x in filtered_ranking_table]

def process_batch(batch_item):
    datapoint_n, clusters = batch_item
    dp_n = datapoint_n.numpy()
    if dp_n not in datapoint_indices:
        return []
    dp_col = f'datapoint_{dp_n}'
    rows = []
    for i, cluster in enumerate(clusters):
        cluster_arr = cluster.numpy()
        cluster_arr = cluster_arr[cluster_arr != 0]
        for val in cluster_arr:
            rows.append({'class': i, dp_col: val})
    return rows

# 将TFDS数据集转为Dask Bag,分10个分区处理
dask_bag = db.from_sequence(dataset.take(114003), npartitions=10)
dask_df = dask_bag.flat_map(process_batch).to_dataframe()

# 写入Parquet文件
with ProgressBar():
    dask_df.to_parquet('dask_output.parquet', compression='snappy')

(2)转置数据结构(按需选择)

如果后续分析以列维度为主,可以将数据转置为"行少列多"的结构,减少单条数据的内存占用,但需确保转置后符合分析需求。

(3)过滤不必要的列

如果部分datapoint_indices对应的列没有有效数据,可以提前过滤,减少列数。

内容的提问来源于stack exchange,提问作者lurum28

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 14:54:51