JAX @jit嵌套类方法调用报错:参数重复TypeError问题求解
问题分析与解决
错误原因
One类参数传递顺序错误:
One的__init__方法定义为def __init__(self, plant, x),但_tree_unflatten中调用cls(*children, **aux_data)时,children对应(self.x,),aux_data对应{'plant': self.plant},这会导致调用cls(2, plant=plant_instance)——位置参数2被传给第一个形参plant,同时又通过关键字参数传递plant,引发"multiple values"冲突错误。Plant类的
_tree_unflatten签名不符合规范:
JAX要求pytree的_tree_unflatten类方法必须接收三个参数:cls、aux_data、children,但你定义的Plant._tree_unflatten只接收cls和children,这会导致JAX在反序列化时调用出错。
修正后的代码
import jax.numpy as jnp from jax import grad, jit, vmap from jax import random import numpy as np from functools import partial from jax import tree_util class One: def __init__(self, plant, x): self.plant = plant self.x = x @jit def call_plant_func(self, y): out = self.plant.func(y) + self.x return out def _tree_flatten(self): children = (self.x,) # arrays / dynamic values aux_data = {'plant': self.plant} # static values return (children, aux_data) @classmethod def _tree_unflatten(cls, aux_data, children): # 按__init__参数顺序传递:先传plant,再传x return cls(aux_data['plant'], *children) tree_util.register_pytree_node(One, One._tree_flatten, One._tree_unflatten) class Plant: def __init__(self, z, kk): self.z = z @jit def func(self, y): y = y + self.z return y def _tree_flatten(self): children = (self.z,) # arrays / dynamic values aux_data = None # static values return (children, aux_data) @classmethod def _tree_unflatten(cls, aux_data, children): # 保留规范签名,即使aux_data未使用 # 原__init__的kk参数未被使用,这里传默认值或按需调整 return cls(*children, kk=0) tree_util.register_pytree_node(Plant, Plant._tree_flatten, Plant._tree_unflatten) plant = Plant(5, 2) one = One(plant, 2) print(one.call_plant_func(10)) # 输出: 17
额外优化建议
- 原Plant类的
__init__中kk参数未被使用,可直接简化为def __init__(self, z),此时_tree_unflatten可改为return cls(*children),代码更简洁。 - 确认pytree分类逻辑:动态参数(children)是JAX会追踪、参与编译/微分的数组,静态参数(aux_data)是不被JAX处理的Python对象,你的分类逻辑是正确的。
内容的提问来源于stack exchange,提问作者user1168149
相关产品推荐
相关产品推荐

