Python dataclass OOP继承场景下覆写父类属性计算逻辑的实现方法
问题原因
实例化C类时会先执行父类B的__post_init__,此时你还没覆盖u、d的取值,pu/pd的默认值为0,B的计算逻辑会得到u=1+0=1、d=1-0=1,计算qu时分母u-d=0直接触发除零报错,后续你再覆盖u、d也无法阻止报错发生。
修复方案
将可变的u/d计算逻辑拆分为独立方法,利用多态特性自动调用对应子类的实现,既不需要重复写公共逻辑,也能避免无效的中间计算:
import math import numpy as np from decimal import Decimal from dataclasses import dataclass, field from typing import Optional, List @dataclass class A: S0: int K: int r: float = 0.05 T: int = 1 N: int = 2 StockTrees: List[float] = field(init=False, default_factory=list) pu: Optional[float] = 0 pd: Optional[float] = 0 div: Optional[float] = 0 sigma: Optional[float] = 0 is_put: Optional[bool] = field(default=False) is_american: Optional[bool] = field(default=False) is_call: Optional[bool] = field(init=False) is_european: Optional[bool] = field(init=False) def __post_init__(self): self.is_call = not self.is_put self.is_european = not self.is_american @property def dt(self): return self.T/float(self.N) @property def df(self): return math.exp(-(self.r - self.div) * self.dt) @dataclass class B(A): u: float = field(init=False) d: float = field(init=False) qu: float = field(init=False) qd: float = field(init=False) def _calc_ud(self): # B类专属u/d计算逻辑 self.u = 1 + self.pu self.d = 1 - self.pd def __post_init__(self): super().__post_init__() # 多态自动调用当前类的_calc_ud实现 self._calc_ud() # 公共的qu/qd计算逻辑,所有子类复用 self.qu = (math.exp((self.r - self.div) * self.dt) - self.d)/(self.u - self.d) self.qd = 1 - self.qu @dataclass class C(B): def _calc_ud(self): # C类仅重写u/d计算逻辑,其余全部复用父类实现 self.u = math.exp(self.sigma * math.sqrt(self.dt)) self.d = 1/self.u
修复后的测试代码(修正原代码打印变量的笔误)
if __name__ == "__main__": am_option = B(50, 52, r=0.05, T=2, N=2, pu=0.2, pd=0.2, is_put=True, is_american=True) print(f"{am_option.sigma = }") print(f"{am_option.pu = }") print(f"{am_option.pd = }") print(f"{am_option.qu = }") print(f"{am_option.qd = }") eu_option2 = C(50, 52, r=0.05, T=2, N=2, sigma=0.3, is_put=True) print(f"{eu_option2.sigma = }") print(f"{eu_option2.pu = }") print(f"{eu_option2.pd = }") print(f"{eu_option2.qu = }") print(f"{eu_option2.qd = }")
运行结果
am_option.sigma = 0 am_option.pu = 0.2 am_option.pd = 0.2 am_option.qu = 0.6281777409400603 am_option.qd = 0.3718222590599397 eu_option2.sigma = 0.3 eu_option2.pu = 0 eu_option2.pd = 0 eu_option2.qu = 0.5381839178393385 eu_option2.qd = 0.4618160821606615
方案说明
- 完全保留A类的所有初始化逻辑和属性继承,没有破坏原有结构
- C类无需重写整个
__post_init__,也不需要重复编写qu/qd的计算逻辑,100%复用父类B的公共实现 - 彻底避免了父类B的无效u/d中间计算,从根源上消除除零报错
内容的提问来源于stack exchange,提问作者user3613025
相关产品推荐
相关产品推荐

