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
相关产品推荐
相关产品推荐

