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

使用JAX fori_loop时遇TypeError:'aval_method'对象不可调用求解决

问题分析与解决方案

错误根源

  1. 可变列表的副作用问题:fori_loop要求状态是纯函数式的不可变结构,直接修改Python列表元素ums[0] = u属于副作用操作,JAX无法正确追踪这种可变状态。
  2. Pytree重建时的抽象值处理错误:UB类的__init__中,当JAX用抽象值(AbstractValue)重建Pytree时,调用self.arr.reshape会触发aval_method不可调用的错误——抽象值的reshape方法在JAX追踪阶段无法直接执行。
  3. 类型判断逻辑失效:type(arr) is not object的判断无法识别JAX抽象数组类型,导致在错误时机执行reshape操作。

修复步骤

1. 替换可变列表为不可元组

JAX控制流要求状态不可变,将Python列表改为元组,通过创建新元组的方式实现元素更新:

def s_w(ub, ums):
    e = jnp.identity(2)
    u = UM(e, [2])
    # 用切片创建新元组替代列表修改操作
    ums = (u,) + ums[1:]
    return ub, ums

2. 修正UB类初始化逻辑

将reshape操作限定在实际数组初始化阶段,避免在Pytree重建时执行,同时修正类型判断:

class UB():
    def __init__(self, arr, new_shape):
        self.shape = new_shape
        # 用isinstance准确识别JAX数组类型
        if isinstance(arr, jnp.ndarray):
            self.arr = arr.reshape(new_shape + new_shape)
        else:
            self.arr = arr
    
    def _tree_flatten(self):
        children = (self.arr,)
        aux_data = {'new_shape': self.shape}
        return (children, aux_data)

    @classmethod
    def _tree_unflatten(cls, aux_data, children):
        # 直接返回实例,不在unflatten阶段执行reshape
        return cls(*children, **aux_data)

3. 调整s_c函数的状态初始化

将ums初始化为元组而非列表,确保JAX能正确追踪状态:

def s_c(t, uns):
    n = 20
    # 直接创建元组,避免使用可变列表
    ums = tuple(UM(un, [2]) for un in uns)

    tub = UB(t.arr, t.r)
    
    s_loop_body = lambda i, x: s_w(ub=x[0], ums=x[1])
    
    tub, ums = jax.lax.fori_loop(0, n, s_loop_body, (tub, ums))

    return jnp.array([u.arr.flatten() for u in ums])

4. 优化UM类参数处理

确保r参数始终为元组,避免初始化时的类型不一致:

class UM():
    def __init__(self, arr, r=None):
        self.arr = arr
        # 统一转换为元组,兼容列表输入
        self.r = tuple(r) if r is not None else None
    
    def _tree_flatten(self):
        children = (self.arr,)
        aux_data = {'r': self.r}
        return (children, aux_data)

    @classmethod
    def _tree_unflatten(cls, aux_data, children):
        return cls(*children, **aux_data)

完整修复代码

import jax.numpy as jnp
import jax

class UB():
    def __init__(self, arr, new_shape):
        self.shape = new_shape
        if isinstance(arr, jnp.ndarray):
            self.arr = arr.reshape(new_shape + new_shape)
        else:
            self.arr = arr
    
    def _tree_flatten(self):
        children = (self.arr,)
        aux_data = {'new_shape': self.shape}
        return (children, aux_data)

    @classmethod
    def _tree_unflatten(cls, aux_data, children):
        return cls(*children, **aux_data)


class UM():
    def __init__(self, arr, r=None):
        self.arr = arr
        self.r = tuple(r) if r is not None else None
    
    def _tree_flatten(self):
        children = (self.arr,)
        aux_data = {'r': self.r}
        return (children, aux_data)

    @classmethod
    def _tree_unflatten(cls, aux_data, children):
        return cls(*children, **aux_data)


for C in [UB, UM]:
    jax.tree_util.register_pytree_node(
        C,
        C._tree_flatten,
        C._tree_unflatten,
    )


def s_w(ub, ums):
    e = jnp.identity(2)
    u = UM(e, [2])
    ums = (u,) + ums[1:]
    return ub, ums

def s_c(t, uns):
    n = 20
    ums = tuple(UM(un, [2]) for un in uns)

    tub = UB(t.arr, t.r)
    
    s_loop_body = lambda i, x: s_w(ub=x[0], ums=x[1])
    
    tub, ums = jax.lax.fori_loop(0, n, s_loop_body, (tub, ums))

    return jnp.array([u.arr.flatten() for u in ums])


uns = jnp.array([jnp.array([1, 2, 3, 4]) for _ in range(6)])
t = UM(jnp.array([1, 0, 0, 1]), r=[2])
uns = s_c(t, uns)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 08:05:24