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

JAX @jit嵌套类方法调用报错:参数重复TypeError问题求解

问题分析与解决

错误原因

  1. 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"冲突错误。

  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 03:26:15