M1 Max上jax-metal的anec.reshape不支持问题咨询
anec.reshape警告问题解答 背景信息
在M1 Max款MacBook Pro(系统:MacOS Sonoma v14.1,jax-metal版本:0.0.4)上测试jax-metal时,运行以下加载MNIST数据集的代码:
import gzip import os import struct import urllib.request import jax.numpy as jnp def mnist(): url_dir = "https://storage.googleapis.com/cvdf-datasets/mnist" target_dir = os.getcwd() + "/data/mnist" # download images and labels into data folder url = f"{url_dir}/train-images-idx3-ubyte.gz" target = f"{target_dir}/train-images-idx3-ubyte.gz" if not os.path.exists(target): os.makedirs(target_dir, exist_ok=True) urllib.request.urlretrieve(url, target) print(f"Downloaded {url} to {target}") # load images into memory target = f"{target_dir}/train-images-idx3-ubyte.gz" with gzip.open(target, "rb") as fh: _, batch, rows, cols = struct.unpack(">IIII", fh.read(16)) shape = (batch, 1, rows, cols) images = jnp.frombuffer(fh.read(), dtype=jnp.uint8).reshape(shape) return images if __name__ == "__main__": images = mnist() print(images.shape)
出现如下警告:
loc("jit(reshape)/jit(main)/reshapenew_sizes=(60000, 1, 28, 28) dimensions=None): error: 'anec.reshape' op failed: input tensor dimensions are not supported on ANEs.
程序可正常运行,但存在以下疑问,以下是对应解答:
1. anec.reshape不支持会带来哪些后果?
这个警告的核心是:JAX尝试将reshape操作调度到Apple Neural Engine(ANE)执行,但当前jax-metal版本不支持该维度的reshape操作。此时JAX会自动回退到CPU或Metal GPU执行该reshape步骤,程序逻辑不会受影响,唯一的影响是这个reshape操作无法利用ANE的硬件加速,执行速度可能比预期稍慢,但不会导致程序报错或功能异常。
2. 能否通过规避reshape操作来消除该警告,还是只能等待anec.reshape获得完整支持?
完全可以通过调整数据加载方式规避该警告,无需等待官方更新。核心思路是将reshape操作从JAX的计算图中转移到CPU端的NumPy处理流程中,避免触发ANE上的reshape尝试。
修改后的代码示例:
import gzip import os import struct import urllib.request import numpy as np import jax.numpy as jnp def mnist(): url_dir = "https://storage.googleapis.com/cvdf-datasets/mnist" target_dir = os.path.join(os.getcwd(), "data", "mnist") # 下载数据集 url = f"{url_dir}/train-images-idx3-ubyte.gz" target = os.path.join(target_dir, "train-images-idx3-ubyte.gz") if not os.path.exists(target): os.makedirs(target_dir, exist_ok=True) urllib.request.urlretrieve(url, target) print(f"已下载 {url} 到 {target}") # 加载并处理数据:先通过NumPy完成reshape,再转为JAX数组 with gzip.open(target, "rb") as fh: _, batch, rows, cols = struct.unpack(">IIII", fh.read(16)) shape = (batch, 1, rows, cols) # 先用NumPy读取并reshape,再转成JAX数组 images_np = np.frombuffer(fh.read(), dtype=np.uint8).reshape(shape) images = jnp.array(images_np) return images if __name__ == "__main__": images = mnist() print(images.shape)
这样修改后,reshape操作在CPU上由NumPy完成,JAX仅负责将NumPy数组转换为JAX数组,不会触发ANE的reshape操作请求,警告自然会消失。
内容的提问来源于stack exchange,提问作者Hericks

