waldburger/farood-pilot / embed_job.py
waldburger's picture
download
raw
6.61 kB
#!/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.