From 206b27fe7244870176c2ef11f207d396f55414e4 Mon Sep 17 00:00:00 2001 From: Ubuntu Date: Thu, 27 Aug 2026 02:01:43 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E4=BB=A5=E8=87=AA=E5=AE=9E?= =?UTF-8?q?=E7=8E=B0=E6=8C=87=E6=A0=87=E6=9B=BF=E6=8D=A2=20talib=20?= =?UTF-8?q?=E4=B8=8E=20technical=20=E4=BE=9D=E8=B5=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit chanlun/indicators/ta.py 接口兼容 talib.abstract,实现代码实际用到的 SMA/MA/EMA/RSI/ATR/MACD/BBANDS;chanlun/pipeline/resample.py 替代 technical.util.resample_to_interval。调用点只改 import,逻辑未动。 暖机长度与平滑种子按 TA-Lib 的约定实现,差一根 K 线就会让下游所有 笔/线段/中枢整体位移。其中 MACD 需特别处理:TA-Lib 让快慢两条 EMA 在同一根 K 线出首值,因而快线的种子取 x[slow-fast:slow] 的均值,而非 从 fastperiod-1 一路递推——两者在百元价位上相差约 0.17。 BBANDS 是有意的分歧:TA-Lib 用 sumsq/n - mean² 求方差,短窗口远离零 时灾难性抵消(timeperiod=2 误差 8.7e-7),本实现用 rolling std,对 50 位精度基准误差为 0。项目实际使用的周期两者一致到 1e-10。 顺带清理 12 个文件中 16 处从未调用的 talib/technical 导入。 验证:9440 组随机对拨;真实 K 线端到端比对 add_indicators 全部 33 个 指标列,NaN 模式一致、MACD 柱符号 100% 相同;屏蔽两个包后 60 个模块 均可导入。新增 test_ta_compat.py 将输出逐 bar 钉在 TA-Lib 上,但该文件 在 TA-Lib 缺失时静默跳过,改动 ta.py 需在装有 TA-Lib 的环境复跑。 Co-authored-by: Cursor --- chanlun/analysis/ChanLun_Classifier.py | 2 - chanlun/analysis/Find_Trend.py | 7 +- chanlun/indicators/ta.py | 196 ++++++++++++++++++++++ chanlun/pipeline/builders/bi.py | 2 - chanlun/pipeline/builders/bsp.py | 2 - chanlun/pipeline/builders/incremental.py | 2 +- chanlun/pipeline/builders/indicators.py | 3 +- chanlun/pipeline/builders/kline.py | 2 - chanlun/pipeline/builders/seg.py | 2 - chanlun/pipeline/builders/zs.py | 2 - chanlun/pipeline/orchestrator.py | 2 - chanlun/pipeline/resample.py | 52 ++++++ chanlun/pipeline/timeframe.py | 3 +- chanlun/tests/test_ta_compat.py | 198 +++++++++++++++++++++++ web/services/cn_stock.py | 2 +- web/services/runtime/analyze.py | 2 +- web/services/runtime/indicators.py | 2 +- 17 files changed, 456 insertions(+), 25 deletions(-) create mode 100644 chanlun/indicators/ta.py create mode 100644 chanlun/pipeline/resample.py create mode 100644 chanlun/tests/test_ta_compat.py diff --git a/chanlun/analysis/ChanLun_Classifier.py b/chanlun/analysis/ChanLun_Classifier.py index f4a40a3..affe93c 100644 --- a/chanlun/analysis/ChanLun_Classifier.py +++ b/chanlun/analysis/ChanLun_Classifier.py @@ -14,12 +14,10 @@ from chanlun.core.ChanSBI import ChanSBI from chanlun.core.ChanSEG import ChanSEG from chanlun.core.ChanZS import ChanZS from chanlun.core.ChanBSP import ChanBSP -import talib.abstract as ta import pandas as pd import matplotlib.pyplot as plt from matplotlib.dates import DateFormatter, date2num import matplotlib.patches as patches -from technical.util import resample_to_interval from decimal import Decimal from chanlun.pipeline.orchestrator import ChanLun import xgboost as xgb diff --git a/chanlun/analysis/Find_Trend.py b/chanlun/analysis/Find_Trend.py index 8d392a3..56956c2 100644 --- a/chanlun/analysis/Find_Trend.py +++ b/chanlun/analysis/Find_Trend.py @@ -2,7 +2,7 @@ import ccxt import pandas as pd import numpy as np import mplfinance as mpf -from talib import MACD, SMA +from chanlun.indicators import ta from datetime import datetime, timedelta import logging import datetime as dt @@ -249,8 +249,9 @@ def analyze_higher_timeframe(df_30m): # 8. Back-divergence detection (enhanced) def detect_back_divergence(df, strokes, higher_trend): try: - macd, signal, hist = MACD(df['Close'], fastperiod=12, slowperiod=26, signalperiod=9) - sma20 = SMA(df['Close'], timeperiod=20) + macd_df = ta.MACD(df['Close'], fastperiod=12, slowperiod=26, signalperiod=9) + macd, hist = macd_df['macd'], macd_df['macdhist'] + sma20 = ta.SMA(df['Close'], timeperiod=20) df['macd'] = macd df['hist'] = hist df['sma20'] = sma20 diff --git a/chanlun/indicators/ta.py b/chanlun/indicators/ta.py new file mode 100644 index 0000000..2aa28ce --- /dev/null +++ b/chanlun/indicators/ta.py @@ -0,0 +1,196 @@ +"""Drop-in replacement for the `talib.abstract` calls this project makes. + +Same call signatures, same column names, same NaN warm-up lengths, so call +sites only change their import line. + +Only what the codebase actually uses is implemented: SMA, MA, EMA, RSI, ATR, +MACD, BBANDS. Numerical agreement with TA-Lib is enforced by +`chanlun/tests/test_ta_compat.py`, which skips when talib is absent. + +The warm-up conventions below are TA-Lib's, not the textbook ones, and they +differ between functions — getting them wrong shifts every downstream Chan +structure by a bar: + + SMA/BBANDS first value at index period-1 + EMA seeded with the SMA of the first `period` values, at index period-1 + RSI/ATR Wilder smoothing (alpha = 1/period), first value at index period +""" + +from __future__ import annotations + +import numpy as np +import pandas as pd + +__all__ = ["SMA", "MA", "EMA", "RSI", "ATR", "MACD", "BBANDS"] + + +def _series(data, price: str = "close") -> pd.Series: + """Accept the abstract-API shapes: DataFrame, Series, or ndarray.""" + if isinstance(data, pd.DataFrame): + return data[price].astype(float) + if isinstance(data, pd.Series): + return data.astype(float) + return pd.Series(np.asarray(data, dtype=float)) + + +def _recursive(values: np.ndarray, seed: float, start: int, alpha: float, n: int) -> np.ndarray: + """out[start] = seed; out[i] = alpha*values[i] + (1-alpha)*out[i-1]. + + Delegates the recursion to pandas' C implementation rather than a Python + loop — `research/` runs this over long histories. + """ + out = np.full(n, np.nan) + if start >= n: + return out + tail = values[start:].astype(float).copy() + tail[0] = seed + out[start:] = pd.Series(tail).ewm(alpha=alpha, adjust=False).mean().to_numpy() + return out + + +def SMA(data, timeperiod: int = 30, price: str = "close") -> pd.Series: + s = _series(data, price) + return s.rolling(window=timeperiod, min_periods=timeperiod).mean() + + +def MA(data, timeperiod: int = 30, matype: int = 0, price: str = "close") -> pd.Series: + if matype != 0: + raise NotImplementedError(f"MA matype={matype} is not used by this codebase") + return SMA(data, timeperiod, price=price) + + +def _ema(x: np.ndarray, period: int, start: int) -> np.ndarray: + """EMA whose first output lands on `start`, seeded by the SMA of the + `period` values ending there. + + `start` is a parameter because MACD needs the fast EMA to begin later than + it naturally would; see the note in MACD(). + """ + n = x.size + if n <= start or start < period - 1: + return np.full(n, np.nan) + seed = x[start - period + 1: start + 1].mean() + return _recursive(x, seed, start, 2.0 / (period + 1.0), n) + + +def EMA(data, timeperiod: int = 30, price: str = "close") -> pd.Series: + s = _series(data, price) + x = s.to_numpy(dtype=float) + return pd.Series(_ema(x, timeperiod, timeperiod - 1), index=s.index) + + +def RSI(data, timeperiod: int = 14, price: str = "close") -> pd.Series: + s = _series(data, price) + x = s.to_numpy(dtype=float) + n = x.size + out = np.full(n, np.nan) + if n <= timeperiod: + return pd.Series(out, index=s.index) + + delta = np.diff(x) + gain = np.where(delta > 0.0, delta, 0.0) + loss = np.where(delta < 0.0, -delta, 0.0) + + # delta[k] corresponds to bar k+1, so the first `timeperiod` deltas seed bar `timeperiod`. + alpha = 1.0 / timeperiod + avg_gain = _recursive(gain, gain[:timeperiod].mean(), timeperiod - 1, alpha, n - 1) + avg_loss = _recursive(loss, loss[:timeperiod].mean(), timeperiod - 1, alpha, n - 1) + + ag = avg_gain[timeperiod - 1:] + al = avg_loss[timeperiod - 1:] + with np.errstate(divide="ignore", invalid="ignore"): + rsi = np.where(al == 0.0, 100.0, 100.0 - 100.0 / (1.0 + ag / al)) + out[timeperiod:] = rsi + return pd.Series(out, index=s.index) + + +def ATR(data, timeperiod: int = 14) -> pd.Series: + if not isinstance(data, pd.DataFrame): + raise TypeError("ATR needs a DataFrame with high/low/close") + high = data["high"].to_numpy(dtype=float) + low = data["low"].to_numpy(dtype=float) + close = data["close"].to_numpy(dtype=float) + n = high.size + out = np.full(n, np.nan) + if n <= timeperiod: + return pd.Series(out, index=data.index) + + prev_close = close[:-1] + tr = np.maximum.reduce([ + high[1:] - low[1:], + np.abs(high[1:] - prev_close), + np.abs(low[1:] - prev_close), + ]) + + # tr[k] is bar k+1; the first `timeperiod` true ranges seed bar `timeperiod`. + smoothed = _recursive(tr, tr[:timeperiod].mean(), timeperiod - 1, 1.0 / timeperiod, n - 1) + out[timeperiod:] = smoothed[timeperiod - 1:] + return pd.Series(out, index=data.index) + + +def MACD( + data, + fastperiod: int = 12, + slowperiod: int = 26, + signalperiod: int = 9, + price: str = "close", +) -> pd.DataFrame: + if slowperiod < fastperiod: + fastperiod, slowperiod = slowperiod, fastperiod + + s = _series(data, price) + x = s.to_numpy(dtype=float) + n = x.size + + macd = np.full(n, np.nan) + signal = np.full(n, np.nan) + hist = np.full(n, np.nan) + empty = pd.DataFrame({"macd": macd, "macdsignal": signal, "macdhist": hist}, index=s.index) + + # Both EMAs emit their first value on the same bar. That makes the slow one + # ordinary, but re-seeds the fast one from the SMA of the `fastperiod` + # values ending there instead of carrying the recursion forward from bar + # fastperiod-1 — the two disagree by ~0.2 on a 100-priced series. + macd_start = slowperiod - 1 + if n <= macd_start: + return empty + line = _ema(x, fastperiod, macd_start) - _ema(x, slowperiod, macd_start) + + # The signal EMA runs over the MACD line, so everything shifts by another + # signalperiod-1 bars, and TA-Lib trims the MACD line to match. + valid = line[macd_start:] + if valid.size < signalperiod: + return empty + sig = _recursive( + valid, valid[:signalperiod].mean(), signalperiod - 1, 2.0 / (signalperiod + 1.0), valid.size + ) + start = macd_start + signalperiod - 1 + macd[start:] = line[start:] + signal[macd_start:] = sig + hist = macd - signal + + return pd.DataFrame({"macd": macd, "macdsignal": signal, "macdhist": hist}, index=s.index) + + +def BBANDS( + data, + timeperiod: int = 5, + nbdevup: float = 2.0, + nbdevdn: float = 2.0, + matype: int = 0, + price: str = "close", +) -> pd.DataFrame: + if matype != 0: + raise NotImplementedError(f"BBANDS matype={matype} is not used by this codebase") + s = _series(data, price) + middle = s.rolling(window=timeperiod, min_periods=timeperiod).mean() + # TA-Lib uses the population standard deviation. + std = s.rolling(window=timeperiod, min_periods=timeperiod).std(ddof=0) + return pd.DataFrame( + { + "upperband": middle + nbdevup * std, + "middleband": middle, + "lowerband": middle - nbdevdn * std, + }, + index=s.index, + ) diff --git a/chanlun/pipeline/builders/bi.py b/chanlun/pipeline/builders/bi.py index 53600ac..ba82710 100644 --- a/chanlun/pipeline/builders/bi.py +++ b/chanlun/pipeline/builders/bi.py @@ -6,9 +6,7 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/builders/bsp.py b/chanlun/pipeline/builders/bsp.py index dcf1183..61145e8 100644 --- a/chanlun/pipeline/builders/bsp.py +++ b/chanlun/pipeline/builders/bsp.py @@ -6,9 +6,7 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/builders/incremental.py b/chanlun/pipeline/builders/incremental.py index 42e99fe..1080b24 100644 --- a/chanlun/pipeline/builders/incremental.py +++ b/chanlun/pipeline/builders/incremental.py @@ -9,7 +9,7 @@ from datetime import datetime import pandas as pd from pandas import DataFrame -from technical.util import resample_to_interval +from chanlun.pipeline.resample import resample_to_interval from chanlun.core.ChanEnum import Chan_FX_TYPE, Chan_KLC_FX, Chan_KLC_STATE from chanlun.core.ChanKLU import ChanKLU diff --git a/chanlun/pipeline/builders/indicators.py b/chanlun/pipeline/builders/indicators.py index 900175e..610af48 100644 --- a/chanlun/pipeline/builders/indicators.py +++ b/chanlun/pipeline/builders/indicators.py @@ -6,9 +6,8 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta +from chanlun.indicators import ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/builders/kline.py b/chanlun/pipeline/builders/kline.py index 2f887a2..807ff2f 100644 --- a/chanlun/pipeline/builders/kline.py +++ b/chanlun/pipeline/builders/kline.py @@ -6,9 +6,7 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/builders/seg.py b/chanlun/pipeline/builders/seg.py index b1b700d..1a42d8a 100644 --- a/chanlun/pipeline/builders/seg.py +++ b/chanlun/pipeline/builders/seg.py @@ -6,9 +6,7 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/builders/zs.py b/chanlun/pipeline/builders/zs.py index 084b83d..d2aca1a 100644 --- a/chanlun/pipeline/builders/zs.py +++ b/chanlun/pipeline/builders/zs.py @@ -6,9 +6,7 @@ from decimal import Decimal import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/pipeline/orchestrator.py b/chanlun/pipeline/orchestrator.py index ce32177..788ef14 100644 --- a/chanlun/pipeline/orchestrator.py +++ b/chanlun/pipeline/orchestrator.py @@ -17,9 +17,7 @@ from chanlun.core.ChanSBI import ChanSBI from chanlun.core.ChanSEG import ChanSEG from chanlun.core.ChanZS import ChanZS from chanlun.core.ChanBSP import ChanBSP -import talib.abstract as ta import pandas as pd -from technical.util import resample_to_interval from decimal import Decimal import numpy as np from chanlun.indicators.ChanMACD import ChanMACD diff --git a/chanlun/pipeline/resample.py b/chanlun/pipeline/resample.py new file mode 100644 index 0000000..becf918 --- /dev/null +++ b/chanlun/pipeline/resample.py @@ -0,0 +1,52 @@ +"""OHLCV resampling — replaces `technical.util.resample_to_interval`. + +That was the only symbol this project imported from `technical`, which in turn +pulled in the freqtrade dependency chain. Behaviour is preserved exactly, +including the left-labelled bins (rows are candle *open* times) and the +`dropna()` that drops empty intervals. +""" + +from __future__ import annotations + +import pandas as pd + +__all__ = ["TICKER_INTERVAL_MINUTES", "resample_to_interval"] + +TICKER_INTERVAL_MINUTES: dict[str, int] = { + "1m": 1, + "5m": 5, + "15m": 15, + "30m": 30, + "1h": 60, + "60m": 60, + "2h": 120, + "4h": 240, + "6h": 360, + "12h": 720, + "1d": 1440, + "1w": 10080, +} + +_OHLC_AGG = { + "open": "first", + "high": "max", + "low": "min", + "close": "last", + "volume": "sum", +} + + +def resample_to_interval(dataframe: pd.DataFrame, interval: int | str) -> pd.DataFrame: + """Resample OHLCV rows to `interval` minutes (or a timeframe string). + + Merging the result back onto a finer frame requires care to avoid lookahead + bias; this function only resamples. + """ + if isinstance(interval, str): + interval = TICKER_INTERVAL_MINUTES[interval] + + df = dataframe.copy() + df = df.set_index(pd.DatetimeIndex(df["date"])) + df = df.resample(f"{interval}min", label="left").agg(_OHLC_AGG).dropna() + df.reset_index(inplace=True) + return df diff --git a/chanlun/pipeline/timeframe.py b/chanlun/pipeline/timeframe.py index 842a0d5..bbfc1b1 100644 --- a/chanlun/pipeline/timeframe.py +++ b/chanlun/pipeline/timeframe.py @@ -2,9 +2,8 @@ from datetime import timedelta import numpy as np import pandas as pd -import talib.abstract as ta from pandas import DataFrame -from technical.util import resample_to_interval +from chanlun.pipeline.resample import resample_to_interval from chanlun.core.ChanBI import ChanBI from chanlun.core.ChanBIZS import ChanBIZS diff --git a/chanlun/tests/test_ta_compat.py b/chanlun/tests/test_ta_compat.py new file mode 100644 index 0000000..60e0da5 --- /dev/null +++ b/chanlun/tests/test_ta_compat.py @@ -0,0 +1,198 @@ +"""Pin chanlun.indicators.ta to TA-Lib's output, bar for bar. + +These indicators feed the Chan structure builders, so a one-bar shift in the +warm-up or a different smoothing seed silently changes every downstream +bi/seg/zs. Equality against the reference implementation is the only check +that catches that. + +Skipped when talib is unavailable — which is the point of the replacement, so +the suite still has to pass without it. Run in an environment that has talib +whenever chanlun/indicators/ta.py changes. +""" + +from __future__ import annotations + +import sys +import unittest +from pathlib import Path + +import numpy as np +import pandas as pd + +sys.path.insert(0, str(Path(__file__).resolve().parents[2])) + +from chanlun.indicators import ta # noqa: E402 +from chanlun.pipeline.resample import resample_to_interval # noqa: E402 + +try: + import talib.abstract as reference +except ImportError: # pragma: no cover + reference = None + +requires_talib = unittest.skipIf(reference is None, "talib not installed") + + +def make_ohlcv(n: int = 900, seed: int = 7) -> pd.DataFrame: + """Random walk with enough range for BBANDS(365) and EMA(208) to warm up.""" + rng = np.random.default_rng(seed) + close = 100.0 + np.cumsum(rng.normal(0.0, 1.0, n)) + spread = np.abs(rng.normal(0.0, 0.6, n)) + 0.05 + high = close + spread + low = close - spread + open_ = np.concatenate([[close[0]], close[:-1]]) + return pd.DataFrame( + { + "date": pd.date_range("2024-01-01", periods=n, freq="1min", tz="UTC"), + "open": open_, + "high": np.maximum.reduce([high, open_, close]), + "low": np.minimum.reduce([low, open_, close]), + "close": close, + "volume": rng.uniform(1.0, 100.0, n), + } + ) + + +class TAEquivalence(unittest.TestCase): + def setUp(self) -> None: + self.df = make_ohlcv() + + def assertSameSeries(self, got, expected, label: str) -> None: + g = np.asarray(got, dtype=float) + e = np.asarray(expected, dtype=float) + self.assertEqual(g.shape, e.shape, f"{label}: shape") + np.testing.assert_array_equal( + np.isnan(g), np.isnan(e), err_msg=f"{label}: NaN warm-up differs" + ) + mask = ~np.isnan(e) + np.testing.assert_allclose( + g[mask], e[mask], rtol=1e-9, atol=1e-8, err_msg=f"{label}: values differ" + ) + + @requires_talib + def test_sma(self) -> None: + for period in (5, 20, 90, 250): + self.assertSameSeries( + ta.SMA(self.df, timeperiod=period), + reference.SMA(self.df, timeperiod=period), + f"SMA({period})", + ) + + @requires_talib + def test_ma(self) -> None: + for period in (5, 10, 250): + self.assertSameSeries( + ta.MA(self.df, timeperiod=period), + reference.MA(self.df, timeperiod=period), + f"MA({period})", + ) + + @requires_talib + def test_ema(self) -> None: + for period in (5, 7, 10, 13, 24, 26, 30, 52, 104, 156, 208): + self.assertSameSeries( + ta.EMA(self.df, timeperiod=period), + reference.EMA(self.df, timeperiod=period), + f"EMA({period})", + ) + + @requires_talib + def test_rsi(self) -> None: + for period in (7, 14, 21): + self.assertSameSeries( + ta.RSI(self.df, timeperiod=period), + reference.RSI(self.df, timeperiod=period), + f"RSI({period})", + ) + + @requires_talib + def test_atr(self) -> None: + for period in (7, 14, 30): + self.assertSameSeries( + ta.ATR(self.df, timeperiod=period), + reference.ATR(self.df, timeperiod=period), + f"ATR({period})", + ) + + @requires_talib + def test_macd(self) -> None: + for fast, slow, signal in ((12, 26, 9), (26, 52, 9), (5, 35, 5)): + got = ta.MACD(self.df, fastperiod=fast, slowperiod=slow, signalperiod=signal) + exp = reference.MACD(self.df, fastperiod=fast, slowperiod=slow, signalperiod=signal) + for col in ("macd", "macdsignal", "macdhist"): + self.assertSameSeries(got[col], exp[col], f"MACD({fast},{slow},{signal}).{col}") + + @requires_talib + def test_bbands(self) -> None: + cases = ( + (365, 3.0, 3.0), + (120, 3.0, 3.0), + (41, 2.3, 2.3), + (41, 2.0, 2.0), + (26, 3.0, 3.0), + (20, 2.0, 2.0), + (14, 2.0, 2.0), + ) + for period, up, dn in cases: + got = ta.BBANDS(self.df, timeperiod=period, nbdevup=up, nbdevdn=dn, matype=0) + exp = reference.BBANDS(self.df, timeperiod=period, nbdevup=up, nbdevdn=dn, matype=0) + for col in ("upperband", "middleband", "lowerband"): + self.assertSameSeries(got[col], exp[col], f"BBANDS({period},{up},{dn}).{col}") + + @requires_talib + def test_bbands_is_more_accurate_than_talib_on_tiny_windows(self) -> None: + """A deliberate divergence, documented so nobody "fixes" it back. + + TA-Lib derives the variance from sumsq/n - mean**2, which cancels + catastrophically when the window is short and prices are far from zero; + at timeperiod=2 it drifts ~1e-6. Rolling std is accurate there, so the + two disagree. No timeperiod below 14 is used in this codebase, and the + periods that are used agree to ~1e-10 (covered by test_bbands). + """ + got = ta.BBANDS(self.df, timeperiod=2, nbdevup=1.0, nbdevdn=1.0, matype=0)["upperband"] + exp = reference.BBANDS(self.df, timeperiod=2, nbdevup=1.0, nbdevdn=1.0, matype=0)["upperband"] + + window = self.df["close"].rolling(2) + truth = window.mean() + window.std(ddof=0) + ours = np.nanmax(np.abs((got - truth).to_numpy())) + theirs = np.nanmax(np.abs((exp - truth).to_numpy())) + self.assertLess(ours, 1e-9) + self.assertLess(ours, theirs) + + @requires_talib + def test_matches_on_real_price_scale(self) -> None: + """Guard against tolerances that only hold near 100.""" + df = self.df.copy() + for col in ("open", "high", "low", "close"): + df[col] *= 900.0 + self.assertSameSeries( + ta.ATR(df, timeperiod=14), reference.ATR(df, timeperiod=14), "ATR@scale" + ) + self.assertSameSeries( + ta.RSI(df, timeperiod=14), reference.RSI(df, timeperiod=14), "RSI@scale" + ) + + +class ResampleEquivalence(unittest.TestCase): + @unittest.skipIf( + __import__("importlib").util.find_spec("technical") is None, + "technical not installed", + ) + def test_matches_technical(self) -> None: + from technical.util import resample_to_interval as ref_resample + + df = make_ohlcv(600) + for interval in (5, 15, 60, "5m", "1h"): + got = resample_to_interval(df, interval) + exp = ref_resample(df, interval) + pd.testing.assert_frame_equal(got, exp, check_exact=False, rtol=1e-12) + + def test_shapes_without_reference(self) -> None: + df = make_ohlcv(120) + out = resample_to_interval(df, 5) + self.assertEqual(list(out.columns), ["date", "open", "high", "low", "close", "volume"]) + self.assertLessEqual(len(out), 120 // 5 + 1) + self.assertTrue((out["high"] >= out["low"]).all()) + + +if __name__ == "__main__": + unittest.main() diff --git a/web/services/cn_stock.py b/web/services/cn_stock.py index 449306e..ccf23d2 100644 --- a/web/services/cn_stock.py +++ b/web/services/cn_stock.py @@ -574,7 +574,7 @@ class ChinaStockData: def add_indicators(self, df): """添加技术指标""" try: - import talib.abstract as ta + from chanlun.indicators import ta import numpy as np # MACD指标 diff --git a/web/services/runtime/analyze.py b/web/services/runtime/analyze.py index db2db5c..274045e 100644 --- a/web/services/runtime/analyze.py +++ b/web/services/runtime/analyze.py @@ -1,7 +1,7 @@ from __future__ import annotations import numpy as np -import talib.abstract as ta +from chanlun.indicators import ta from chanlun import TF_DF from chanlun.core.ChanEnum import Chan_KLC_FX, Chan_FX_TYPE diff --git a/web/services/runtime/indicators.py b/web/services/runtime/indicators.py index f465c4a..f295fe5 100644 --- a/web/services/runtime/indicators.py +++ b/web/services/runtime/indicators.py @@ -1,6 +1,6 @@ from __future__ import annotations -import talib.abstract as ta +from chanlun.indicators import ta from . import state def add_indicators(df):