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

无法创建/保存/加载超大磁盘数组,求高效方案及tensorflow.keras适配方法

问题描述

出于学习需求,我需要创建、保存并加载(导入tensorflow.keras)一个规模达10^10量级的超大int数组。

已尝试方案

  1. NumPy创建失败:
x=np.ones((274576,200,200,1),dtype='int')
  1. Dask创建成功但HDF5保存/加载效率低:
import dask.array as da
x=da.ones((274576,200,200,1),dtype='float')#create
da.to_hdf5('x.hdf5',{'x':x})#save
import h5py
y=h5py.File('x.hdf5','r')['x'][:]#read
print(y)

待解决问题

  1. 是否存在速度更快、空间效率更高的替代方案?
  2. 直接将Dask数组传入model.fit触发警告:
WARNING:tensorflow:Keras is training/fitting/evaluating on array-like data. Keras may not be optimized for this format, so if your input data format is supported by TensorFlow I/O (https://github.com/tensorflow/io) we recommend using that to load a Dataset instead.

需要执行哪些转换步骤才能避免该警告?


解决方案

一、超大数组的高效创建、保存与加载方案

针对10^10量级的int数组,推荐以下几种更高效的方案:

1. Dask结合Zarr格式替代HDF5

Zarr专为分布式/超大数组设计,并行读写性能优于HDF5,内存占用更低,完美适配Dask:

import dask.array as da
# 创建超大int数组,指定更小的int类型节省空间
x = da.ones((274576, 200, 200, 1), dtype='int32')
# 保存为Zarr
x.to_zarr('x.zarr', overwrite=True)
# 加载Zarr数组
y = da.from_zarr('x.zarr')
  • 核心优势:分块存储机制,读写仅处理当前所需块,内存占用可控;并行IO效率远高于HDF5。
  • 空间优化:根据实际数值范围选择最小可用int类型(如int8/int16/int32),避免默认int64的空间浪费。

2. 直接生成TensorFlow Dataset(训练场景优先)

如果数组用于Keras训练,无需完整保存文件,可直接用Dask生成Dataset,边生成边训练:

import dask.array as da
import tensorflow as tf

# 创建Dask数组
x = da.ones((274576, 200, 200, 1), dtype='int32')
# 转换为TensorFlow Dataset,按批次加载
def dask_to_tf_dataset(dask_arr, batch_size=32):
    return tf.data.Dataset.from_generator(
        lambda: dask_arr.to_delayed().ravel(),
        output_signature=tf.TensorSpec(shape=(200,200,1), dtype=tf.int32)
    ).batch(batch_size)

train_dataset = dask_to_tf_dataset(x)
  • 核心优势:无需一次性加载整个数组,节省磁盘和内存资源,直接对接训练流程。

3. TileDB格式存储

TileDB是面向多维数组的高效存储引擎,支持并行读写与高压缩比,兼容Dask和TensorFlow:

import dask.array as da
import tiledb

# 创建并保存Dask数组到TileDB
x = da.ones((274576, 200, 200, 1), dtype='int32')
da.to_tiledb('x_tiledb', x, overwrite=True)
# 加载TileDB数组
y = da.from_tiledb('x_tiledb')

二、避免Keras警告的转换步骤

消除警告的核心是将Dask数组转换为TensorFlow原生的tf.data.Dataset格式,有两种可靠方式:

1. 通过Dask-TensorFlow桥接转换

Dask提供直接转换工具,无需手动编写生成器:

from dask_tensorflow import from_dask

# 创建Dask数组(含特征和标签)
x = da.ones((274576, 200, 200, 1), dtype='int32')
y = da.zeros((274576, 1), dtype='int32')

# 转换为TensorFlow Dataset并指定批次
dataset = from_dask((x, y)).batch(32)

# 直接传入model.fit训练
model.fit(dataset, epochs=10)

2. 手动构建tf.data.Dataset

若不想依赖额外库,可利用Dask分块特性手动构建Dataset:

import dask.array as da
import tensorflow as tf

x = da.ones((274576, 200, 200, 1), dtype='int32')
y = da.zeros((274576, 1), dtype='int32')
# 将Dask数组拆分为延迟计算的块
x_blocks = x.to_delayed()
y_blocks = y.to_delayed()

# 定义生成器,逐个返回特征和标签块
def generator():
    for x_block, y_block in zip(x_blocks.ravel(), y_blocks.ravel()):
        yield x_block.compute(), y_block.compute()

# 构建Dataset并指定输出签名
dataset = tf.data.Dataset.from_generator(
    generator,
    output_signature=(
        tf.TensorSpec(shape=x_blocks.shape[1:], dtype=tf.int32),
        tf.TensorSpec(shape=y_blocks.shape[1:], dtype=tf.int32)
    )
)
# 打乱、分批(按需配置)
dataset = dataset.shuffle(1000).batch(32)

# 传入model.fit训练
model.fit(dataset, epochs=10)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 05:37:07