Files

120 lines
4.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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')}")