Buckets:
| #!/usr/bin/env python | |
| """HF Job: embed far-OOD pools with ESM-2 t12_35M on GPU + run per-arm analysis. | |
| Reads /workspace/pools/*.csv, writes /workspace/out35/. Self-contained. | |
| """ | |
| import os, time, json | |
| import numpy as np | |
| import pandas as pd | |
| MODEL = "esm2_t12_35M_UR50D" | |
| MAX_TOKENS = 8192 # fp32 attention on 16GB T4: safe | |
| POOLS, OUT = "/workspace/pools", "/workspace/out35" | |
| os.makedirs(OUT, exist_ok=True) | |
| DATASETS = [ | |
| ("meltA_trainpool", "meltA_trainpool.csv"), | |
| ("meltA_A1_test", "meltA_A1_test.csv"), | |
| ("meltA_A2_test", "meltA_A2_test.csv"), | |
| ("meltA_A3_test", "meltA_A3_test.csv"), | |
| ("aav_B2_train", "aav_B2_train.csv"), | |
| ("aav_B2_test", "aav_B2_test.csv"), | |
| ("aav_B3_train", "aav_B3_train.csv"), | |
| ("aav_B3_test", "aav_B3_test.csv"), | |
| ] | |
| ARMS = { | |
| "A1_Saccharomyces": ("meltA_trainpool", "meltA_A1_test"), | |
| "A2_Ecoli": ("meltA_trainpool", "meltA_A2_test"), | |
| "A3_Thermus": ("meltA_trainpool", "meltA_A3_test"), | |
| "B2_seven_vs_many": ("aav_B2_train", "aav_B2_test"), | |
| "B3_des_mut": ("aav_B3_train", "aav_B3_test"), | |
| } | |
| def embed_all(): | |
| import torch, esm | |
| model, alphabet = getattr(esm.pretrained, MODEL)() | |
| model = model.eval().cuda() | |
| bc = alphabet.get_batch_converter() | |
| for tag, csv in DATASETS: | |
| cache = os.path.join(OUT, f"emb_{tag}_{MODEL}.npz") | |
| if os.path.exists(cache): | |
| print(f"[{tag}] cached, skip", flush=True); continue | |
| df = pd.read_csv(os.path.join(POOLS, csv)) | |
| seqs = df.sequence.tolist() | |
| order = np.argsort([len(s) for s in seqs]) | |
| ss = [seqs[i] for i in order] | |
| batches, cur, tok = [], [], 0 | |
| for i, s in enumerate(ss): | |
| t = len(s) + 2 | |
| if cur and tok + t > MAX_TOKENS: | |
| batches.append(cur); cur, tok = [], 0 | |
| cur.append(i); tok += t | |
| if cur: batches.append(cur) | |
| embs = np.zeros((len(seqs), model.embed_dim), dtype=np.float32) | |
| t0, done = time.time(), 0 | |
| with torch.no_grad(): | |
| for bidx in batches: | |
| chunk = [ss[i] for i in bidx] | |
| _, _, toks = bc([(f"p{j}", s) for j, s in enumerate(chunk)]) | |
| toks = toks.cuda() | |
| rep = model(toks, repr_layers=[model.num_layers])["representations"][model.num_layers] | |
| for j, s in enumerate(chunk): | |
| embs[order[bidx[j]]] = rep[j, 1:len(s) + 1].mean(0).cpu().numpy() | |
| del toks, rep | |
| done += len(chunk) | |
| if done % 1000 < 64: | |
| print(f"[{tag}] {done}/{len(seqs)} ({done/(time.time()-t0):.1f}/s)", flush=True) | |
| np.savez_compressed(cache, emb=embs) | |
| print(f"[{tag}] saved {len(seqs)}", flush=True) | |
| def load(tag): | |
| return np.load(os.path.join(OUT, f"emb_{tag}_{MODEL}.npz"))["emb"] | |
| def analyze(arm, tr_tag, te_tag, n_ens=10, k=10, seed=0): | |
| from sklearn.linear_model import Ridge, LinearRegression | |
| from sklearn.preprocessing import StandardScaler | |
| from sklearn.decomposition import PCA | |
| from sklearn.covariance import LedoitWolf | |
| from sklearn.metrics import roc_auc_score | |
| from scipy.stats import spearmanr, rankdata | |
| tr = pd.read_csv(os.path.join(POOLS, f"{tr_tag}.csv")) | |
| te = pd.read_csv(os.path.join(POOLS, f"{te_tag}.csv")) | |
| Xtr, Xte = load(tr_tag), load(te_tag) | |
| ytr, yte = tr.target.to_numpy(), te.target.to_numpy() | |
| sc = StandardScaler().fit(Xtr) | |
| Ztr, Zte = sc.transform(Xtr), sc.transform(Xte) | |
| rng = np.random.default_rng(seed) | |
| preds = np.zeros((n_ens, len(Zte))) | |
| for b in range(n_ens): | |
| idx = rng.choice(len(Ztr), len(Ztr), replace=True) | |
| preds[b] = Ridge(alpha=1.0).fit(Ztr[idx], ytr[idx]).predict(Zte) | |
| mu, var = preds.mean(0), preds.var(0, ddof=1) | |
| err = np.abs(yte - mu) | |
| Ztr_n = Ztr / np.linalg.norm(Ztr, axis=1, keepdims=True) | |
| Zte_n = Zte / np.linalg.norm(Zte, axis=1, keepdims=True) | |
| knn = 1 - np.sort(Zte_n @ Ztr_n.T, axis=1)[:, -k:].mean(1) | |
| pca = PCA(n_components=64, random_state=seed).fit(Ztr) | |
| mahal = LedoitWolf().fit(pca.transform(Ztr)).mahalanobis(pca.transform(Zte)) ** 0.5 | |
| idx = rng.permutation(len(Ztr)) | |
| c, t = idx[:2000], idx[2000:] | |
| m = Ridge(alpha=1.0).fit(Ztr[t], ytr[t]) | |
| q90 = np.quantile(np.abs(ytr[c] - m.predict(Ztr[c])), 0.9) | |
| cov = np.abs(yte - m.predict(Zte)) <= q90 | |
| bins = np.quantile(knn, [0, .25, .5, .75, 1.0]) | |
| bid = np.clip(np.digitize(knn, bins[1:-1]), 0, 3) | |
| def sp(a, b): return float(spearmanr(a, b).statistic) | |
| top = err >= np.quantile(err, 0.75) | |
| rv, rk, re_ = rankdata(var), rankdata(knn), rankdata(err) | |
| rk_r = rk - LinearRegression().fit(rv.reshape(-1, 1), rk).predict(rv.reshape(-1, 1)) | |
| re_r = re_ - LinearRegression().fit(rv.reshape(-1, 1), re_).predict(rv.reshape(-1, 1)) | |
| rv_r = rv - LinearRegression().fit(rk.reshape(-1, 1), rv).predict(rk.reshape(-1, 1)) | |
| re_r2 = re_ - LinearRegression().fit(rk.reshape(-1, 1), re_).predict(rk.reshape(-1, 1)) | |
| # shift detection: kNN of test vs IID train slice | |
| iid = rng.choice(len(Ztr_n), 2000, replace=False) | |
| knn_iid = 1 - np.sort(Ztr_n[iid] @ Ztr_n.T, axis=1)[:, -k:].mean(1) | |
| y_det = np.r_[np.zeros(len(knn_iid)), np.ones(len(knn))] | |
| det_auroc = float(roc_auc_score(y_det, np.r_[knn_iid, knn])) | |
| return { | |
| "arm": arm, "model": MODEL, "n_train": int(len(Ztr)), "n_test": int(len(Zte)), | |
| "mae": float(err.mean()), | |
| "spearman_err_var": sp(err, var), "spearman_err_knn": sp(err, knn), | |
| "spearman_err_mahal": sp(err, mahal), | |
| "auroc_top25_var": float(roc_auc_score(top, var)), | |
| "auroc_top25_knn": float(roc_auc_score(top, knn)), | |
| "auroc_top25_mahal": float(roc_auc_score(top, mahal)), | |
| "partial_knn_given_var": float(np.corrcoef(rk_r, re_r)[0, 1]), | |
| "partial_var_given_knn": float(np.corrcoef(rv_r, re_r2)[0, 1]), | |
| "conformal_cov_ood": float(cov.mean()), | |
| "coverage_by_knn_quartile": [float(cov[bid == i].mean()) for i in range(4)], | |
| "mae_by_knn_quartile": [float(err[bid == i].mean()) for i in range(4)], | |
| "shift_detection_auroc": det_auroc, | |
| "knn_median": float(np.median(knn)), | |
| } | |
| if __name__ == "__main__": | |
| print("embedding all datasets with", MODEL, flush=True) | |
| embed_all() | |
| results = [] | |
| for arm, (tr, te) in ARMS.items(): | |
| print("analyzing", arm, flush=True) | |
| results.append(analyze(arm, tr, te)) | |
| with open(os.path.join(OUT, "far_metrics_35m.json"), "w") as f: | |
| json.dump(results, f, indent=2) | |
| print(json.dumps(results, indent=2), flush=True) | |
| print("JOB DONE", flush=True) | |
Xet Storage Details
- Size:
- 6.61 kB
- Xet hash:
- a8138843d69fb1e08d90069bd1038d3f83623fe196f0c660c334fcc9065a9c44
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.