Files
OpenFin/backend/routers/causal_lab.py
OpenSquared 9ac3ebb4b5 causal lab
2026-06-27 21:45:10 +02:00

525 lines
21 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.
"""
Causal Lab — Validation empirique du graphe causal EUR/USD.
Compare les prédictions du modèle aux mouvements réels autour des publications
macro (CPI, NFP, FOMC, ECB...). Détecte également les pré-drifts (fuites d'info).
Endpoints:
GET /api/causal-lab/events — liste des événements analysables
GET /api/causal-lab/event/{id}/analyze — analyse complète (avec cache)
GET /api/causal-lab/summary — agrégat toutes séries
"""
import json
import logging
import math
from datetime import datetime, timedelta
from typing import Optional
from fastapi import APIRouter, HTTPException, Query
router = APIRouter()
logger = logging.getLogger(__name__)
# ── Paramètres neutres de référence ──────────────────────────────────────────
# Tous les inputs à zéro / valeur neutre — on n'injecte que la surprise de l'event
NEUTRAL: dict = {
"fed_rate": 4.25, "ecb_rate": 3.65,
"fed_tone": 0, "ecb_tone": 0,
"cpi_us_surprise": 0, "cpi_eu_surprise": 0,
"nfp_surprise": 0,
"pmi_us": 50, "pmi_eu": 50,
"vix": 18.0, "oil": 80.0, "real_yield_us": 2.1,
"us_2y": 4.50, "us_10y": 4.30,
"eu_2y": 2.80, "eu_10y": 2.60,
"eurusd": 1.1450,
}
# ── Mapping série FRED → nœud causal ─────────────────────────────────────────
SERIES_MAP: dict = {
# Canal forward guidance FED
"CPIAUCSL": {"name": "CPI US", "side": "US", "ch": "fed_fwd", "key": "cpi_us_surprise"},
"CPILFESL": {"name": "Core CPI US", "side": "US", "ch": "fed_fwd", "key": "cpi_us_surprise"},
"PAYEMS": {"name": "NFP", "side": "US", "ch": "fed_fwd", "key": "nfp_surprise"},
"MANEMP": {"name": "Emploi Manuf.","side": "US", "ch": "fed_fwd", "key": "nfp_surprise"},
"UNRATE": {"name": "Chômage US", "side": "US", "ch": "fed_fwd", "key": "nfp_surprise", "flip": True},
# Canal taux directeur FED
"FEDFUNDS": {"name": "Fed Funds", "side": "US", "ch": "fed_rate", "key": "fed_rate", "level": True},
"DFF": {"name": "Fed Funds D", "side": "US", "ch": "fed_rate", "key": "fed_rate", "level": True},
# Canal BCE
"ECBDFR": {"name": "ECB Rate", "side": "EU", "ch": "ecb_rate", "key": "ecb_rate", "level": True},
# CPI EU
"HICP": {"name": "CPI EU", "side": "EU", "ch": "ecb_fwd", "key": "cpi_eu_surprise"},
"HICPEI": {"name": "Core CPI EU", "side": "EU", "ch": "ecb_fwd", "key": "cpi_eu_surprise"},
# Canal secondaire
"DFII10": {"name": "Taux réel US", "side": "US", "ch": "secondary","key": "real_yield_us", "level": True},
# PMI (niveau absolu)
"ISMMAN": {"name": "ISM US", "side": "US", "ch": "pmi_us", "key": "pmi_us", "level_pmi": True},
"NAPM": {"name": "PMI US", "side": "US", "ch": "pmi_us", "key": "pmi_us", "level_pmi": True},
}
RELEVANT_SERIES = set(SERIES_MAP.keys())
# ── Port Python exact du compute() TypeScript (EuroSimulator) ─────────────────
def _compute(p: dict, base: dict) -> dict:
def g(k): return p.get(k, base.get(k, NEUTRAL[k]))
def gb(k): return base.get(k, NEUTRAL[k])
fed_rp = (g("fed_rate") - gb("fed_rate")) / 0.25
ecb_rp = (g("ecb_rate") - gb("ecb_rate")) / 0.25
fed_fwd = (
-g("fed_tone") * 1.2
+ g("cpi_us_surprise") / 0.1 * 0.50
+ g("nfp_surprise") / 100 * 0.35
+ (g("pmi_us") - 50) / 5 * 0.18
)
ecb_fwd = (
-g("ecb_tone") * 1.2
+ g("cpi_eu_surprise") / 0.1 * 0.45
+ (g("pmi_eu") - 50) / 5 * 0.22
)
us2d = fed_rp * 0.085 + fed_fwd * 0.030
eu2d = ecb_rp * 0.080 + ecb_fwd * 0.025
us10d = fed_rp * 0.035 + fed_fwd * 0.070 + (g("pmi_us") - 50) / 50 * 0.012
eu10d = ecb_rp * 0.030 + ecb_fwd * 0.060 + (g("pmi_eu") - 50) / 50 * 0.012
dd2 = us2d - eu2d
dd10 = us10d - eu10d
pmi_diff = (g("pmi_eu") - 50) - (g("pmi_us") - 50)
vix_dev = g("vix") - gb("vix")
ry_dev = g("real_yield_us") - gb("real_yield_us")
oil_dev = g("oil") - gb("oil")
c_2y = round(-dd2 * 500)
c_10y = round(-dd10 * 200)
c_pmi = round(pmi_diff * 7)
c_vix = round(-vix_dev * 4)
c_ry = round(-ry_dev * 100)
c_oil = round(oil_dev * 0.5)
total = c_2y + c_10y + c_pmi + c_vix + c_ry + c_oil
return {
"fed_rate_pressure": round(fed_rp, 3),
"ecb_rate_pressure": round(ecb_rp, 3),
"fed_fwd_signal": round(fed_fwd, 3),
"ecb_fwd_signal": round(ecb_fwd, 3),
"us_2y_delta": round(us2d, 3),
"eu_2y_delta": round(eu2d, 3),
"us_10y_delta": round(us10d, 3),
"eu_10y_delta": round(eu10d, 3),
"delta_diff_2y": round(dd2, 3),
"delta_diff_10y": round(dd10, 3),
"c_2y": c_2y, "c_10y": c_10y, "c_pmi": c_pmi,
"c_vix": c_vix, "c_ry": c_ry, "c_oil": c_oil,
"total_pips": total,
}
def _build_scenario(series_id: str, actual: float, forecast: Optional[float], previous: Optional[float]):
"""Construit le scénario (p, base) isolant l'effet marginal de la surprise."""
cfg = SERIES_MAP.get(series_id)
if not cfg:
return None, None
p = {**NEUTRAL}
base = {**NEUTRAL}
if cfg.get("level"):
# Événement de taux : surprise = actual vs niveau attendu
expected = forecast if forecast is not None else previous
if expected is None:
return None, None
p[cfg["key"]] = actual
base[cfg["key"]] = expected
elif cfg.get("level_pmi"):
# PMI : niveau absolu vs consensus
expected = forecast if forecast is not None else 50.0
p[cfg["key"]] = actual
base[cfg["key"]] = expected
else:
# Surprise additive : actual forecast (en unités du modèle)
if forecast is None:
return None, None
surprise = actual - forecast
if cfg.get("flip"):
surprise = -surprise
p[cfg["key"]] = surprise
return p, base
# ── Récupération des prix autour d'un événement ───────────────────────────────
def _fetch_prices(event_date_str: str) -> dict:
out: dict = {"mode": "none", "eurusd": [], "us10y": [], "eu10y": [], "us2y": []}
try:
import yfinance as yf
event_dt = datetime.strptime(event_date_str[:10], "%Y-%m-%d")
days_ago = (datetime.utcnow() - event_dt).days
# Intraday 5min — disponible ~55 jours
if days_ago < 55:
start = (event_dt - timedelta(days=1)).strftime("%Y-%m-%d")
end = (event_dt + timedelta(days=2)).strftime("%Y-%m-%d")
df = yf.download("EURUSD=X", start=start, end=end,
interval="5m", progress=False, auto_adjust=True)
if df is not None and len(df) > 0:
if hasattr(df.columns, "levels"):
df.columns = df.columns.get_level_values(0)
day_df = df[df.index.date == event_dt.date()]
rows = [(idx, float(row["Close"])) for idx, row in day_df.iterrows()
if not (math.isnan(row["Close"]) if isinstance(row["Close"], float) else False)]
if len(rows) >= 6:
out["mode"] = "intraday_5m"
out["eurusd"] = [{"t": idx.isoformat(), "c": round(c, 5)} for idx, c in rows]
# Fallback journalier
if out["mode"] == "none":
start = (event_dt - timedelta(days=5)).strftime("%Y-%m-%d")
end = (event_dt + timedelta(days=5)).strftime("%Y-%m-%d")
df = yf.download("EURUSD=X", start=start, end=end,
interval="1d", progress=False, auto_adjust=True)
if df is not None and len(df) > 0:
if hasattr(df.columns, "levels"):
df.columns = df.columns.get_level_values(0)
out["mode"] = "daily"
out["eurusd"] = [
{"t": str(idx.date()), "c": round(float(row["Close"]), 5)}
for idx, row in df.iterrows()
if not (math.isnan(float(row["Close"])) if row["Close"] == row["Close"] else True)
]
# Taux journaliers (toujours)
start = (event_dt - timedelta(days=4)).strftime("%Y-%m-%d")
end = (event_dt + timedelta(days=4)).strftime("%Y-%m-%d")
for sym, key in [("^TNX", "us10y"), ("GE10YT=RR", "eu10y"), ("US2YT=RR", "us2y")]:
try:
ydf = yf.download(sym, start=start, end=end,
interval="1d", progress=False, auto_adjust=True)
if ydf is None or len(ydf) == 0:
continue
if hasattr(ydf.columns, "levels"):
ydf.columns = ydf.columns.get_level_values(0)
out[key] = [
{"t": str(idx.date()), "c": round(float(row["Close"]), 3)}
for idx, row in ydf.iterrows()
if row["Close"] == row["Close"] # not NaN
]
except Exception as e:
logger.debug(f"[causal_lab] yield {sym}: {e}")
except Exception as e:
logger.error(f"[causal_lab] _fetch_prices: {e}")
return out
# ── Métriques de drift ────────────────────────────────────────────────────────
def _drift_metrics(prices: dict, event_date_str: str) -> dict:
eurusd = prices.get("eurusd", [])
mode = prices.get("mode", "none")
edate = event_date_str[:10]
empty = {"pre_pips": None, "post_pips": None, "drift_ratio": None, "leak": "unknown"}
if not eurusd:
return empty
if mode == "intraday_5m":
n = len(eurusd)
if n < 8:
return empty
mid = n // 2
pre_pips = round((eurusd[mid - 1]["c"] - eurusd[0]["c"]) * 10000)
post_pips = round((eurusd[-1]["c"] - eurusd[mid]["c"]) * 10000)
else:
pre = [b for b in eurusd if b["t"] < edate]
same = [b for b in eurusd if b["t"] == edate]
if not pre or not same:
return empty
pre_pips = None
post_pips = round((same[0]["c"] - pre[-1]["c"]) * 10000)
ratio: Optional[float] = None
if pre_pips is not None and post_pips and post_pips != 0:
ratio = round(pre_pips / post_pips, 3)
leak = "unknown"
if ratio is not None:
ar = abs(ratio)
leak = "high" if ar > 0.5 else "medium" if ar > 0.25 else "low"
return {"pre_pips": pre_pips, "post_pips": post_pips, "drift_ratio": ratio, "leak": leak}
def _yield_delta(series: list, event_date: str) -> Optional[float]:
if not series:
return None
pre = [b for b in series if b["t"] < event_date]
post = [b for b in series if b["t"] >= event_date]
if not pre or not post:
return None
return round(post[0]["c"] - pre[-1]["c"], 3)
# ── Score d'activation ────────────────────────────────────────────────────────
def _activation(model: dict, us10y: Optional[float], eu10y: Optional[float],
post_pips: Optional[int]) -> dict:
def match(pred: float, act: Optional[float], thr: float = 0.01) -> str:
if act is None:
return "unknown"
if abs(pred) < thr:
return "neutral"
return "correct" if (pred > 0) == (act > 0) else "wrong"
nodes = {
"us_10y": {"pred": model["us_10y_delta"], "act": us10y, "status": match(model["us_10y_delta"], us10y)},
"eu_10y": {"pred": model["eu_10y_delta"], "act": eu10y, "status": match(model["eu_10y_delta"], eu10y)},
"eurusd": {"pred": model["total_pips"], "act": post_pips, "status": match(model["total_pips"], post_pips, 5)},
}
active = [v for v in nodes.values() if v["status"] not in ("unknown", "neutral")]
correct = sum(1 for v in active if v["status"] == "correct")
score = round(correct / len(active), 2) if active else None
return {"score": score, "nodes": nodes, "correct": correct, "total": len(active)}
# ── Cache ─────────────────────────────────────────────────────────────────────
def _ensure_cache(conn):
conn.execute("""
CREATE TABLE IF NOT EXISTS causal_lab_cache (
id INTEGER PRIMARY KEY,
event_id INTEGER UNIQUE,
result_json TEXT,
analyzed_at TEXT
)
""")
conn.commit()
# ── Endpoints ─────────────────────────────────────────────────────────────────
@router.get("/api/causal-lab/events")
def list_events(limit: int = Query(120, le=500), series: str = Query("")):
"""Liste des événements macro pertinents pour le modèle EUR/USD."""
try:
from services.database import get_conn
conn = get_conn()
_ensure_cache(conn)
series_filter = (
f" AND e.series_id = '{series}'" if series
else " AND e.series_id IN ({})".format(
",".join(f"'{s}'" for s in RELEVANT_SERIES))
)
rows = conn.execute(f"""
SELECT e.id, e.event_name, e.series_id, e.event_date,
e.actual_value, e.forecast_value, e.previous_value,
e.surprise_pct, e.surprise_direction,
c.analyzed_at, c.result_json
FROM economic_events e
LEFT JOIN causal_lab_cache c ON c.event_id = e.id
WHERE e.actual_value IS NOT NULL
AND e.forecast_value IS NOT NULL
{series_filter}
ORDER BY e.event_date DESC
LIMIT {limit}
""").fetchall()
conn.close()
result = []
for r in rows:
item = dict(r)
if item.get("result_json"):
try:
cached = json.loads(item["result_json"])
item["activation_score"] = cached.get("activation", {}).get("score")
item["predicted_pips"] = cached.get("model", {}).get("total_pips")
item["actual_pips"] = cached.get("drift", {}).get("post_pips")
item["leak"] = cached.get("drift", {}).get("leak")
item["drift_ratio"] = cached.get("drift", {}).get("drift_ratio")
except Exception:
pass
item.pop("result_json", None)
item["cfg"] = SERIES_MAP.get(item["series_id"], {})
result.append(item)
return result
except Exception as e:
logger.error(f"[causal_lab] list_events: {e}")
raise HTTPException(500, str(e))
@router.get("/api/causal-lab/event/{event_id}/analyze")
def analyze_event(event_id: int, force: bool = Query(False)):
"""Analyse complète d'un événement : prédiction modèle + données réelles + activation."""
try:
from services.database import get_conn
conn = get_conn()
_ensure_cache(conn)
# Vérifier le cache
if not force:
cached = conn.execute(
"SELECT result_json FROM causal_lab_cache WHERE event_id = ?", (event_id,)
).fetchone()
if cached:
conn.close()
return json.loads(cached["result_json"])
# Charger l'événement
row = conn.execute("SELECT * FROM economic_events WHERE id = ?", (event_id,)).fetchone()
if not row:
conn.close()
raise HTTPException(404, "Événement introuvable")
ev = dict(row)
cfg = SERIES_MAP.get(ev["series_id"])
if not cfg:
conn.close()
raise HTTPException(400, f"Série {ev['series_id']} non mappée dans le modèle causal")
# Construire le scénario isolé
p, base = _build_scenario(
ev["series_id"], ev["actual_value"],
ev.get("forecast_value"), ev.get("previous_value"),
)
if p is None:
conn.close()
raise HTTPException(400, "Impossible de construire le scénario (forecast manquant)")
model = _compute(p, base)
prices = _fetch_prices(ev["event_date"])
drift = _drift_metrics(prices, ev["event_date"])
edate = ev["event_date"][:10]
us10y = _yield_delta(prices["us10y"], edate)
eu10y = _yield_delta(prices["eu10y"], edate)
us2y = _yield_delta(prices["us2y"], edate)
activ = _activation(model, us10y, eu10y, drift["post_pips"])
surprise_abs = round(
ev["actual_value"] - (ev.get("forecast_value") or ev.get("previous_value") or 0), 4
)
result = {
"event": {
"id": ev["id"],
"name": ev["event_name"],
"series_id": ev["series_id"],
"date": ev["event_date"],
"actual": ev["actual_value"],
"forecast": ev.get("forecast_value"),
"previous": ev.get("previous_value"),
"surprise": surprise_abs,
"surprise_pct": ev.get("surprise_pct"),
"direction": ev.get("surprise_direction", "neutral"),
"cfg": cfg,
},
"model": model,
"drift": drift,
"yields": {"actual_us10y": us10y, "actual_eu10y": eu10y, "actual_us2y": us2y},
"activation": activ,
"prices": prices,
"analyzed_at": datetime.utcnow().strftime("%Y-%m-%dT%H:%M:%SZ"),
}
conn.execute(
"INSERT OR REPLACE INTO causal_lab_cache (event_id, result_json, analyzed_at) VALUES (?,?,?)",
(event_id, json.dumps(result, default=str), result["analyzed_at"])
)
conn.commit()
conn.close()
return result
except HTTPException:
raise
except Exception as e:
logger.error(f"[causal_lab] analyze {event_id}: {e}")
raise HTTPException(500, str(e))
@router.get("/api/causal-lab/summary")
def get_summary():
"""Agrégat de toutes les analyses — taux d'activation par nœud et par série."""
try:
from services.database import get_conn
conn = get_conn()
_ensure_cache(conn)
rows = conn.execute(
"SELECT result_json FROM causal_lab_cache WHERE result_json IS NOT NULL"
).fetchall()
conn.close()
if not rows:
return {"n": 0, "by_series": {}, "by_node": {}, "drift_dist": []}
by_series: dict = {}
node_stats: dict = {"us_10y": [0, 0], "eu_10y": [0, 0], "eurusd": [0, 0]}
drift_vals: list = []
for row in rows:
try:
d = json.loads(row["result_json"])
sid = d["event"]["series_id"]
name = d["event"]["name"]
if sid not in by_series:
by_series[sid] = {
"name": name, "n": 0,
"correct": 0, "total_nodes": 0,
"_preds": [], "_actuals": [],
}
s = by_series[sid]
s["n"] += 1
s["correct"] += d["activation"].get("correct", 0)
s["total_nodes"] += d["activation"].get("total", 0)
s["_preds"].append(d["model"].get("total_pips", 0))
if d["drift"].get("post_pips") is not None:
s["_actuals"].append(d["drift"]["post_pips"])
for nk, nv in d["activation"].get("nodes", {}).items():
if nk in node_stats:
node_stats[nk][1] += 1
if nv.get("status") == "correct":
node_stats[nk][0] += 1
if d["drift"].get("drift_ratio") is not None:
drift_vals.append(d["drift"]["drift_ratio"])
except Exception:
pass
by_node = {
k: {"correct": v[0], "total": v[1],
"rate": round(v[0] / v[1], 2) if v[1] > 0 else None}
for k, v in node_stats.items()
}
for s in by_series.values():
s["activation_rate"] = round(s["correct"] / s["total_nodes"], 2) if s["total_nodes"] else None
s["avg_pred_pips"] = round(sum(s["_preds"]) / len(s["_preds"])) if s["_preds"] else None
s["avg_actual_pips"] = round(sum(s["_actuals"]) / len(s["_actuals"])) if s["_actuals"] else None
del s["_preds"]
del s["_actuals"]
return {
"n": len(rows),
"by_series": by_series,
"by_node": by_node,
"drift_dist": drift_vals,
}
except Exception as e:
logger.error(f"[causal_lab] summary: {e}")
raise HTTPException(500, str(e))