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.

169 lines
6.7 KiB
Python

#!/usr/bin/env python3
"""Train LogReg for ml_v3 scorer on ftb_fix_v2 snapshot.
Emits Python tuple literals for FEATURES, MEDIANS, MEANS, SCALES, COEFFICIENTS
plus INTERCEPT - ready to paste into libs/backtest/scoring.py.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
import pandas as pd
import pyarrow.parquet as pq
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import roc_auc_score
CANDIDATE_FEATURES = [
"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_entropy_60d", "pre_event_gravitational_pull", "pre_event_hurst_60d",
"pre_event_market_temperature", "pre_event_ou_theta_60d",
"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",
# additional features available in ftb_fix_v2 we want to consider:
"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",
]
TARGET_AUX = ["mfe_10d", "mae_10d", "fwd_return_10d"]
def _load(snapshot_root: Path, split: str, present_cols: list[str]) -> pd.DataFrame:
path = snapshot_root / f"{split}.parquet"
schema = pq.read_schema(path)
cols = [c for c in present_cols + TARGET_AUX if c in schema.names]
df = pq.read_table(path, columns=cols).to_pandas()
for c in present_cols:
if c not in df.columns:
df[c] = np.nan
df[c] = pd.to_numeric(df[c], errors="coerce")
df["target"] = ((df["mfe_10d"] - df["mae_10d"]) > 0.08).astype(int)
return df
def _format_tuple(values, name, width=4) -> str:
lines = [f"{name} = ("]
for v in values:
if isinstance(v, str):
lines.append(f' "{v}",')
else:
lines.append(f" {v!r},")
lines.append(")")
return "\n".join(lines)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--snapshot-root",
default="data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_ftb_fix_v2")
ap.add_argument("--output", default="runs/ml_v3_logreg_constants.py")
ap.add_argument("--report", default="runs/ml_v3_logreg_metrics.json")
args = ap.parse_args()
root = Path(args.snapshot_root)
schema = pq.read_schema(root / "train.parquet")
available = [c for c in CANDIDATE_FEATURES if c in schema.names]
train = _load(root, "train", available)
valid = _load(root, "valid", available)
test = _load(root, "test", available)
# filter to features with at least 1 non-null value in train AND non-zero variance
coverage = train[available].notna().mean()
variance = train[available].var(skipna=True)
keep = [
c for c in available
if coverage.get(c, 0) > 0.05 and (variance.get(c, 0) or 0) > 1e-12
]
print(f"Available: {len(available)}, kept: {len(keep)}")
print(f"Dropped (low coverage / zero variance): {sorted(set(available) - set(keep))}")
medians = train[keep].median(skipna=True)
train_filled = train[keep].fillna(medians)
means = train_filled.mean()
stds = train_filled.std(ddof=0).replace(0, 1.0)
Xtr = (train_filled - means) / stds
ytr = train["target"].values
model = LogisticRegression(max_iter=5000, class_weight=None, C=1.0)
model.fit(Xtr, ytr)
coefs = model.coef_[0]
intercept = float(model.intercept_[0])
# eval
def _pred(df: pd.DataFrame) -> np.ndarray:
Xf = (df[keep].fillna(medians) - means) / stds
return model.predict_proba(Xf)[:, 1]
metrics = {}
for name, df in [("train", train), ("valid", valid), ("test", test)]:
prob = _pred(df)
try:
auc = roc_auc_score(df["target"], prob)
except ValueError:
auc = None
prob_series = pd.Series(prob, index=df.index)
target_series = df["target"]
try:
buckets = pd.qcut(prob_series, 5, labels=False, duplicates="drop")
top = target_series[buckets == buckets.max()].mean()
bot = target_series[buckets == buckets.min()].mean()
spread = float(top - bot)
except Exception:
top = bot = spread = None
metrics[name] = {
"auc": float(auc) if auc is not None else None,
"top_quintile": float(top) if top is not None else None,
"bottom_quintile": float(bot) if bot is not None else None,
"spread": spread,
"mean_pred": float(prob.mean()),
"std_pred": float(prob.std()),
"p10": float(np.percentile(prob, 10)),
"p50": float(np.percentile(prob, 50)),
"p90": float(np.percentile(prob, 90)),
}
print(f"{name}: auc={metrics[name]['auc']} spread={spread} pred_std={prob.std():.4f} pred_mean={prob.mean():.4f}")
# write outputs
out_py = Path(args.output)
out_py.parent.mkdir(parents=True, exist_ok=True)
out_py.write_text(
'"""Auto-generated by scripts/train_ml_v3_logreg.py - paste into scoring.py"""\n\n'
+ _format_tuple(keep, "_RETURN_MAX_LONG_ML_V3_FEATURES") + "\n"
+ _format_tuple([float(medians[c]) for c in keep], "_RETURN_MAX_LONG_ML_V3_MEDIANS") + "\n"
+ _format_tuple([float(means[c]) for c in keep], "_RETURN_MAX_LONG_ML_V3_MEANS") + "\n"
+ _format_tuple([float(stds[c]) for c in keep], "_RETURN_MAX_LONG_ML_V3_SCALES") + "\n"
+ _format_tuple([float(c) for c in coefs], "_RETURN_MAX_LONG_ML_V3_COEFFICIENTS") + "\n"
+ f"_RETURN_MAX_LONG_ML_V3_INTERCEPT = {intercept!r}\n"
)
Path(args.report).write_text(json.dumps({
"metrics": metrics,
"n_features": len(keep),
"features": keep,
"snapshot": str(root),
}, indent=2))
print(f"\nWrote {out_py}")
print(f"Wrote {args.report}")
if __name__ == "__main__":
main()