7 Sprints 全部完成: Sprint 0: 基础设施 (DataManager + MariaDB) Sprint 1: 因子引擎 (34因子/12分类) Sprint 2: VectorBT 回测 (5策略+截面) Sprint 3: Optuna 优化 (+Walk-Forward) Sprint 4: ML 模型 (LightGBM+CatBoost) Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐) Sprint 6: Agent 系统 (4Agent+日报.md/.html) 生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping, save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复, RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4, CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化, indexDatas API修正 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
230 lines
8.1 KiB
Python
230 lines
8.1 KiB
Python
"""
|
|
Qwen API 封装(DashScope / 本地 Ollama)。
|
|
|
|
文本 → 情绪分析 → 标准化 JSON 输出。
|
|
"""
|
|
|
|
import json
|
|
import os
|
|
import time
|
|
from typing import Optional
|
|
|
|
import requests
|
|
|
|
# 确保 .env 已加载(无论从哪个路径导入)
|
|
import config.settings # noqa: F401
|
|
|
|
|
|
class QwenClient:
|
|
"""
|
|
Qwen 情绪分析客户端。
|
|
|
|
支持两种后端:
|
|
1. DashScope API(云)- 需要 QWEN_API_KEY
|
|
2. 本地 Ollama(本地)- 需要 QWEN_LOCAL_BASE_URL
|
|
|
|
优先级:本地 > API
|
|
|
|
参数:
|
|
api_key: DashScope API Key,默认从环境变量 QWEN_API_KEY 读取
|
|
model: 模型名称,默认从 QWEN_MODEL 读取
|
|
local_base_url: Ollama 地址,默认从 QWEN_LOCAL_BASE_URL 读取
|
|
local_model: Ollama 模型名,默认从 QWEN_LOCAL_MODEL 读取
|
|
max_retries: 失败重试次数
|
|
cache_enabled: 是否启用结果缓存(按文本 hash)
|
|
"""
|
|
|
|
# 情绪分析 Prompt
|
|
SYSTEM_PROMPT = """你是一个专业的金融情绪分析专家。分析以下 A 股相关的新闻/公告文本,返回 JSON 格式的分析结果。
|
|
|
|
分析维度:
|
|
1. sentiment_score: -1.0(极度利空) ~ 0(中性) ~ +1.0(极度利好),保留1位小数
|
|
2. impact_duration: "short"(1-3个交易日) | "medium"(1-2周) | "long"(1个月以上)
|
|
3. confidence: 0.0~1.0 置信度
|
|
4. key_topics: 涉及的关键主题列表(最多5个)
|
|
5. affected_factors: 可能影响的基本面/技术面因子类型列表
|
|
|
|
重要规则:
|
|
- 只考虑对股价的直接影响
|
|
- 中性新闻(常规报道、例行公告)给 0 分
|
|
- 重大利好(业绩超预期、政策支持、大订单)给 >0.5 分
|
|
- 重大利空(亏损、处罚、减持、诉讼)给 <-0.5 分
|
|
- 如果不确定,confidences 要低 (<0.5)
|
|
|
|
返回格式(仅 JSON):
|
|
{"sentiment_score": 0.0, "impact_duration": "short", "confidence": 0.5, "key_topics": [], "affected_factors": []}"""
|
|
|
|
def __init__(
|
|
self,
|
|
api_key: str | None = None,
|
|
model: str | None = None,
|
|
local_base_url: str | None = None,
|
|
local_model: str | None = None,
|
|
max_retries: int = 3,
|
|
cache_enabled: bool = True,
|
|
):
|
|
self.api_key = api_key or os.getenv("QWEN_API_KEY", "")
|
|
self.model = model or os.getenv("QWEN_MODEL", "qwen-turbo")
|
|
self.local_base_url = local_base_url or os.getenv("QWEN_LOCAL_BASE_URL", "")
|
|
self.local_model = local_model or os.getenv("QWEN_LOCAL_MODEL", "qwen2.5:7b")
|
|
self.max_retries = max_retries
|
|
self.cache_enabled = cache_enabled
|
|
self._cache: dict[str, dict] = {}
|
|
self._is_local = bool(self.local_base_url)
|
|
|
|
# ── 单条分析 ──────────────────────────────────────────
|
|
|
|
def analyze_sentiment(self, text: str) -> dict:
|
|
"""
|
|
分析单条文本的情绪。
|
|
|
|
返回:
|
|
{"sentiment_score": float, "impact_duration": str, "confidence": float,
|
|
"key_topics": list, "affected_factors": list}
|
|
"""
|
|
if not text or not text.strip():
|
|
return self._empty_result()
|
|
|
|
cache_key = str(hash(text))
|
|
if self.cache_enabled and cache_key in self._cache:
|
|
return self._cache[cache_key]
|
|
|
|
for attempt in range(self.max_retries):
|
|
try:
|
|
if self._is_local:
|
|
result = self._call_ollama(text)
|
|
else:
|
|
result = self._call_dashscope(text)
|
|
|
|
if result and "sentiment_score" in result:
|
|
if self.cache_enabled:
|
|
self._cache[cache_key] = result
|
|
return result
|
|
|
|
except Exception as e:
|
|
if attempt < self.max_retries - 1:
|
|
time.sleep(1 * (attempt + 1))
|
|
else:
|
|
print(f" [ERROR] Qwen 分析失败(重试{self.max_retries}次): {e}")
|
|
|
|
return self._empty_result()
|
|
|
|
# ── 批量分析 ──────────────────────────────────────────
|
|
|
|
def analyze_batch(
|
|
self,
|
|
texts: list[str],
|
|
batch_size: int = 10,
|
|
progress_callback: callable = None, # type: ignore
|
|
) -> list[dict]:
|
|
"""
|
|
批量分析(逐条调用,节省并发成本)。
|
|
|
|
参数:
|
|
texts: 文本列表
|
|
batch_size: 每批次数量
|
|
progress_callback: 进度回调 fn(done, total)
|
|
|
|
返回:
|
|
[{"sentiment_score": ..., ...}, ...]
|
|
"""
|
|
results = []
|
|
total = len(texts)
|
|
for i, text in enumerate(texts):
|
|
result = self.analyze_sentiment(text)
|
|
results.append(result)
|
|
if progress_callback:
|
|
progress_callback(i + 1, total)
|
|
# 批次间短暂休息,避免触发限频
|
|
if (i + 1) % batch_size == 0 and i < total - 1:
|
|
time.sleep(0.5)
|
|
return results
|
|
|
|
# ── DashScope API ──────────────────────────────────────
|
|
|
|
def _call_dashscope(self, text: str) -> dict:
|
|
"""调用 DashScope API。"""
|
|
resp = requests.post(
|
|
"https://dashscope.aliyuncs.com/compatible-mode/v1/chat/completions",
|
|
headers={
|
|
"Authorization": f"Bearer {self.api_key}",
|
|
"Content-Type": "application/json",
|
|
},
|
|
json={
|
|
"model": self.model,
|
|
"messages": [
|
|
{"role": "system", "content": self.SYSTEM_PROMPT},
|
|
{"role": "user", "content": text[:4000]}, # 截断长文本
|
|
],
|
|
"temperature": 0.1,
|
|
"max_tokens": 500,
|
|
},
|
|
timeout=30,
|
|
)
|
|
resp.raise_for_status()
|
|
body = resp.json()
|
|
content = body["choices"][0]["message"]["content"]
|
|
return self._parse_json(content)
|
|
|
|
# ── Ollama 本地调用 ────────────────────────────────────
|
|
|
|
def _call_ollama(self, text: str) -> dict:
|
|
"""调用本地 Ollama。"""
|
|
resp = requests.post(
|
|
f"{self.local_base_url.rstrip('/')}/chat/completions",
|
|
json={
|
|
"model": self.local_model,
|
|
"messages": [
|
|
{"role": "system", "content": self.SYSTEM_PROMPT},
|
|
{"role": "user", "content": text[:4000]},
|
|
],
|
|
"temperature": 0.1,
|
|
"max_tokens": 500,
|
|
},
|
|
timeout=60,
|
|
)
|
|
resp.raise_for_status()
|
|
body = resp.json()
|
|
content = body["choices"][0]["message"]["content"]
|
|
return self._parse_json(content)
|
|
|
|
# ── JSON 解析 ──────────────────────────────────────────
|
|
|
|
@staticmethod
|
|
def _parse_json(content: str) -> dict:
|
|
"""从模型输出中提取 JSON。"""
|
|
# 尝试直接解析
|
|
try:
|
|
return json.loads(content)
|
|
except json.JSONDecodeError:
|
|
pass
|
|
# 尝试提取 ```json ... ``` 代码块
|
|
if "```json" in content:
|
|
start = content.find("```json") + 7
|
|
end = content.find("```", start)
|
|
if end > start:
|
|
try:
|
|
return json.loads(content[start:end].strip())
|
|
except json.JSONDecodeError:
|
|
pass
|
|
# 尝试找 { ... }
|
|
brace_start = content.find("{")
|
|
brace_end = content.rfind("}")
|
|
if brace_start >= 0 and brace_end > brace_start:
|
|
try:
|
|
return json.loads(content[brace_start:brace_end + 1])
|
|
except json.JSONDecodeError:
|
|
pass
|
|
|
|
return QwenClient._empty_result()
|
|
|
|
@staticmethod
|
|
def _empty_result() -> dict:
|
|
return {
|
|
"sentiment_score": 0.0,
|
|
"impact_duration": "short",
|
|
"confidence": 0.0,
|
|
"key_topics": [],
|
|
"affected_factors": [],
|
|
}
|