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

JAX版本升级后TypeError: 'Device'对象不可调用的解决方法

解决JAX中TypeError: 'Device' object is not callable的问题

错误原因分析

这个错误通常是因为自定义装饰器do_on_cpu中存在对Device对象的错误调用(比如直接写cpu()),或者试图给jax.random.PRNGKey传递它不支持的device参数——JAX 0.4.31版本中,PRNGKey并不接受device参数,且Device对象本身不可调用,直接调用会触发该错误。

正确的装饰器写法

以下两种方式可以实现让seed2key在CPU上运行:

方式1:用jax.jit指定设备执行函数

import jax
import jax.random as rand

def do_on_cpu(func):
    def wrapper(*args, **kwargs):
        # 获取第一个CPU设备
        cpu_device = jax.devices("cpu")[0]
        # 通过jax.jit指定函数在CPU上执行
        return jax.jit(func, device=cpu_device)(*args, **kwargs)
    return wrapper

@do_on_cpu
def seed2key(seed):
    return rand.PRNGKey(seed)

方式2:先执行函数再将结果转移到CPU

如果不需要强制函数在CPU上执行,只是确保最终结果位于CPU,可以用这种更轻量的方式:

import jax
import jax.random as rand

def do_on_cpu(func):
    def wrapper(*args, **kwargs):
        cpu_device = jax.devices("cpu")[0]
        result = func(*args, **kwargs)
        # 将结果转移到CPU设备
        return jax.device_put(result, cpu_device)
    return wrapper

@do_on_cpu
def seed2key(seed):
    return rand.PRNGKey(seed)

关键注意事项

  • 不要直接调用Device对象(比如jax.devices("cpu")[0]()是错误写法)
  • jax.random.PRNGKey本身不接受device参数,不要在装饰器中给它传递该参数
  • JAX 0.4.x版本中,指定函数执行设备的标准方式是通过jax.jit的device参数,或用jax.device_put转移数据所在设备

内容的提问来源于stack exchange,提问作者Raptor

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 19:49:59