#!/usr/bin/env python
"""
retrain_gate_monotone.py — retrain the MEME-power gate per veg/datamonkey3#151.

Fixes:
  #1 monotonicity: HistGradientBoostingClassifier(monotonic_cst=[1,1,1]) — non-decreasing in all 3 live
     features. (Issue suggested GradientBoostingClassifier, but only Hist* supports monotonic_cst in sklearn.)
  #2 drop frac_p_defined: it was provably inert (0 splits) AND a post-hoc feature. Input is now [1,3].
  #3 median_pos_dist floor: floor missing/zero distances so "no branch lengths" cannot outrank short branches.

Verifies monotonicity with a VECTORIZED sweep (single predict_proba over a grid), reports CV AUC vs the old
0.887, and saves the new pickle. Re-export to ONNX via the existing skl2onnx path is a separate DM3 step.
"""
import csv, numpy as np, pickle, sys
from collections import defaultdict
from sklearn.ensemble import HistGradientBoostingClassifier
from sklearn.metrics import roc_auc_score
from sklearn.model_selection import GroupKFold

FT = "/storage/xl-data/users/sweaver/data/axomeme/finetune-data"
DIST_FLOOR = 0.001

def build():
    jobs = defaultdict(lambda: {"lrts": [], "ns": 0}); med = {}
    for r in csv.DictReader(open(f"{FT}/dm_sites.csv")):
        j = r["job_id"]; d = jobs[j]
        d["lrts"].append(float(r["lrt"])); d["ns"] = int(r.get("num_seqs", 0) or 0)
    for r in csv.DictReader(open(f"{FT}/um_manifest.csv")):
        try: med[r["job_id"]] = float(r["median_pos_dist"])
        except Exception: pass
    X, y, g = [], [], []
    for j, d in jobs.items():
        md = med.get(j, -1.0); md = max(md, DIST_FLOOR) if md >= 0 else DIST_FLOOR
        X.append([d["ns"], len(d["lrts"]), md]); y.append(int(np.array(d["lrts"]).max() >= 5.0)); g.append(j)
    return np.array(X, float), np.array(y), np.array(g)

def mk():
    return HistGradientBoostingClassifier(max_iter=200, max_depth=4, learning_rate=0.08,
                                          monotonic_cst=[1, 1, 1], random_state=0)

def verify_monotone(clf, X):
    """Vectorized: for each feature, sweep 200 values over a batch of held-fixed other-feature combos,
    assert the probability never decreases as the feature increases."""
    names = ["num_seqs", "num_sites", "median_pos_dist"]
    combos = [(ns, nsit, md) for ns in [10, 50, 150, 400] for nsit in [50, 300, 700] for md in [0.005, 0.02, 0.1]]
    total_viol = 0
    for fi, nm in enumerate(names):
        grid = np.linspace(X[:, fi].min(), X[:, fi].max(), 200)
        viol = 0
        for base in combos:
            B = np.tile(np.array(base, float), (len(grid), 1)); B[:, fi] = grid
            p = clf.predict_proba(B)[:, 1]           # one vectorized call per combo
            viol += int((np.diff(p) < -1e-9).sum())
        print(f"  {nm}: {viol} decreasing steps over {len(combos)} sweeps (want 0)")
        total_viol += viol
    # zero-distance inversion check
    p_floor = clf.predict_proba([[50, 200, DIST_FLOOR]])[0, 1]
    p_short = clf.predict_proba([[50, 200, 0.005]])[0, 1]
    print(f"  zero-dist: floor({DIST_FLOOR})={p_floor:.3f} vs short(0.005)={p_short:.3f}  "
          f"({'OK: floor<=short' if p_floor <= p_short + 1e-9 else 'STILL INVERTED'})")
    return total_viol

def main():
    X, y, g = build()
    print(f"data: {X.shape[0]} jobs, {y.mean()*100:.1f}% informative", flush=True)
    gkf = GroupKFold(5); oof = np.zeros(len(y))
    for tr, te in gkf.split(X, y, g):
        oof[te] = mk().fit(X[tr], y[tr]).predict_proba(X[te])[:, 1]
    auc = roc_auc_score(y, oof)
    print(f"MONOTONE gate CV ROC-AUC: {auc:.4f}  (old unconstrained: 0.887)", flush=True)
    clf = mk().fit(X, y)
    print("monotonicity verification:")
    viol = verify_monotone(clf, X)
    out = f"{FT}/meme_power_gate_monotone.pkl"
    pickle.dump({"model": clf, "features": ["num_seqs", "num_sites", "median_pos_dist"],
                 "monotonic_cst": [1, 1, 1], "dist_floor": DIST_FLOOR, "auc_cv": round(auc, 4),
                 "model_class": "HistGradientBoostingClassifier", "label": "max_LRT>=5",
                 "note": "monotone retrain per veg/datamonkey3#151: monotonic_cst[1,1,1], dropped frac_p_defined, dist floored"},
                open(out, "wb"))
    print(f"SAVED {out}  ({'CLEAN — 0 violations' if viol == 0 else f'WARNING {viol} violations'})", flush=True)

if __name__ == "__main__":
    main()
