From f87b3dc555e69a564b4a67d948ff5080e20985bd Mon Sep 17 00:00:00 2001 From: Kevin Johnson Date: Sun, 3 May 2020 09:13:55 -0700 Subject: [PATCH] BUG #39 ATR fix --- pandas_ta/volatility/atr.py | 5 ++++- tests/test_indicator_momentum.py | 2 +- tests/test_indicator_overlap.py | 2 +- tests/test_indicator_performance_ext.py | 4 +--- tests/test_indicator_statistics.py | 2 +- tests/test_indicator_trend.py | 2 +- tests/test_indicator_volatility.py | 2 +- tests/test_indicator_volume.py | 2 +- 8 files changed, 11 insertions(+), 10 deletions(-) diff --git a/pandas_ta/volatility/atr.py b/pandas_ta/volatility/atr.py index e56495a..c6474da 100644 --- a/pandas_ta/volatility/atr.py +++ b/pandas_ta/volatility/atr.py @@ -1,4 +1,6 @@ # -*- coding: utf-8 -*- +# from ..overlap.ema import ema +from ..overlap.rma import rma from .true_range import true_range from ..utils import get_drift, get_offset, verify_series @@ -17,7 +19,8 @@ def atr(high, low, close, length=None, mamode=None, drift=None, offset=None, **k # Calculate Result tr = true_range(high=high, low=low, close=close, drift=drift) if mamode == 'ema': - atr = tr.ewm(span=length, min_periods=min_periods).mean() + alpha = (1.0 / length) if length > 0 else 0.5 + atr = tr.ewm(alpha=alpha, min_periods=min_periods).mean() else: atr = tr.rolling(length, min_periods=min_periods).mean() diff --git a/tests/test_indicator_momentum.py b/tests/test_indicator_momentum.py index d0b2479..8b5c006 100644 --- a/tests/test_indicator_momentum.py +++ b/tests/test_indicator_momentum.py @@ -2,7 +2,7 @@ from .config import error_analysis, sample_data, CORRELATION, CORRELATION_THRESH from .context import pandas_ta from unittest import TestCase, skip -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import talib as tal diff --git a/tests/test_indicator_overlap.py b/tests/test_indicator_overlap.py index 3656e4f..0cd1691 100644 --- a/tests/test_indicator_overlap.py +++ b/tests/test_indicator_overlap.py @@ -2,7 +2,7 @@ from .config import CORRELATION, CORRELATION_THRESHOLD, error_analysis, sample_d from .context import pandas_ta from unittest import TestCase -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import pandas as pd diff --git a/tests/test_indicator_performance_ext.py b/tests/test_indicator_performance_ext.py index 4b9b4ac..8ea3e59 100644 --- a/tests/test_indicator_performance_ext.py +++ b/tests/test_indicator_performance_ext.py @@ -2,9 +2,7 @@ from .config import sample_data from .context import pandas_ta from unittest import TestCase -# import numpy.testing as npt -# import pandas.util.testing as pdt -from pandas import DataFrame#, Series +from pandas import DataFrame diff --git a/tests/test_indicator_statistics.py b/tests/test_indicator_statistics.py index dde4af3..c03efea 100644 --- a/tests/test_indicator_statistics.py +++ b/tests/test_indicator_statistics.py @@ -2,7 +2,7 @@ from .config import error_analysis, sample_data, CORRELATION, CORRELATION_THRESH from .context import pandas_ta from unittest import TestCase, skip -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import talib as tal diff --git a/tests/test_indicator_trend.py b/tests/test_indicator_trend.py index 664f053..43d3997 100644 --- a/tests/test_indicator_trend.py +++ b/tests/test_indicator_trend.py @@ -2,7 +2,7 @@ from .config import error_analysis, sample_data, CORRELATION, CORRELATION_THRESH from .context import pandas_ta from unittest import TestCase, skip -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import talib as tal diff --git a/tests/test_indicator_volatility.py b/tests/test_indicator_volatility.py index 3e6a1a6..3b38d1e 100644 --- a/tests/test_indicator_volatility.py +++ b/tests/test_indicator_volatility.py @@ -2,7 +2,7 @@ from .config import error_analysis, sample_data, CORRELATION, CORRELATION_THRESH from .context import pandas_ta from unittest import TestCase, skip -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import talib as tal diff --git a/tests/test_indicator_volume.py b/tests/test_indicator_volume.py index f22aad4..6ac1153 100644 --- a/tests/test_indicator_volume.py +++ b/tests/test_indicator_volume.py @@ -2,7 +2,7 @@ from .config import error_analysis, sample_data, CORRELATION, CORRELATION_THRESH from .context import pandas_ta from unittest import TestCase, skip -import pandas.util.testing as pdt +import pandas.testing as pdt from pandas import DataFrame, Series import talib as tal