#!/usr/bin/env python3
import json, sys, glob, warnings, argparse
import numpy as np, pandas as pd
from scipy.stats import norm
from sklearn.model_selection import KFold
from sklearn.preprocessing import StandardScaler, LabelEncoder
from sklearn.neighbors import NearestNeighbors
import torch, torch.nn as nn, torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset
warnings.filterwarnings("ignore"); torch.set_num_threads(4)

DATA="/mnt/user-data/uploads/tsunami_dataset.csv"
FEAT=["EQ_MAGNITUDE","log_depth","LATITUDE","LONGITUDE","abs_lat","region_code","has_landslide"]
TGT="TS_INTENSITY"; THR=3.0

def load():
    df=pd.read_csv(DATA)
    df=df[df["EVENT_VALIDITY"].isin(["Definite Tsunami","Probable Tsunami"])]
    df=df[df["CAUSE"].str.contains("Earthquake",na=False)]
    df=df.dropna(subset=["EQ_MAGNITUDE","TS_INTENSITY","LATITUDE","LONGITUDE"])
    df=df[df["EQ_MAGNITUDE"]>=5.0].reset_index(drop=True)
    df["EQ_DEPTH"]=df["EQ_DEPTH"].fillna(df["EQ_DEPTH"].median())
    df["log_depth"]=np.log(df["EQ_DEPTH"]+1); df["abs_lat"]=np.abs(df["LATITUDE"])
    df["has_landslide"]=df["CAUSE"].str.contains("Landslide",na=False).astype(float)
    df["region_code"]=LabelEncoder().fit_transform(df["REGION"].fillna("Unknown"))
    df["K_pred"]=2.0*df["EQ_MAGNITUDE"]-12.0; df["residual"]=df[TGT]-df["K_pred"]; return df

class RNet(nn.Module):
    def __init__(s,n,h=64,L=2,p=0.1):
        super().__init__(); ly=[]; d=n
        for _ in range(L): ly+=[nn.Linear(d,h),nn.GELU(),nn.Dropout(p)]; d=h
        ly.append(nn.Linear(d,1)); s.net=nn.Sequential(*ly); s.skip=nn.Linear(n,1)
    def forward(s,x): return s.net(x)+s.skip(x)

def train(m,Xtr,Ytr,Xv,Yv,ep=200,pat=20):
    opt=optim.AdamW(m.parameters(),lr=1e-3,weight_decay=1e-3)
    ld=DataLoader(TensorDataset(torch.FloatTensor(Xtr),torch.FloatTensor(Ytr.reshape(-1,1))),batch_size=64,shuffle=True)
    Xv_,Yv_=torch.FloatTensor(Xv),torch.FloatTensor(Yv.reshape(-1,1)); best,st,w=1e9,None,0
    for _ in range(ep):
        m.train()
        for xb,yb in ld: opt.zero_grad(); nn.MSELoss()(m(xb),yb).backward(); opt.step()
        m.eval()
        with torch.no_grad(): vl=nn.MSELoss()(m(Xv_),Yv_).item()
        if vl<best: best,st,w=vl,{k:v.clone() for k,v in m.state_dict().items()},0
        else:
            w+=1
            if w>=pat: break
    m.load_state_dict(st); m.eval(); return m

def detect(Y,p,th=THR):
    a,q=Y>=th,p>=th; tp=int((a&q).sum()); fp=int((~a&q).sum()); fn=int((a&~q).sum())
    pr=tp/(tp+fp+1e-9); rc=tp/(tp+fn+1e-9); return tp,fp,fn,2*pr*rc/(pr+rc+1e-9)

def gate(Xtr,Xte,k=5,pct=99):
    d,_=NearestNeighbors(n_neighbors=k+1).fit(Xtr).kneighbors(Xtr); tau=np.percentile(d[:,1:].mean(1),pct)
    dt,_=NearestNeighbors(n_neighbors=k).fit(Xtr).kneighbors(Xte); return np.minimum(1.0,(tau/(dt.mean(1)+1e-9))**2)

df=load(); N=len(df)
Y=df[TGT].values.astype(np.float32); Yr=df["residual"].values.astype(np.float32)
K=df["K_pred"].values.astype(np.float32); Xraw=df[FEAT].values.astype(np.float32)
reg=df["REGION"].fillna("Unknown").values; sb=float((Y-K).std())

