You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
fithia2/scripts/probe_lgbm_topn_vs_v13e.py

108 lines
4.4 KiB
Python

#!/usr/bin/env python3
"""Decisive A/B: LGBM top-N vs v13e top-N realized fwd_return_10d on v13e-eligible subset.
If LGBM top-N realized >= v13e top-N realized on test, integration is worth building.
Else Option A is dead.
"""
from __future__ import annotations
import sys
from pathlib import Path
import numpy as np
import pandas as pd
import pyarrow.parquet as pq
from lightgbm import LGBMClassifier
sys.path.insert(0, ".")
from libs.backtest.scoring import compute_return_max_long_score_v13e
NUMERIC = [
"reaction_day_return", "gap_size", "close_location", "volume_ratio_20d",
"document_quality_score", "parse_confidence_overall", "oneoff_penalty",
"market_cap_proxy", "avg_dollar_volume_20d",
"pre_event_volatility_20d", "pre_event_rsi_14", "pre_event_bb_position",
"pre_event_obv_slope_20d",
"macro_vix", "macro_hy_spread", "macro_t10y2y", "prior_event_fwd5d",
"lm_positive_pct", "lm_negative_pct", "lm_net_sentiment",
"earnings_surprise_pct",
"sue_lag_1_pct", "sue_lag_2_pct", "sue_lag_3_pct",
"sue_hist_mean_4q", "sue_hist_mean_8q", "sue_hist_mean_12q",
"sue_hist_pos_rate_4q", "sue_hist_pos_rate_12q",
"sue_hist_latest_pct", "sue_hist_streak_pos",
"peer_sector_event_count_365d", "peer_sector_surprise_median_365d",
"peer_sector_surprise_mean_365d", "peer_sector_surprise_pos_rate_365d",
"peer_relative_surprise_pct_365d",
"peer_sector_sue_hist_mean_4q_median_365d",
"peer_sector_sue_hist_mean_4q_mean_365d",
"peer_relative_sue_hist_mean_4q_365d",
"peer_sector_sue_hist_pos_rate_4q_mean_365d",
"pre_event_short_ratio", "pre_event_sector_momentum_20d",
"prior_catalyst_count_60d", "prior_catalyst_type_diversity_60d",
"price_vs_sma20", "pre_event_momentum_20d",
]
CAT = ["event_type", "event_direction", "guidance_status", "sector"]
def load(root: Path, split: str) -> pd.DataFrame:
schema = pq.read_schema(root / f"{split}.parquet")
cols = [c for c in NUMERIC + CAT + ["mfe_10d", "mae_10d", "fwd_return_10d"] if c in schema.names]
df = pq.read_table(root / f"{split}.parquet", columns=cols).to_pandas()
for c in NUMERIC:
if c not in df.columns:
df[c] = np.nan
df[c] = pd.to_numeric(df[c], errors="coerce")
for c in CAT:
if c not in df.columns:
df[c] = "missing"
df[c] = df[c].fillna("missing").astype("category")
df["target"] = ((df["mfe_10d"] - df["mae_10d"]) > 0.08).astype(int)
return df
def main():
root = Path("data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_ftb_fix_v2")
train = load(root, "train")
valid = load(root, "valid")
test = load(root, "test")
feats = NUMERIC + CAT
model = LGBMClassifier(
n_estimators=400, learning_rate=0.03, num_leaves=31,
min_child_samples=50, subsample=0.8, colsample_bytree=0.8,
objective="binary", random_state=42, verbosity=-1,
)
model.fit(train[feats], train["target"], categorical_feature=CAT)
for split_name, df in [("valid", valid), ("test", test)]:
df = df.copy()
df["lgbm"] = model.predict_proba(df[feats])[:, 1]
df["v13e"] = [compute_return_max_long_score_v13e(r) for r in df.to_dict("records")]
elig = df[df["v13e"] >= 0.42].copy()
n_trades = {"valid": 65, "test": 37}.get(split_name, 30)
N = min(n_trades, len(elig))
v13e_top = elig.nlargest(N, "v13e")
lgbm_top = elig.nlargest(N, "lgbm")
elig_mean = elig["fwd_return_10d"].mean()
v13e_mean = v13e_top["fwd_return_10d"].mean()
lgbm_mean = lgbm_top["fwd_return_10d"].mean()
v13e_hit = v13e_top["target"].mean()
lgbm_hit = lgbm_top["target"].mean()
print(f"== {split_name} (eligible n={len(elig)}, top-N={N}) ==")
print(f" eligible mean fwd10d: {elig_mean:+.4f}")
print(f" v13e top-N mean fwd10d: {v13e_mean:+.4f} hit={v13e_hit:.3f}")
print(f" lgbm top-N mean fwd10d: {lgbm_mean:+.4f} hit={lgbm_hit:.3f}")
delta = lgbm_mean - v13e_mean
verdict = "LGBM WINS" if delta > 0.005 else ("LGBM TIES" if abs(delta) < 0.005 else "LGBM LOSES")
print(f" delta (lgbm - v13e): {delta:+.4f} -> {verdict}")
# also overlap
v13e_set = set(v13e_top.index)
lgbm_set = set(lgbm_top.index)
overlap = len(v13e_set & lgbm_set)
print(f" index overlap: {overlap}/{N} ({100*overlap/N:.1f}%)")
print()
if __name__ == "__main__":
main()