使用JAX fori_loop时遇TypeError:'aval_method'对象不可调用求解决
问题分析与解决方案
错误根源
- 可变列表的副作用问题:
fori_loop要求状态是纯函数式的不可变结构,直接修改Python列表元素ums[0] = u属于副作用操作,JAX无法正确追踪这种可变状态。 - Pytree重建时的抽象值处理错误:
UB类的__init__中,当JAX用抽象值(AbstractValue)重建Pytree时,调用self.arr.reshape会触发aval_method不可调用的错误——抽象值的reshape方法在JAX追踪阶段无法直接执行。 - 类型判断逻辑失效:
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
相关产品推荐
相关产品推荐

