diff --git a/ta/volatility.py b/ta/volatility.py index 3bbf57e..ec5a562 100644 --- a/ta/volatility.py +++ b/ta/volatility.py @@ -47,7 +47,8 @@ def _run(self): close_shift = self._close.shift(1) true_range = self._true_range(self._high, self._low, close_shift) atr = np.zeros(len(self._close)) - atr[self._window - 1] = true_range[0 : self._window].mean() + if len(atr) >= self._window: + atr[self._window - 1] = true_range[0 : self._window].mean() for i in range(self._window, len(atr)): atr[i] = (atr[i - 1] * (self._window - 1) + true_range.iloc[i]) / float( self._window diff --git a/test/test_atr_short_input.py b/test/test_atr_short_input.py new file mode 100644 index 0000000..17bd9a3 --- /dev/null +++ b/test/test_atr_short_input.py @@ -0,0 +1,43 @@ +import unittest + +import pandas as pd + +from ta.volatility import AverageTrueRange, average_true_range + + +class TestATRShortInput(unittest.TestCase): + def test_short_inputs(self): + for size in (0, 1, 5): + for fillna in (False, True): + with self.subTest(size=size, fillna=fillna): + index = pd.date_range("2021-01-01", periods=size) + close = pd.Series(range(2, size + 2), index=index, dtype=float) + kwargs = dict( + high=close + 1, + low=close - 1, + close=close, + window=6, + fillna=fillna, + ) + expected = pd.Series(0.0, index=index, name="atr") + pd.testing.assert_series_equal( + average_true_range(**kwargs), expected + ) + pd.testing.assert_series_equal( + AverageTrueRange(**kwargs).average_true_range(), expected + ) + + def test_complete_window(self): + for size in (6, 7): + with self.subTest(size=size): + index = pd.date_range("2021-01-01", periods=size) + close = pd.Series(range(2, size + 2), index=index, dtype=float) + expected = pd.Series( + [0.0] * 5 + [2.0] * (size - 5), index=index, name="atr" + ) + result = average_true_range(close + 1, close - 1, close, window=6) + pd.testing.assert_series_equal(result, expected) + + +if __name__ == "__main__": + unittest.main()