Flax 数据加载指南如何将 Torchvision、TensorFlow 与 Hugging Face 数据集转换为 JAX 输入【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax导读在 Flax 中编写神经网络第一步就是让数据进入jax.numpy的世界。本指南以手写数字识别数据集 MNIST 为例系统讲解如何分别通过 Torchvision、TensorFlow DatasetsTFDS和 Hugging Facedatasets三大生态加载数据并完成类型转换、像素归一化与维度重塑使数据满足 Flax 模型(B, 28, 28, 1)的输入约定。读完本文你将掌握一套通用的加载 → 转 NumPy → 转 JAX 数组 → 整形数据接入范式并能结合 Flax 官方示例中的完整输入管线含 shuffle、batch、prefetch与多设备评估时的 padding 技巧把任意来源的数据平滑接入 JAXFlax 训练与评估流程。本文对应的原始文档为 loading_datasets.md含同内容的可执行 Notebook loading_datasets.ipynb并在此基础上结合 examples/mnist/train.py、examples/imagenet/input_pipeline.py 等仓库源码进行纵深扩充。核心思想一切数据最终都要变成jax.numpy数组用 JAX Flax 编写的神经网络其输入数据必须是jax.numpy数组实例即jnp.ndarray。因此从任何来源加载数据集本质上都只做两件事转换把数据无论是 NumPy 数组、tf.Tensor还是 PIL Image转换为jax.numpy类型整形把数据 reshape/expand 到网络期望的维度。本指南选用 MNIST 作为贯穿案例因为它足够简单且信息明确MNIST 由28×28 像素的灰度手写数字图像组成官方划分60k 训练 / 10k 测试任务是预测每张图像所属的类别数字 0~9。假设我们要训练一个 CNN 分类器则输入数据应满足形状(B, 28, 28, 1)其中末尾的单一维度表示灰度图像的通道数channel标签则是与图像一一对应的整数0~9形状应为(B,)。标签使用整数而非 one-hot 向量是为了直接配合optax.softmax_cross_entropy_with_integer_labels计算损失——这一点在仓库的 MNIST 示例中得到了印证examples/mnist/train.py 的loss_fn正是用该损失函数将batch[image]与整数标签batch[label]计算交叉熵def loss_fn(model: CNN, batch, rngs): logits model(batch[image], rngs) loss optax.softmax_cross_entropy_with_integer_labels( logitslogits, labelsbatch[label] ).mean() return loss, logits先导入两个基础库后面的三种加载方式都会用到import numpy as np import jax.numpy as jnp关于内存的说明本指南演示的是将整个数据集一次性载入内存的做法MNIST 全集约 32 MB完全可行。对于内存装不下的数据集处理流程是类似的只是需要改为按批次batchwise流式处理这部分在文末会结合仓库源码展开。从torchvision.datasets加载Torchvision 是 PyTorch 生态的视觉工具库内置了 MNIST、CIFAR 等常见视觉数据集的下载与管理接口。import torchvision def get_dataset_torch(): mnist { train: torchvision.datasets.MNIST(./data, trainTrue, downloadTrue), test: torchvision.datasets.MNIST(./data, trainFalse, downloadTrue) } ds {} for split in [train, test]: ds[split] { image: mnist[split].data.numpy(), label: mnist[split].targets.numpy() } # cast from np to jnp and rescale the pixel values from [0,255] to [0,1] ds[split][image] jnp.float32(ds[split][image]) / 255 ds[split][label] jnp.int16(ds[split][label]) # torchvision returns shape (B, 28, 28). # hence, append the trailing channel dimension. ds[split][image] jnp.expand_dims(ds[split][image], 3) return ds[train], ds[test]逐行拆解这段代码它完整展示了源生态 → NumPy → JAX的转换链路步骤代码说明下载/加载torchvision.datasets.MNIST(./data, trainTrue/False, downloadTrue)首次运行会下载到本地./data目录之后直接复用train参数控制取训练集还是测试集转 NumPymnist[split].data.numpy()/mnist[split].targets.numpy()TorchVision 的 MNIST 对象内部是torch.Tensor通过.numpy()转成 NumPy 数组转 JAX 归一化jnp.float32(...) / 255用jnp.float32显式转换类型同时把像素值从[0, 255]缩放到[0, 1]标签类型jnp.int16(...)标签保持整数类型避免与 softmax 交叉熵的浮点计算混淆补通道维jnp.expand_dims(..., 3)TorchVision 返回的是(B, 28, 28)在第 3 轴axis3上追加灰度通道得到(B, 28, 28, 1)验证加载结果train, test get_dataset_torch() print(train[image].shape, train[image].dtype) print(train[label].shape, train[label].dtype) print(test[image].shape, test[image].dtype) print(test[label].shape, test[label].dtype) # 预期输出 # (60000, 28, 28, 1) float32 # (60000,) int16 # (10000, 28, 28, 1) float32 # (10000,) int16从tensorflow_datasets加载TensorFlow DatasetsTFDS是 TensorFlow 生态的数据集仓库提供统一的tfds.builder/tfds.load接口和标准化的 split 语义。Flax 仓库的多数示例MNIST、ImageNet、LM1B、WMT 等都基于 TFDS 构建数据管线。import tensorflow_datasets as tfds def get_dataset_tf(): mnist tfds.builder(mnist) mnist.download_and_prepare() ds {} for split in [train, test]: ds[split] tfds.as_numpy(mnist.as_dataset(splitsplit, batch_size-1)) # cast to jnp and rescale pixel values ds[split][image] jnp.float32(ds[split][image]) / 255 ds[split][label] jnp.int16(ds[split][label]) return ds[train], ds[test]关键点解析tfds.builder(mnist)创建数据集构建器download_and_prepare()负责下载并预处理已下载过则直接复用缓存可从~/.cache/ 指定data_dir读取as_dataset(splitsplit, batch_size-1)返回tf.data.Dataset其中batch_size-1表示一次性把整个 split 打包成一个 batch即整体载入内存tfds.as_numpy()是把tf.data.Dataset转成 NumPy 数组的关键 API——这正是本文一切数据转 NumPy思想的体现它会把数据集中的tf.Tensor全部物化为 NumPy 数组之后再用jnp.float32(...) / 255、jnp.int16(...)完成向 JAX 类型的转换与归一化注意TFDS 的 MNIST 返回的图像本身已经是(B, 28, 28, 1)带通道维因此这里不需要再调用expand_dims。仓库实战MNIST 官方示例的 TFDS 输入管线examples/mnist/train.py 中的get_datasets展示了更贴近真实训练需求的 TFDS 用法——不再整体载入而是保留tf.data.Dataset的流式能力并串联map → shuffle → batch → prefetchdef get_datasets( config: ml_collections.ConfigDict, ) - tuple[tf.data.Dataset, tf.data.Dataset]: Load MNIST train and test datasets into memory. batch_size config.batch_size train_ds: tf.data.Dataset tfds.load(mnist, splittrain) test_ds: tf.data.Dataset tfds.load(mnist, splittest) train_ds train_ds.map( lambda sample: { image: tf.cast(sample[image], tf.float32) / 255, label: sample[label], } ) # normalize train set test_ds test_ds.map( lambda sample: { image: tf.cast(sample[image], tf.float32) / 255, label: sample[label], } ) # normalize the test set. # Create a shuffled dataset by allocating a buffer size of 1024 to randomly # draw elements from. train_ds train_ds.shuffle(1024) # Group into batches of batch_size and skip incomplete batches, prefetch the # next sample to improve latency. train_ds train_ds.batch(batch_size, drop_remainderTrue).prefetch(1) # Group into batches of batch_size and skip incomplete batches, prefetch the # next sample to improve latency. test_ds test_ds.batch(batch_size, drop_remainderTrue).prefetch(1) return train_ds, test_ds这里的map用tf.cast(sample[image], tf.float32) / 255完成了与文档一致的归一化只是把jnp换成tf随后在训练循环中通过train_ds.as_numpy_iterator()逐 batch 取出 NumPy 数据喂给nnx.jit编译的train_step见 examples/mnist/train.py。batch_size等超参来自 examples/mnist/configs/default.py默认batch_size 128、num_epochs 10。这一实践路径说明对于可流式消费的数据集不必先用tfds.as_numpy整体物化直接在tf.data.Dataset上完成归一化、打乱、分批最后在消费端用.as_numpy_iterator()转 NumPy 即可——JAX/Flax 与tf.data的配合是官方示例中的标准姿势。从 Hugging Facedatasets加载Hugging Face 的datasets库提供统一的load_dataset接口覆盖图像、文本、语音等多种模态。MNIST 在该生态中同样可用一行代码加载。#!pip install datasets # datasets isnt preinstalled on Colab; uncomment to install from datasets import load_dataset def get_dataset_hf(): mnist load_dataset(mnist) ds {} for split in [train, test]: ds[split] { image: np.array([np.array(im) for im in mnist[split][image]]), label: np.array(mnist[split][label]) } # cast to jnp and rescale pixel values ds[split][image] jnp.float32(ds[split][image]) / 255 ds[split][label] jnp.int16(ds[split][label]) # append trailing channel dimension ds[split][image] jnp.expand_dims(ds[split][image], 3) return ds[train], ds[test]要点说明load_dataset(mnist)返回一个DatasetDict内含train/test两个 split与 TorchVision、TFDS 不同Hugging Face 数据集中的image字段是PIL Image 对象列表因此需要先用列表推导np.array(im)逐张转 NumPy再统一np.array(...)堆叠成(B, 28, 28)数组其余步骤与 TorchVision 版本完全一致jnp.float32 / 255归一化、jnp.int16处理标签、jnp.expand_dims(..., 3)补通道维得到(B, 28, 28, 1)由于datasets默认不随 Colab 环境安装首次使用需先执行pip install datasets。三种加载方式对比与通用范式总结对比项TorchVisionTensorFlow DatasetsHugging Facedatasets加载接口torchvision.datasets.MNIST(...)tfds.builder(...).as_dataset(...)load_dataset(mnist)原始图像类型torch.Tensortf.TensorPIL Image 列表转 NumPy 方式.numpy()tfds.as_numpy()np.array([np.array(im) for im in ...])是否自带通道维否需expand_dims(..., 3)是MNIST 默认(B,28,28,1)否需expand_dims(..., 3)内存友好度整体载入可用as_numpy_iterator()流式整体载入无论走哪条路最终都收敛到同一套四步通用范式取原始数据用各生态自己的 API 拿到数据对象转 NumPy.numpy()、tfds.as_numpy()、np.array(...)任选其一转 JAX 并归一化jnp.float32(x) / 255图像场景标签用jnp.int16整形到模型输入约定jnp.expand_dims(x, 3)或reshape到(B, H, W, C)。进阶数据放不下内存怎么办按批处理与多设备 padding上文提到当数据集超出内存容量时流程不变但必须改为按批次处理。两种可行路径路径一TFDS as_numpy_iterator()流式消费。保持tf.data.Dataset的惰性训练循环里逐 batch 取数。这正是 examples/mnist/train.py 的做法for batch in train_ds.as_numpy_iterator(): train_step(...)归一化、shuffle、batch、prefetch 全部由tf.data在后台完成。路径二多设备/多主机评估时对最后一个不完整 batch 做 padding。当 batch 大小不能被设备数整除时最后一个 batch 形状不同会触发 XLA 的重新编译甚至在多主机 SPMD 场景下导致psum等待挂起。Flax 提供的解决方案是flax.jax_utils.pad_shard_unpad——它在主机内存中把输入 padding 到设备数整除、再 shard 到各设备、计算完成后 unshard 并 unpad。该方法源自 big_vision实现在 flax/jax_utils.py核心逻辑为def pad(x): _, *shape x.shape db, rest divmod(b, d) # d jax.local_device_count() if rest: x np.concatenate([x, np.zeros((d - rest, *shape), x.dtype)], axis0) db 1 ... return x.reshape(d, db, *shape) # shard: (d, db, ...)其典型用法是装饰jax.pmap后的前向函数pad_shard_unpad jax.pmap def forward(params, x): ...并对 padding 前后、多主机分片不均匀、动态过滤导致的 batch 数不一致等边界情况给出了系统解法详见配套文档 full_eval.rst含手写 padding 循环、pad_shard_unpad封装、static_argnums/static_return以及无限 padding等完整讨论。多主机分片则可用tfds.split_for_jax_process()或手动按jax.process_count()/jax.process_index()切分——examples/imagenet/input_pipeline.py 的create_split就展示了按进程数均分训练/验证集、SkipDecoding延迟解码、map(AUTOTUNE) → batch(drop_remainderTrue) → repeat → prefetch的完整多主机 ImageNet 管线。这些实现共同印证了本文的核心结论数据接入 JAX 的关键始终是转成jax.numpy 匹配模型输入形状无论是单机全量载入还是多设备流式、多主机分片都围绕这一不变式展开。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考