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
相关产品推荐
相关产品推荐

