如何将Wilder移动平均函数应用于多列np.ndarray?
解决方案
要让你的smooth_ma函数支持二维(2000,7)数组,核心思路是按列独立处理每个时间序列,利用np.apply_along_axis将原函数应用到数组的每一列上。以下是具体实现步骤:
1. 保留原一维处理逻辑
你的smooth_ma已经能完美处理(2000,)的一维数组,无需修改这个核心逻辑,我们只需要把它复用在二维数组的每一列上。
2. 封装适配二维数组的函数
编写一个包装函数,自动识别输入数组维度,并针对二维数组做列维度的批量处理,同时兼容period参数的两种类型(单个整数/整数列表):
import numpy as np from typing import Union # 保留你原有的smooth_ma函数不变 def smooth_ma( price: np.ndarray, period: Union[int, List[int]], ffill: bool = False, ) -> np.ndarray: """Compute Wilder's Moving Average indicator.""" arr, row, col = utils._push_nan(arr=price) tmp = np.full_like(arr, np.nan) number_nan_in_arr = np.count_nonzero(np.isnan(arr)) first_moving_avg_value = ( np.sum(arr[number_nan_in_arr : period + number_nan_in_arr]) / period ) tmp[number_nan_in_arr + period] = first_moving_avg_value data_series = arr[number_nan_in_arr + period :].T period_factor = (period - 1) / period data_series_factor = period_factor ** np.arange(data_series.shape[1] - 1, -1, -1) wilders_moving_average = np.cumsum( data_series_factor * data_series ) / data_series_factor / period + first_moving_avg_value * ( period_factor ** np.arange(1, data_series.shape[1] + 1) ) wilders_moving_average_transpose = np.array([wilders_moving_average]).T out = np.vstack( ( tmp[: number_nan_in_arr + period - 1], first_moving_avg_value, wilders_moving_average_transpose, ) ) out = utils._pull_nan(arr=out, row=row, col=col) if ffill: return utils._ffill_arr(arr=out) else: return out # 新增适配二维数组的函数 def smooth_ma_multi_col( price: np.ndarray, period: Union[int, List[int]], ffill: bool = False, ) -> np.ndarray: if price.ndim == 1: return smooth_ma(price, period, ffill) elif price.ndim != 2: raise ValueError("仅支持一维或二维数组输入") # 处理单个周期的情况 if isinstance(period, int): # 沿axis=0处理每一列,返回形状为(2000,7)的结果 return np.apply_along_axis( lambda col: smooth_ma(col, period, ffill), axis=0, arr=price ) # 处理多周期的情况 elif isinstance(period, list): results = [] for p in period: col_result = np.apply_along_axis( lambda col: smooth_ma(col, p, ffill), axis=0, arr=price ) results.append(col_result) # 将多周期结果堆叠为三维数组(2000,7,len(period)) # 若需合并周期到列维度,可改用np.concatenate(results, axis=1),形状变为(2000,7*len(period)) return np.stack(results, axis=-1) else: raise TypeError("period参数仅支持int或List[int]类型")
3. 调用示例
# 一维数组调用 arr_1d = np.random.randn(2000) result_1d = smooth_ma_multi_col(arr_1d, period=14, ffill=True) # 二维数组调用 arr_2d = np.random.randn(2000,7) # 单个周期 result_2d_single = smooth_ma_multi_col(arr_2d, period=14, ffill=True) # 多个周期 result_2d_multi = smooth_ma_multi_col(arr_2d, period=[7,14,21], ffill=True)
关键说明
np.apply_along_axis(func, axis=0, arr)会将二维数组的每一列作为独立的一维序列传入smooth_ma,确保每列的NaN处理、MA计算都是独立进行的,解决了原函数全局处理NaN的问题。- 多周期场景下,默认将结果堆叠为三维数组,你可以根据业务需求调整维度格式。
- 确保你的
utils._push_nan、utils._pull_nan、utils._ffill_arr工具函数能兼容一维数组输入——如果它们原本只支持二维,只需添加简单的维度判断逻辑即可适配。
内容的提问来源于stack exchange,提问作者SebastianHeeg
相关产品推荐
相关产品推荐

