Files
MoFin/deploy/profile-scripts/temp_band.py
T

112 lines
3.6 KiB
Python
Raw 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
DB = Path("/home/hmo/MoFin/data/mofin.db")
INDEX = "sh000001"
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(db_path=None):
"""读取最新市场温度(rsi + 档位 + regime)。"""
db = db_path or DB
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,)
).fetchall()
# 最新 regime
reg = conn.execute(
"SELECT date, regime FROM market_regime ORDER BY date DESC LIMIT 1"
).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')}")