You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

M1 Max上jax-metal的anec.reshape不支持问题咨询

JAX-Metal在M1 Max上加载MNIST时的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.04 22:06:11