Python分段常数时间序列类的运算符重载实现优化问询
Python分段常数函数类(PiecewiseConstant)的运算符优化方案
一、优雅判断数值类型
替代if type(other) == int or type(other) == float:的生硬写法,使用numbers.Number抽象基类,能覆盖所有数值类型(包括int、float、numpy数值类型等),避免遗漏合法类型:
import numbers def is_scalar(other): return isinstance(other, numbers.Number)
二、通用运算符逻辑复用
不要用eval这种不安全的方式,直接传入operator模块的运算符函数,编写通用二元运算处理函数,复用核心逻辑:
- 常数运算:直接对原
values执行对应运算,保留原时间点 - 同类型实例运算:合并两个序列的时间点并集,逐点计算运算结果
三、numpy数组与常数运算语法
numpy支持广播机制,直接用常规运算符即可完成常数与数组的运算,示例:
# 常数乘numpy数组 new_values = self.values * scalar # 常数加numpy数组 new_values = self.values + scalar
四、完整优化后的代码实现
from bisect import bisect_left import numpy as np import numbers import operator class PiecewiseConstant: def __init__(self, times: np.array, values: np.array): assert len(times) == len(values), "Times and values must have the same length" assert len(times) > 0, "Time series cannot be empty" # 用numpy高效检查时间序列严格递增 assert np.all(np.diff(times) > 0), "Time series must be strictly increasing" self.times = times self.values = values def __getitem__(self, time): if time < self.times[0]: return self.values[0] time_index = bisect_left(self.times, time) if time_index == len(self.times) or time < self.times[time_index]: return self.values[time_index - 1] return self.values[time_index] def __setitem__(self, time): raise NotImplementedError("Setting values via indexing is not allowed") def _binary_operation(self, other, op_func): # 处理常数运算 if isinstance(other, numbers.Number): new_values = op_func(self.values, other) return PiecewiseConstant(self.times.copy(), new_values) # 处理同类型实例运算 elif isinstance(other, PiecewiseConstant): # 合并时间点并集,自动去重排序 combined_times = np.union1d(self.times, other.times) combined_values = [] for t in combined_times: val1 = self[t] val2 = other[t] combined_values.append(op_func(val1, val2)) return PiecewiseConstant(combined_times, np.array(combined_values)) else: raise TypeError(f"Unsupported operand type(s): {type(self)} and {type(other)}") def __add__(self, other): return self._binary_operation(other, operator.add) def __mul__(self, other): return self._binary_operation(other, operator.mul) def __sub__(self, other): return self._binary_operation(other, operator.sub) def __truediv__(self, other): return self._binary_operation(other, operator.truediv) # 支持反向运算(比如常数 + 实例) def __radd__(self, other): return self.__add__(other) def __rmul__(self, other): return self.__mul__(other) def __rsub__(self, other): return PiecewiseConstant(self.times.copy(), operator.sub(other, self.values)) def __rtruediv__(self, other): return PiecewiseConstant(self.times.copy(), operator.truediv(other, self.values))
五、原代码的问题修复点
- 初始化检查优化:用
np.diff替代循环判断时间序列递增,更高效简洁 - 移除不安全的eval:直接传入运算符函数,提升代码安全性与可读性
- 合并时间点简化:用
np.union1d直接获取时间点并集,替代手动迭代合并,减少代码复杂度 - 支持反向运算:实现
__radd__、__rmul__等方法,支持常数 + 实例这类反向操作 - 类型判断严谨:用
numbers.Number覆盖所有合法数值类型,避免遗漏
内容的提问来源于stack exchange,提问作者Lost1
相关产品推荐
相关产品推荐

