120 lines
4.0 KiB
Python
120 lines
4.0 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""temp_band.py — 市场温度分档模块(2026-08-13)
|
||
|
||
结论(数据验证):三态骨架 + 温度(rsi)连续维度。
|
||
- 三态(trend_up/choppy/trend_down)= market_regime.py 已有(策略类型分类)
|
||
- 温度(rsi 分档)= 本模块(力度维度,仓位乘数)
|
||
|
||
数据依据(2026-08-13 分市场状态检验):
|
||
v_weak choppy: 恐慌<45 -> 71%胜率 / 中性45-60 -> 46% / 亢奋>60 -> 37%
|
||
v_oversold trend_down: 恐慌<35 -> 64% / 偏弱35-50 -> 53% / 中性>50 -> 0%
|
||
s2_panic trend_down 恐慌<35: 86%胜率 +16.6%
|
||
|
||
档位:panic(<35) / fear(35-45) / neutral(45-60) / greed(60-70) / euphoria(>70)
|
||
用法:
|
||
from temp_band import get_market_temp, temp_multiplier
|
||
t = get_market_temp() # {rsi, band, regime}
|
||
mult = temp_multiplier(t["band"], strategy_family="mr")
|
||
"""
|
||
import sqlite3
|
||
from pathlib import Path
|
||
|
||
from market_config import MARKETS
|
||
|
||
DB = Path("/home/hmo/MoFin/data/mofin.db")
|
||
INDEX = "sh000001" # A股默认指数(兜底);实际以 MARKETS[market]['index_code'] 为准
|
||
|
||
|
||
def calc_rsi(series, n=14):
|
||
result = [None] * len(series)
|
||
if len(series) < n + 1:
|
||
return result
|
||
gains, losses = [], []
|
||
for i in range(1, len(series)):
|
||
ch = series[i] - series[i - 1]
|
||
gains.append(max(ch, 0))
|
||
losses.append(max(-ch, 0))
|
||
if i >= n:
|
||
avg_g = sum(gains[i - n:i]) / n
|
||
avg_l = sum(losses[i - n:i]) / n
|
||
rs = avg_g / avg_l if avg_l > 0 else 100
|
||
result[i] = 100 - 100 / (1 + rs)
|
||
return result
|
||
|
||
|
||
def temp_band(rsi):
|
||
"""连续 rsi -> 温度档位"""
|
||
if rsi is None:
|
||
return "unknown"
|
||
if rsi < 35:
|
||
return "panic"
|
||
if rsi < 45:
|
||
return "fear"
|
||
if rsi < 60:
|
||
return "neutral"
|
||
if rsi < 70:
|
||
return "greed"
|
||
return "euphoria"
|
||
|
||
|
||
def temp_multiplier(band, strategy_family="mr"):
|
||
"""温度 -> 仓位乘数。mr=均值回复(恐慌重仓),trend=趋势(恐慌回避)"""
|
||
if strategy_family == "trend":
|
||
return {
|
||
"panic": 0.0, "fear": 0.0, "neutral": 0.5,
|
||
"greed": 1.0, "euphoria": 0.5, "unknown": 0.5,
|
||
}.get(band, 0.5)
|
||
return {
|
||
"panic": 1.5, "fear": 1.0, "neutral": 0.8,
|
||
"greed": 0.5, "euphoria": 0.3, "unknown": 0.5,
|
||
}.get(band, 0.5)
|
||
|
||
|
||
def get_market_temp(market='a', db_path=None):
|
||
"""读取指定市场最新温度(rsi + 档位 + regime)。
|
||
|
||
market: 'a'(默认, A股 sh000001) / 'hk'(港股 hkHSI)。
|
||
指数代码从 market_config.MARKETS 读取,不写死。
|
||
"""
|
||
db = db_path or DB
|
||
index_code = MARKETS.get(market, MARKETS['a'])['index_code']
|
||
conn = sqlite3.connect(str(db), timeout=5)
|
||
try:
|
||
# 指数最近 30 日收盘算 rsi
|
||
rows = conn.execute(
|
||
"SELECT date, close FROM stock_daily WHERE code=? ORDER BY date DESC LIMIT 30",
|
||
(index_code,)
|
||
).fetchall()
|
||
# 最新 regime(按市场过滤——否则 A 股温度会误读港股 regime)
|
||
reg = conn.execute(
|
||
"SELECT date, regime FROM market_regime WHERE market=? ORDER BY date DESC LIMIT 1",
|
||
(market,)
|
||
).fetchone()
|
||
finally:
|
||
conn.close()
|
||
if not rows or len(rows) < 15:
|
||
return {"rsi": None, "band": "unknown", "regime": "unknown", "date": ""}
|
||
rows = list(reversed(rows))
|
||
closes = [r[1] for r in rows]
|
||
dates = [r[0] for r in rows]
|
||
rsi_series = calc_rsi(closes)
|
||
rsi_now = rsi_series[-1] if rsi_series else None
|
||
band = temp_band(rsi_now)
|
||
regime = reg[1] if reg else "unknown"
|
||
reg_date = reg[0] if reg else ""
|
||
return {
|
||
"rsi": round(rsi_now, 1) if rsi_now is not None else None,
|
||
"band": band,
|
||
"regime": regime,
|
||
"date": dates[-1],
|
||
"regime_date": reg_date,
|
||
}
|
||
|
||
|
||
if __name__ == "__main__":
|
||
t = get_market_temp()
|
||
print(f"温度: rsi={t['rsi']} band={t['band']} regime={t['regime']} ({t['date']})")
|
||
print(f" mr 乘数: {temp_multiplier(t['band'], 'mr')}")
|
||
print(f" trend 乘数: {temp_multiplier(t['band'], 'trend')}")
|