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 Datasets(TFDS)和 Hugging Facedatasets三大生态加载数据,并完成类型转换、像素归一化与维度重塑,使数据满足 Flax 模型(B, 28, 28, 1)的输入约定。读完本文,你将掌握一套通用的"加载 → 转 NumPy → 转 JAX 数组 → 整形"数据接入范式,并能结合 Flax 官方示例中的完整输入管线(含 shuffle、batch、prefetch)与多设备评估时的 padding 技巧,把任意来源的数据平滑接入 JAX+Flax 训练与评估流程。
本文对应的原始文档为 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( logits=logits, labels=batch['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', train=True, download=True), 'test': torchvision.datasets.MNIST('./data', train=False, download=True) } 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', train=True/False, download=True) | 首次运行会下载到本地./data目录,之后直接复用;train参数控制取训练集还是测试集 |
| 转 NumPy | mnist[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 轴(axis=3)上追加灰度通道,得到(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 Datasets(TFDS)是 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(split=split, 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(split=split, 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 → prefetch:
def 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', split='train') test_ds: tf.data.Dataset = tfds.load('mnist', split='test') 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_remainder=True).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_remainder=True).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 isn't 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。
三种加载方式对比与通用范式总结
| 对比项 | TorchVision | TensorFlow Datasets | Hugging Facedatasets |
|---|---|---|---|
| 加载接口 | torchvision.datasets.MNIST(...) | tfds.builder(...).as_dataset(...) | load_dataset("mnist") |
| 原始图像类型 | torch.Tensor | tf.Tensor | PIL 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)], axis=0) 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_remainder=True) → 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),仅供参考