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

如何避免创建Ragged Tensor时内存溢出?实现类似flow_from_dataframe的流式加载

解决方案:流式加载并创建Ragged Tensor

你遇到的内存崩溃问题源于一次性将10000个张量全部加载到内存中,TensorFlow的tf.data.Dataset API可以实现类似flow_from_dataframe的流式加载逻辑,避免一次性占用过多内存。以下是具体实现步骤:

1. 将文件路径整理为DataFrame

先把所有.npy文件的路径存入pandas DataFrame,方便后续流式读取:

import pandas as pd
import glob

# 获取并排序文件路径
train_tensors_paths = sorted(glob.glob('/content/drive/MyDrive/dataset/*.npy'), key=lambda x: x.split('/')[-1])

# 转为DataFrame
df = pd.DataFrame({'file_path': train_tensors_paths})

2. 定义流式加载函数

编写一个加载单个.npy文件并转换为张量的函数,用tf.py_function包装以兼容TensorFlow Dataset:

import tensorflow as tf
import numpy as np

def load_npy_file(file_path):
    # 将TensorFlow字符串转为Python字符串
    file_path_str = file_path.numpy().decode('utf-8')
    # 加载npy文件(可根据需求选择是否用mmap_mode)
    np_array = np.load(file_path_str)
    # 转换为TensorFlow张量
    return tf.convert_to_tensor(np_array)

3. 创建流式Dataset并处理Ragged Tensor

利用tf.data.Dataset从DataFrame读取路径,映射加载函数,实现惰性加载。如果需要处理变长张量(对应Ragged Tensor的场景),可以直接在Dataset中保留变长数据,或通过批量转换为Ragged Tensor:

# 从DataFrame创建Dataset
dataset = tf.data.Dataset.from_tensor_slices(df['file_path'].values)

# 映射加载函数,注意用tf.py_function包装
dataset = dataset.map(lambda x: tf.py_function(load_npy_file, [x], tf.float32))  # 根据你的数据类型调整dtype

# 可选:批量转换为Ragged Tensor(适合变长样本)
dataset = dataset.batch(32).map(lambda batch: tf.ragged.stack(batch))

# 迭代验证(不会一次性加载所有数据)
for batch in dataset:
    print(batch.shape)
    # 这里可以加入你的训练逻辑

关键说明

  • 惰性加载:Dataset只会在迭代(或训练)时才加载对应批次的文件,不会一次性将所有10000个张量存入内存,从根源避免内存崩溃。
  • 变长数据适配:tf.ragged.stack可以将批次内的变长张量转换为Ragged Tensor,完美替代一次性创建tf.ragged.constant的逻辑。
  • 性能优化:可以进一步添加prefetch(tf.data.AUTOTUNE)或cache()(如果内存允许缓存部分数据)来提升加载效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 04:01:42