如何优化同参数Python函数调用?消除传参冗余的最佳实践
消除函数调用与计算冗余的最佳实践
针对你遇到的两个函数重复传参、重复计算的问题,以下是几种实用的解决方案,按场景推荐:
1. 合并为单一函数,返回多结果(最简洁直接)
把两个函数的逻辑合并,一次调用同时返回「是否存在断点」和「断点位置」,从根源上消除重复传参和重复计算:
import numpy as np def checkBroken(time : np.ndarray, tol:float = 3): finiteDifference = np.diff(time) mask = finiteDifference >= tol has_broken = np.any(mask) if not has_broken: # 返回空数组表示无断点 return has_broken, np.array([]) brokenWhere = np.argwhere(mask) brokenStack = np.hstack((brokenWhere, brokenWhere+1)) return has_broken, brokenStack
调用时只需一次传参,直接获取两个结果:
Time = np.array([0,1,2,3,4,5,6,7,8,13,14,15,16,17,24,25,26,27]) has_broken, broken_locations = checkBroken(Time) if has_broken: # 处理断点位置 print(broken_locations)
2. 提取共享逻辑为辅助函数(保持单一职责)
如果想保留两个函数的独立职责,可把重复计算的np.diff和阈值判断抽成私有辅助函数,避免重复计算:
import numpy as np def _get_diff_mask(time : np.ndarray, tol:float = 3): # 私有函数,封装共享计算逻辑 finite_diff = np.diff(time) return finite_diff, finite_diff >= tol def isBroken(time : np.ndarray, tol:float = 3): _, mask = _get_diff_mask(time, tol) return np.any(mask) def brokenLocation(time : np.ndarray, tol:float = 3): _, mask = _get_diff_mask(time, tol) brokenWhere = np.argwhere(mask) return np.hstack((brokenWhere, brokenWhere+1))
调用时若需先判断再获取位置,可预计算一次共享结果,避免两次调用辅助函数:
Time = np.array([0,1,2,3,4,5,6,7,8,13,14,15,16,17,24,25,26,27]) _, mask = _get_diff_mask(Time) if np.any(mask): broken_locations = np.hstack((np.argwhere(mask), np.argwhere(mask)+1)) # 处理逻辑
3. 类封装状态(适合复用场景)
如果这个断点检测逻辑需要在多处复用,或后续要扩展相关功能,用类封装状态,初始化时传入参数,后续调用方法无需重复传参:
import numpy as np class BreakpointChecker: def __init__(self, time : np.ndarray, tol:float = 3): self.time = time self.tol = tol # 初始化时一次性完成所有共享计算 self._finite_diff = np.diff(time) self._mask = self._finite_diff >= tol self._has_broken = np.any(self._mask) def isBroken(self): return self._has_broken def brokenLocation(self): if not self._has_broken: return np.array([]) brokenWhere = np.argwhere(self._mask) return np.hstack((brokenWhere, brokenWhere+1))
调用实例化一次后,可多次调用方法:
Time = np.array([0,1,2,3,4,5,6,7,8,13,14,15,16,17,24,25,26,27]) checker = BreakpointChecker(Time) if checker.isBroken(): broken_locations = checker.brokenLocation() # 处理逻辑
内容的提问来源于stack exchange,提问作者The Mastermage
相关产品推荐
相关产品推荐

