使用flax.jax_utils.prefetch_to_device预取128维数组迭代器报错
解决
flax.jax_utils.prefetch_to_device在CPU环境下的维度匹配错误 问题原因
flax.jax_utils.prefetch_to_device默认会尝试把输入数据按第一维度分片分配到所有可用设备上。你的每个样本是128维的一维数组,函数误将这个128当成了需要分片的设备数量,但当前只有1个CPU设备,所以触发了len(shards) = 128 must equal len(devices) = 1的报错。
解决方案
这里提供两种可行的修复方式:
方案一:将单个样本包装成单元素批次
通过给样本增加一个维度,让函数识别这是单个批次而非分片数据,代码修改如下:
import tensorflow_datasets as tfds import tensorflow as tf import jax import jax.numpy as jnp import flax def _sift1m_iter(): def prepare_tf_data(xs): def _prepare(x): dl_arr = tf.experimental.dlpack.to_dlpack(x) jax_arr = jax.dlpack.from_dlpack(dl_arr) return jax_arr # 给单个样本增加一个批次维度(变成形状(1,128)) embedding = jax.tree_util.tree_map(_prepare, xs['embedding']) return (embedding[None, :],) # 也可以用字典格式:{'embedding': embedding[None, :]} ds = tfds.load('sift1m', split='database') it = map(prepare_tf_data, ds) it = flax.jax_utils.prefetch_to_device(it, 2) # 现在可以正常运行 # 注意:迭代时需要去掉额外的批次维度,比如:for batch in it: embedding = batch[0] return it
方案二:CPU环境下手动实现预取逻辑
如果不需要多设备适配,手动写轻量的预取逻辑更直接,避免flax工具函数的多设备默认行为:
import tensorflow_datasets as tfds import tensorflow as tf import jax import jax.numpy as jnp from collections import deque import flax def _sift1m_iter(): def prepare_tf_data(xs): def _prepare(x): dl_arr = tf.experimental.dlpack.to_dlpack(x) jax_arr = jax.dlpack.from_dlpack(dl_arr) return jax_arr return jax.tree_util.tree_map(_prepare, xs['embedding']) ds = tfds.load('sift1m', split='database') it = map(prepare_tf_data, ds) # 手动维护预取缓冲区,大小设为2 prefetch_buffer = deque(maxlen=2) # 先填充初始缓冲区 for _ in range(2): try: prefetch_buffer.append(next(it)) except StopIteration: break # 迭代输出数据,同时补充缓冲区 while prefetch_buffer: yield prefetch_buffer.popleft() try: prefetch_buffer.append(next(it)) except StopIteration: pass return it
补充说明
- 方案一的核心是让数据的第一维度长度和设备数量(1)匹配,这样
prefetch_to_device就不会触发分片逻辑。 - 方案二更适合单CPU场景,逻辑简单直接,不需要修改样本的结构。
内容的提问来源于stack exchange,提问作者jeffreyveon
相关产品推荐
相关产品推荐

