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

Equinox中通过__new__方法根据可选标志返回不同类实例的异常问题咨询

Equinox中通过__new__方法根据可选标志返回不同类实例的异常问题咨询

我正在使用Equinox实现一系列类,目的是支持对类参数求导。大多数情况下,用户会实例化A类并使用fn函数生成数据(具体细节不重要)。但在需要计算梯度的场景下,用sigmoid函数来表示param_c能确保它被限制在(0,1)范围内,同时我不希望用户感知到类行为的差异。

因此,我实现了另一个类A_sigmoid,将param_c作为property,并通过A_abstract抽象类确保两个类都继承fn方法(该方法的逻辑会调用param_c)。我不想让用户必须区分A和A_sigmoid,而是希望他们无论使用哪个类都传入相同的kwargs字典,转换在后台完成。另外,我还希望实例化A时可以通过一个可选标志来指定是否使用sigmoid版本。

基于此,我实现了如下最小可复现示例(MWE):

class A_abstract(eqx.Module):
    param_a: jax.Array
    param_b: jax.Array
    param_c: eqx.AbstractVar[jax.Array]
    
    def fn(self,*args,**kwargs):
        pass

class A_sigmoid(A_abstract):
    _param_c_sigmoid: jax.Array

    @property
    def param_c(self):
        return 1 / (1 + jnp.exp(-self._param_c_sigmoid))

class A(A_abstract):
    param_c: jax.Array

    def __new__(cls, **kwargs):
        sigmoid_flag = kwargs.pop('use_sigmoid_c',False)
        if sigmoid_flag == True:
            param_c = kwargs.pop('param_c')
            _param_c_sigmoid = jnp.log(param_c / (1 - param_c))
            kwargs['_param_c_sigmoid'] = _param_c_sigmoid
            instance = A_sigmoid.__new__(A_sigmoid)
            instance.__init__(**kwargs)
            print(type(instance))
            return instance
        else:
            return super(A,cls).__new__(cls)

classA = A(param_a = 1.,param_b = 2.,param_c = 0.5,use_sigmoid_c=True)
print(type(classA))

运行这段代码时,__new__方法内打印instance的类型是A_sigmoid,但在外部打印type(classA)时,结果却是A,而且该实例没有param_c属性,反而存在_param_c_sigmoid属性。

这是为什么?我在使用__new__时是不是忽略了什么细节?

我知道原则上用工厂模式是最好的方案,但还有B、C等其他类不需要sigmoid实现,我希望它们和A的实例化方式完全一致,方便替换,因此不想为A单独写一个自定义实例化方法。

我当前的运行环境是Jupyter Notebook,依赖包版本如下:

Python           : 3.12.4
IPython          : 8.30.0
ipykernel        : 6.29.5
jupyter_client   : 8.6.3
jupyter_core     : 5.7.2

备注:内容来源于stack exchange,提问作者JamesVR

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 13:59:29