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

如何仅使用Numpy实现平均真实波动范围(ATR)的最快计算方法?

Fastest NumPy-Only Implementation of Average True Range (ATR)

Here's an optimized pure NumPy implementation of ATR that mirrors the logic of your Pandas-NumPy reference code, while maximizing speed through vectorized operations and efficient cumulative sum calculations.

Key Optimizations

  • Uses vectorized NumPy operations (no Python loops) to leverage optimized C backend execution.
  • Computes rolling averages with np.cumsum for O(n) time complexity, which is far faster than iterative rolling window calculations.
  • Handles edge cases like the first period (where no prior close exists) and optional NaN handling.

Implementation for Clean Data (No Missing Values)

If your high, low, and close arrays have no NaN values, this is the fastest version:

import numpy as np

def calculate_atr(high: np.ndarray, low: np.ndarray, close: np.ndarray, window: int = 14) -> np.ndarray:
    # Calculate the three range components
    high_low = high - low
    # Shift close prices (first element becomes invalid, so we set it to NaN)
    shifted_close = np.roll(close, 1)
    shifted_close[0] = np.nan
    high_close = np.abs(high - shifted_close)
    low_close = np.abs(low - shifted_close)
    
    # Compute True Range: max of the three ranges (ignoring NaN in first period)
    ranges = np.stack([high_low, high_close, low_close], axis=1)
    true_range = np.nanmax(ranges, axis=1)
    
    # Calculate rolling 14-period average using cumulative sum
    cumsum_tr = np.cumsum(true_range)
    # Compute rolling sum for each window
    rolling_sum = cumsum_tr[window:] - cumsum_tr[:-window]
    # Calculate ATR values
    atr_values = rolling_sum / window
    # Pad the first (window-1) elements with NaN
    atr = np.concatenate([np.full(window - 1, np.nan), atr_values])
    
    return atr

Implementation with NaN Handling

If your data may contain missing values, use this version which properly ignores NaNs in calculations:

import numpy as np

def calculate_atr_with_nans(high: np.ndarray, low: np.ndarray, close: np.ndarray, window: int = 14) -> np.ndarray:
    # Calculate range components (NaN propagates where input is NaN)
    high_low = high - low
    shifted_close = np.roll(close, 1)
    shifted_close[0] = np.nan
    high_close = np.abs(high - shifted_close)
    low_close = np.abs(low - shifted_close)
    
    # Compute True Range, ignoring NaNs in each row
    ranges = np.stack([high_low, high_close, low_close], axis=1)
    true_range = np.nanmax(ranges, axis=1)
    
    # Cumulative sum of True Range (ignoring NaNs)
    cumsum_tr = np.nancumsum(true_range)
    # Cumulative count of non-NaN values
    cumsum_count = np.cumsum(~np.isnan(true_range))
    
    # Calculate rolling sum and valid count for each window
    rolling_sum = cumsum_tr[window:] - cumsum_tr[:-window]
    rolling_count = cumsum_count[window:] - cumsum_count[:-window]
    
    # ATR is sum divided by window size (matches original Pandas logic)
    atr_values = rolling_sum / window
    # Set to NaN if fewer than window valid values in the window
    atr_values[rolling_count < window] = np.nan
    
    # Pad initial elements with NaN
    atr = np.concatenate([np.full(window - 1, np.nan), atr_values])
    
    return atr

How It Works

  1. Range Components: We compute the three core ranges used in ATR:
    • high_low: Difference between daily high and low.
    • high_close: Absolute difference between daily high and previous close.
    • low_close: Absolute difference between daily low and previous close.
  2. True Range: For each period, take the maximum of the three ranges (using np.nanmax to ignore NaNs in the first period where no prior close exists).
  3. Rolling Average: Using cumulative sums, we efficiently compute the sum of each 14-period window in linear time, then divide by the window size to get the ATR.

Performance Notes

This implementation is significantly faster than equivalent Pandas code for large datasets because it avoids the overhead of DataFrame/Series objects and uses pure NumPy vectorization. For datasets with millions of rows, this can lead to a 2-5x speedup (depending on hardware and data size).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 02:57:49