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

使用attrs库创建带互斥参数的Python类

解决互斥参数类的实例化问题

针对你遇到的这个问题,我有两种可行的解决方案,先从最彻底的方案说起,再讲保留灵活性的替代方式:

1. 完全禁用常规构造函数(推荐)

这是从根源上解决问题的办法,确保用户只能通过你定义的from_prices和from_returns方法来创建实例,彻底杜绝直接传入不兼容参数的可能。

具体实现步骤:

  • 用@attr.s(init=False)告诉attrs不要自动生成构造函数;
  • 自定义一个__init__方法并抛出错误,阻止外部直接调用;
  • 在备选构造函数里用cls.__new__(cls)创建实例(绕过__init__),然后手动给所有属性赋值。

修改后的代码如下:

import pandas as pd
import attr

@attr.s(init=False)  # 关闭attrs自动生成构造函数
class MutuallyExclusive:
    prices: pd.Series = attr.ib(init=False)
    returns: pd.Series = attr.ib(init=False)
    trading_days_per_year: int = attr.ib(init=False, default=252)

    def __init__(self):
        # 直接调用构造函数就报错,引导用户用正确的方法
        raise NotImplementedError("请使用from_prices或from_returns方法实例化对象")

    @classmethod
    def from_prices(cls, price_series: pd.Series, trading_days: int = 252):
        instance = cls.__new__(cls)  # 创建实例但不触发__init__
        instance.prices = price_series
        instance.returns = price_series.pct_change()
        instance.trading_days_per_year = trading_days
        return instance

    @classmethod
    def from_returns(cls, return_series: pd.Series, trading_days: int = 252):
        instance = cls.__new__(cls)
        instance.returns = return_series
        # 从基准值100开始计算对应的价格序列
        instance.prices = pd.Series(100 * (1 + return_series).cumprod(), index=return_series.index)
        instance.trading_days_per_year = trading_days
        return instance

if __name__ == "__main__":
    prices = pd.Series(data=[100, 101, 98, 104, 102, 108])
    returns = pd.Series(data=[0.01, 0.03, -0.02, 0.01, -0.03, 0.04])
    
    # 正常创建实例
    obj_returns = MutuallyExclusive.from_returns(returns)
    obj_prices = MutuallyExclusive.from_prices(prices, trading_days=100)
    
    # 尝试直接调用常规构造函数会触发错误
    try:
        obj = MutuallyExclusive(prices, returns)
    except NotImplementedError as e:
        print(e)  # 输出:请使用from_prices或from_returns方法实例化对象

2. 保留常规构造函数但添加验证逻辑

如果你需要保留常规构造函数的灵活性(比如允许用户只传其中一个参数),可以在类初始化后添加验证逻辑,检查传入的prices和returns是否兼容,不兼容就抛出错误。

实现思路:

  • 把prices和returns的默认值设为None,允许只传其中一个;
  • 在__attrs_post_init__方法中做三件事:检查是否同时传入了不兼容的参数、检查是否一个参数都没传、如果只传了一个则自动计算另一个。

代码示例:

import pandas as pd
import attr
from attrs import validators

@attr.s
class MutuallyExclusive:
    prices: pd.Series = attr.ib(default=None, validator=validators.optional(validators.instance_of(pd.Series)))
    returns: pd.Series = attr.ib(default=None, validator=validators.optional(validators.instance_of(pd.Series)))
    trading_days_per_year: int = attr.ib(default=252)

    @classmethod
    def from_prices(cls, price_series: pd.Series, trading_days: int = 252):
        return cls(price_series, price_series.pct_change(), trading_days)

    @classmethod
    def from_returns(cls, return_series: pd.Series, trading_days: int = 252):
        prices = pd.Series(100 * (1 + return_series).cumprod(), index=return_series.index)
        return cls(prices, return_series, trading_days)

    def __attrs_post_init__(self):
        # 检查是否同时传入了两个参数
        if self.prices is not None and self.returns is not None:
            # 计算prices对应的returns,忽略第一个NaN值,对比精度到小数点后8位
            calculated_returns = self.prices.pct_change()
            is_match = calculated_returns.iloc[1:].round(8).equals(self.returns.iloc[1:].round(8))
            if not is_match:
                raise ValueError("传入的prices和returns不兼容,returns应为prices的pct_change结果")
        # 检查是否一个参数都没传
        elif self.prices is None and self.returns is None:
            raise ValueError("必须传入prices或returns中的一个参数")
        # 只传了prices,自动计算returns
        elif self.prices is not None:
            self.returns = self.prices.pct_change()
        # 只传了returns,自动计算prices
        elif self.returns is not None:
            self.prices = pd.Series(100 * (1 + self.returns).cumprod(), index=self.returns.index)

if __name__ == "__main__":
    prices = pd.Series(data=[100, 101, 98, 104, 102, 108])
    returns = pd.Series(data=[0.01, 0.03, -0.02, 0.01, -0.03, 0.04])
    
    # 正常实例化
    obj_returns = MutuallyExclusive.from_returns(returns)
    obj_prices = MutuallyExclusive.from_prices(prices, trading_days=100)
    
    # 传入不兼容的参数会报错
    try:
        obj = MutuallyExclusive(prices, returns)
    except ValueError as e:
        print(e)  # 输出:传入的prices和returns不兼容...
    
    # 只传一个参数也能自动补全
    obj_only_prices = MutuallyExclusive(prices=prices)
    obj_only_returns = MutuallyExclusive(returns=returns)

方案选择建议

  • 如果你想完全避免用户误用常规构造函数,优先选第一种方案,彻底封死错误路径;
  • 如果你需要给用户更多灵活性(比如允许只传一个参数直接实例化),第二种方案更合适,既能验证参数兼容性,又能自动补全缺失值。

内容的提问来源于stack exchange,提问作者Andi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 16:07:30