# cell 9:LLM 趋势研判(相似周检索 + 可选离散化)
def llm_forecast(record, config, prev_position=0.0):
api_key, base_url = env("LLM_API_KEY"), env("LLM_BASE_URL")
model = config.model or env("LLM_MODEL") or "deepseek-v4-flash"
if not api_key:
raise RuntimeError("未设置 LLM_API_KEY")
if not base_url:
base_url = "https://api.deepseek.com/v1"
def _has_nan(week):
return any(np.isnan(v) for v in week["rsi5"] + week["zscore5"] + [week["pnl"], week["atr_pct"]])
if _has_nan(record["this_week"]):
raise ValueError("特征含 NaN(预热期未满),该周不调用 LLM")
top_k, dists = knn_history(record["this_week"], record["history"], config.knn_k)
instruction = {
"任务": ("你是周频趋势研判员。结合给定的【相似历史周】,判断当前市场处于:上涨趋势、下跌趋势,"
"还是即将反转/正在反转/已经反转;并决定下周动作:买入、加大仓位、减少仓位或空仓观望。"
"只做多,不做空。研究回测,非投资建议。"),
"说明": ("下方仅列出与当前周特征最相似的 %d 个历史周。每个都带下一周实际结果 next_pnl/next_sharpe;"
"请优先依据这些最相关样本决策,而非泛泛而谈。证据不足保持中等仓位。" % len(top_k)),
"特征含义": {"pnl": "该周涨跌幅", "atr_pct": "该周 5 日 ATR/收盘价 均值",
"rsi5": "每个交易日当天及前 4 天 5 日 RSI", "zscore5": "相对 5 日均线 z-score",
"next_pnl": "下一周实际涨跌幅", "next_sharpe": "下一周日收益年化夏普"},
"仓位自洽规则": ("『买入』『加大仓位』时 position 必须【大于】prev_position;"
"『减少仓位』『空仓观望』时 position 必须【小于】prev_position。"),
"prev_position": prev_position,
"当前周": {k: record["this_week"][k] for k in ["date", "pnl", "atr_pct", "rsi5", "zscore5"]},
"相似历史周(共 %d 个)" % len(top_k): [dict(h) for h in top_k],
"输出格式": {"trend": "上涨趋势|下跌趋势|反转酝酿|反转进行中|反转已确立",
"action": "买入|加大仓位|减少仓位|空仓观望",
"position": "下周目标仓位 [0,1],与 action 自洽", "confidence": "0-1",
"reason": "一句话理由(引用最相似的 1~2 个历史周)"},
}
payload = {"model": model, "stream": False, "max_tokens": 16000, "reasoning_effort": "max",
"messages": [{"role": "system", "content": "只输出 JSON。"},
{"role": "user", "content": json.dumps(instruction, ensure_ascii=False)}]}
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
url = base_url.rstrip("/") + "/chat/completions"
req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/json",
"Authorization": "Bearer " + api_key, "User-Agent": "Mozilla/5.0"}, method="POST")
last_error, text = None, None
for _ in range(3):
try:
with urllib.request.urlopen(req, timeout=180) as resp:
text = json.loads(resp.read().decode("utf-8"))["choices"][0]["message"]["content"]
last_error = None
break
except (urllib.error.URLError, TimeoutError, OSError) as e:
last_error = e
if last_error is not None:
raise RuntimeError("LLM 请求失败:%s" % last_error) from last_error
a, b = text.find("{"), text.rfind("}")
if a < 0 or b < a:
raise ValueError("LLM 未返回 JSON。")
parsed = json.loads(text[a:b + 1])
action = str(parsed.get("action", "")).strip()
position = min(1.0, max(0.0, float(parsed.get("position", 0.5))))
position = resolve_position(prev_position, action, position,
quantize=config.quantize, levels=config.quant_levels)
return {"trend": str(parsed.get("trend", "")), "action": action, "position": position,
"confidence": float(parsed.get("confidence", 0.0)), "reason": str(parsed.get("reason", "")),
"knn_distances": dists}