def one_seed(seed):
    np.random.seed(seed); torch.manual_seed(seed)
    X=StandardScaler().fit_transform(Xraw); cv=np.zeros(N)
    for tri,vli in KFold(10,shuffle=True,random_state=seed).split(X):
        vs=int(0.85*len(tri)); m=train(RNet(len(FEAT)),X[tri[:vs]],Yr[tri[:vs]],X[tri[vs:]],Yr[tri[vs:]])
        with torch.no_grad(): cv[vli]=K[vli]+m(torch.FloatTensor(X[vli])).numpy().flatten()
    loro=np.zeros(N); lg=np.zeros(N); gv=np.zeros(N)
    for r in pd.unique(reg):
        te=np.where(reg==r)[0]; tr=np.where(reg!=r)[0]
        sc=StandardScaler().fit(Xraw[tr]); Xtr=sc.transform(Xraw[tr]); Xte=sc.transform(Xraw[te])
        pm=np.random.permutation(len(tr)); vs=int(0.85*len(tr))
        m=train(RNet(len(FEAT)),Xtr[pm[:vs]],Yr[tr][pm[:vs]],Xtr[pm[vs:]],Yr[tr][pm[vs:]])
        with torch.no_grad(): rr=m(torch.FloatTensor(Xte)).numpy().flatten()
        w=gate(Xtr,Xte); gv[te]=w; loro[te]=K[te]+rr; lg[te]=K[te]+w*rr
    def ab(idx,tgt):
        Xa=X[:,idx]; pr=np.zeros(N)
        for tri,vli in KFold(5,shuffle=True,random_state=seed).split(Xa):
            vs=int(0.85*len(tri)); mm=train(RNet(len(idx)),Xa[tri[:vs]],tgt[tri[:vs]],Xa[tri[vs:]],tgt[tri[vs:]])
            with torch.no_grad(): pr[vli]=mm(torch.FloatTensor(Xa[vli])).numpy().flatten()
        res=(Y-pr) if tgt is Y else (Yr-pr); return float((1-res.std()/sb)*100)
    res=dict(seed=seed,
        cv_sig=float((1-(Y-cv).std()/sb)*100), cv_det=detect(Y,cv),
        loro_sig=float((1-(Y-loro).std()/sb)*100), loro_gsig=float((1-(Y-lg).std()/sb)*100),
        loro_det=detect(Y,loro), loro_gdet=detect(Y,lg), gate_loro=float(gv.mean()),
        full=ab(list(range(len(FEAT))),Yr), ronly=ab(list(range(len(FEAT))),Y))
    np.save(f"cvpred_{seed}.npy", cv)
    json.dump(res, open(f"seed_{seed}.json","w"))
    print(f"seed {seed}: cv_sig={res['cv_sig']:.1f}  loro={res['loro_sig']:.1f}  gated={res['loro_gsig']:.1f}  "
          f"full={res['full']:.1f} ronly={res['ronly']:.1f} gate={res['gate_loro']:.2f}")

def aggregate():
    files=sorted(glob.glob("seed_*.json")); S=[json.load(open(f)) for f in files]
    def stat(key,sub=None):
        v=[ (s[key][sub] if sub is not None else s[key]) for s in S]; return float(np.mean(v)),float(np.std(v))
    cv=np.load("cvpred_1.npy"); sig=float((Y-cv).std())
    p=norm.cdf((cv-THR)/sig); pa=norm.cdf((K-THR)/sb); yb=(Y>=THR).astype(float)
    out=dict(nseeds=len(S),
        cv_sig=stat("cv_sig"),
        cv_tp=stat("cv_det",0), cv_fp=stat("cv_det",1), cv_f1=stat("cv_det",3),
        loro_sig=stat("loro_sig"), loro_gsig=stat("loro_gsig"),
        loro_fp=stat("loro_det",1), loro_f1=stat("loro_det",3),
        loro_gfp=stat("loro_gdet",1), loro_gf1=stat("loro_gdet",3),
        gate_loro=stat("gate_loro"), full=stat("full"), ronly=stat("ronly"),
        physics_contrib=stat("full")[0]-stat("ronly")[0],
        brier_abe=float(np.mean((pa-yb)**2)), brier_pir=float(np.mean((p-yb)**2)))
    json.dump(out,open("extended_results.json","w"),indent=2)
    print(json.dumps(out,indent=2))

if __name__=="__main__":
    ap=argparse.ArgumentParser(); ap.add_argument("--seed",type=int); ap.add_argument("--aggregate",action="store_true")
    a=ap.parse_args()
    if a.aggregate: aggregate()
    else: one_seed(a.seed)
