news 2026/9/16 21:34:05

Flax 数据加载指南:如何将 Torchvision、TensorFlow 与 Hugging Face 数据集转换为 JAX 输入

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Flax 数据加载指南:如何将 Torchvision、TensorFlow 与 Hugging Face 数据集转换为 JAX 输入

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)。因此,从任何来源加载数据集,本质上都只做两件事:

  1. 转换:把数据(无论是 NumPy 数组、tf.Tensor还是 PIL Image)转换为jax.numpy类型;
  2. 整形:把数据 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参数控制取训练集还是测试集
转 NumPymnist[split].data.numpy()/mnist[split].targets.numpy()TorchVision 的 MNIST 对象内部是torch.Tensor,通过.numpy()转成 NumPy 数组
转 JAX + 归一化jnp.float32(...) / 255jnp.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(...) / 255jnp.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

这里的maptf.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 = 128num_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

三种加载方式对比与通用范式总结

对比项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()流式整体载入

无论走哪条路,最终都收敛到同一套四步通用范式

  1. 取原始数据:用各生态自己的 API 拿到数据对象;
  2. 转 NumPy.numpy()tfds.as_numpy()np.array(...)任选其一;
  3. 转 JAX 并归一化jnp.float32(x) / 255(图像场景),标签用jnp.int16
  4. 整形到模型输入约定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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 21:32:47

Wireshark抓包实战:过滤器、TCP重传、TLS解密与RTP还原

抓包这件事,说难不难,说简单也真容易翻车。我第一次打开 Wireshark 的时候,满屏花花绿绿的包在滚,脑子里只有一个念头:这玩意儿到底是给谁看的?后来踩了几次坑才回过味来——Wireshark 本身不是"分析工…

作者头像 李华
网站建设 2026/9/16 21:31:53

System Prompt泄露实战指南:从路径分析到防御与止损

system_prompts_leaks 这个词最近在圈子里被反复刷到,很多人把它当成一场“热闹的抓马”在看。但作为长期做 LLM 应用的人,我第一反应不是吃瓜,而是想起自己踩过的一个坑:某次内部 Agent 试运行,用户只是多问了一句“把…

作者头像 李华
网站建设 2026/9/16 21:31:07

cc-switch 走 TaoToken 通道后,WSL 里 claude 对话验证通过

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/16 21:30:40

AI行业爆发式增长与程序员转型七大黄金赛道

1. AI行业爆发式增长背后的技术驱动力2026年AI岗位预测增长10倍并非空穴来风。从技术演进轨迹来看,三大核心因素正在推动这一变革:首先是算力成本的指数级下降,使得企业部署AI解决方案的门槛大幅降低;其次是开源模型的成熟度提升&…

作者头像 李华
网站建设 2026/9/16 21:30:16

视觉驱动浏览器自动化:Qwen2.5-VL与Claude Computer Use实战

一直觉得,让AI自己看屏幕、自己动鼠标键盘把活干完,才是“AI替我打工”的真正形态。以前搞RPA要写一堆选择器、定位符,页面稍微改个class就废了;后来用脚本调用各种接口,又受限于平台开放程度。直到我把 browser-use、…

作者头像 李华