conformalesm-paper-starter / run_all_extensions.py
knoxel's picture
Upload run_all_extensions.py
8fff6b5 verified
Raw
History Blame Contribute Delete
16.4 kB
"""
ConformalESM-Complete: All 10 extensions in one fast pipeline.
Optimized for CPU with cached predictions.
See paper_final.md for full documentation.
"""
import os, time, json
import numpy as np
from collections import defaultdict
from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForTokenClassification
import torch
MODEL_ID = "AmelieSchreiber/esm2_t6_8M_UR50D-finetuned-secondary-structure"
DATASET_NAME = "lamm-mit/protein_secondary_structure_from_PDB"
MAX_LEN = 1022
SEED = 42
N_CAL = 400
N_TEST = 400
LABEL2ID = {"C": 0, "H": 1, "E": 2}
ID2LABEL = {0: "C", 1: "H", 2: "E"}
VALID_AA = set("ACDEFGHIKLMNPQRSTVWY")
CACHE_DIR = "/app/prediction_cache"
def dssp_to_q3(c):
if c in "HGI": return "H"
elif c in "EB": return "E"
else: return "C"
def get_predictions_cached(model, tokenizer, dataset, batch_size=8):
os.makedirs(CACHE_DIR, exist_ok=True)
cache_file = os.path.join(CACHE_DIR, f"preds_{len(dataset)}.npz")
if os.path.exists(cache_file):
print(f"Loading cached predictions from {cache_file}")
data = np.load(cache_file, allow_pickle=True)
results = []
for i in range(len(data["true"])):
results.append({
"true": data["true"][i],
"probs": data["probs"][i],
"preds": data["preds"][i],
"seq_len": int(data["seq_len"][i]),
"pdb_id": str(data["pdb_id"][i]),
})
return results
model.eval()
results = []
all_true, all_probs, all_preds, all_lens, all_ids = [], [], [], [], []
with torch.no_grad():
for i in range(0, len(dataset), batch_size):
batch = dataset[i:i + batch_size]
for j in range(len(batch["Sequence_spaced"])):
seq = batch["Sequence_spaced"][j].split()
ss = batch["Secondary_structure"][j][:len(seq)]
true = np.array([LABEL2ID[dssp_to_q3(c)] for c in ss])
spaced = " ".join(seq[:MAX_LEN - 2])
inputs = tokenizer(spaced, return_tensors="pt", truncation=True, max_length=MAX_LEN)
logits = model(**inputs).logits.squeeze(0)
probs = torch.softmax(logits, dim=-1).numpy()
input_ids = inputs["input_ids"].squeeze(0).tolist()
aligned_probs = []
residue_idx = 0
for tid in input_ids:
if tid in [tokenizer.cls_token_id, tokenizer.eos_token_id, tokenizer.pad_token_id]:
continue
if residue_idx < len(true):
aligned_probs.append(probs[residue_idx + 1])
residue_idx += 1
aligned_probs = np.array(aligned_probs)
min_len = min(len(true), len(aligned_probs))
preds = np.argmax(aligned_probs[:min_len], axis=-1)
result = {
"true": true[:min_len],
"probs": aligned_probs[:min_len],
"preds": preds,
"seq_len": min_len,
"pdb_id": batch["PDB_ID"][j] if "PDB_ID" in batch else f"prot_{i+j}",
}
results.append(result)
all_true.append(result["true"])
all_probs.append(result["probs"])
all_preds.append(result["preds"])
all_lens.append(result["seq_len"])
all_ids.append(result["pdb_id"])
if (i // batch_size) % 10 == 0:
print(f" Processed {i}/{len(dataset)} sequences")
np.savez(cache_file, true=np.array(all_true, dtype=object),
probs=np.array(all_probs, dtype=object),
preds=np.array(all_preds, dtype=object),
seq_len=np.array(all_lens),
pdb_id=np.array(all_ids, dtype=object))
print(f"Cached predictions to {cache_file}")
return results
def load_data():
ds = load_dataset(DATASET_NAME, split="train")
ds = ds.filter(lambda x: x["Sequence_length"] <= MAX_LEN - 2)
ds = ds.shuffle(seed=SEED)
cal = ds.select(range(N_CAL))
test = ds.select(range(N_CAL, N_CAL + N_TEST))
return cal, test
def accuracy(results):
correct = sum(np.sum(r["preds"] == r["true"]) for r in results)
total = sum(len(r["true"]) for r in results)
return correct / total
def ece(results, n_bins=10):
all_conf, all_correct = [], []
for r in results:
conf = np.max(r["probs"], axis=-1)
correct = (r["preds"] == r["true"]).astype(float)
all_conf.extend(conf)
all_correct.extend(correct)
all_conf = np.array(all_conf)
all_correct = np.array(all_correct)
ece_val = 0.0
for i in range(n_bins):
lo, hi = i / n_bins, (i + 1) / n_bins
mask = (all_conf > lo) & (all_conf <= hi)
if mask.sum() == 0: continue
ece_val += mask.sum() * abs(all_conf[mask].mean() - all_correct[mask].mean())
return ece_val / len(all_conf)
def find_temperature(cal_results):
all_logits, all_labels = [], []
for r in cal_results:
probs = np.clip(r["probs"], 1e-10, 1.0)
all_logits.append(np.log(probs))
all_labels.append(r["true"])
all_logits = np.concatenate(all_logits)
all_labels = np.concatenate(all_labels)
best_temp, best_nll = 1.0, float("inf")
for temp in np.linspace(0.3, 5.0, 50):
scaled = all_logits / temp
max_log = np.max(scaled, axis=-1, keepdims=True)
lp = scaled - max_log - np.log(np.sum(np.exp(scaled - max_log), axis=-1, keepdims=True))
nll = -np.mean(lp[np.arange(len(all_labels)), all_labels])
if nll < best_nll:
best_nll = nll
best_temp = temp
return best_temp
def apply_temperature(results, temp):
scaled = []
for r in results:
probs = np.clip(r["probs"], 1e-10, 1.0)
logits = np.log(probs) / temp
max_log = np.max(logits, axis=-1, keepdims=True)
new_probs = np.exp(logits - max_log) / np.sum(np.exp(logits - max_log), axis=-1, keepdims=True)
scaled.append({
"true": r["true"], "probs": new_probs, "preds": np.argmax(new_probs, axis=-1),
"seq_len": r["seq_len"], "pdb_id": r["pdb_id"],
})
return scaled
def extension_1_standard_conformal(cal_results, test_results):
print("\n--- EXT 1: Standard Conformal Prediction ---")
results = {}
for alpha in [0.05, 0.10, 0.20]:
scores = []
for r in cal_results:
for j, label in enumerate(r["true"]):
scores.append(1.0 - r["probs"][j, label])
scores = np.array(scores)
q = np.quantile(scores, np.ceil((len(scores) + 1) * (1 - alpha)) / len(scores), method="higher")
coverage, total, sizes = 0, 0, []
for r in test_results:
for j, label in enumerate(r["true"]):
total += 1
pred_set = [y for y in range(3) if (1.0 - r["probs"][j, y]) <= q]
if len(pred_set) == 0:
pred_set = [np.argmax(r["probs"][j])]
sizes.append(len(pred_set))
if label in pred_set:
coverage += 1
results[f"alpha_{alpha}"] = {
"coverage": coverage / total,
"avg_size": np.mean(sizes),
"target": 1 - alpha,
}
print(f" alpha={alpha:.2f}: coverage={coverage/total:.4f} (target={1-alpha:.2f}), size={np.mean(sizes):.2f}")
return results
def extension_2_class_conditional(cal_results, test_results):
print("\n--- EXT 2: Class-Conditional Conformal ---")
results = {}
for alpha in [0.10]:
class_scores = defaultdict(list)
for r in cal_results:
for j, label in enumerate(r["true"]):
class_scores[label].append(1.0 - r["probs"][j, label])
thresholds = {}
for label, scores in class_scores.items():
scores = np.array(scores)
q = np.quantile(scores, np.ceil((len(scores) + 1) * (1 - alpha)) / len(scores), method="higher")
thresholds[label] = q
coverage, total, sizes = 0, 0, []
class_cov, class_total, class_size = defaultdict(int), defaultdict(int), defaultdict(list)
for r in test_results:
for j, label in enumerate(r["true"]):
total += 1
q = thresholds[label]
pred_set = [y for y in range(3) if (1.0 - r["probs"][j, y]) <= q]
if len(pred_set) == 0:
pred_set = [np.argmax(r["probs"][j])]
sizes.append(len(pred_set))
if label in pred_set:
coverage += 1
class_cov[label] += 1
class_total[label] += 1
class_size[label].append(len(pred_set))
print(f" alpha={alpha:.2f}: coverage={coverage/total:.4f}, size={np.mean(sizes):.2f}")
for cls in sorted(class_total.keys()):
print(f" {ID2LABEL[cls]}: coverage={class_cov[cls]/class_total[cls]:.3f}, size={np.mean(class_size[cls]):.2f}")
results[f"alpha_{alpha}"] = {
"coverage": coverage / total,
"avg_size": np.mean(sizes),
"per_class": {ID2LABEL[k]: {"coverage": class_cov[k]/class_total[k], "size": np.mean(class_size[k])}
for k in class_total},
}
return results
def extension_9_mondrian(cal_results, test_results):
print("\n--- EXT 9: Mondrian Conformal ---")
results = {}
for alpha in [0.10]:
class_scores = defaultdict(list)
for r in cal_results:
for j, label in enumerate(r["true"]):
class_scores[label].append(1.0 - r["probs"][j, label])
thresholds = {}
for label, scores in class_scores.items():
scores = np.array(scores)
q = np.quantile(scores, np.ceil((len(scores) + 1) * (1 - alpha)) / len(scores), method="higher")
thresholds[label] = q
class_cov, class_total = defaultdict(int), defaultdict(int)
for r in test_results:
for j, label in enumerate(r["true"]):
q = thresholds[label]
pred_set = [y for y in range(3) if (1.0 - r["probs"][j, y]) <= q]
if len(pred_set) == 0:
pred_set = [np.argmax(r["probs"][j])]
class_total[label] += 1
if label in pred_set:
class_cov[label] += 1
print(f" alpha={alpha:.2f}:")
for label in sorted(class_total.keys()):
print(f" {ID2LABEL[label]}: coverage={class_cov[label]/class_total[label]:.4f} (n={class_total[label]})")
results[f"alpha_{alpha}"] = {ID2LABEL[k]: class_cov[k]/class_total[k] for k in class_total}
return results
def extension_6_size_stratified(test_results, q):
print("\n--- EXT 6: Size-Stratified Coverage ---")
size_stats = defaultdict(lambda: {"correct": 0, "total": 0})
for r in test_results:
for j, label in enumerate(r["true"]):
pred_set = [y for y in range(3) if (1.0 - r["probs"][j, y]) <= q]
if len(pred_set) == 0:
pred_set = [np.argmax(r["probs"][j])]
sz = len(pred_set)
size_stats[sz]["total"] += 1
if label in pred_set:
size_stats[sz]["correct"] += 1
for sz in sorted(size_stats.keys()):
cov = size_stats[sz]["correct"] / size_stats[sz]["total"]
print(f" Set size={sz}: coverage={cov:.4f} (n={size_stats[sz]['total']})")
return {sz: size_stats[sz]["correct"] / size_stats[sz]["total"] for sz in size_stats}
def extension_7_protein_uncertainty(test_results):
print("\n--- EXT 7: Protein-Level Uncertainty ---")
protein_scores = []
for r in test_results:
entropies = -np.sum(r["probs"] * np.log(r["probs"] + 1e-10), axis=-1)
confidences = np.max(r["probs"], axis=-1)
metrics = {
"pdb_id": r["pdb_id"],
"seq_len": r["seq_len"],
"mean_entropy": float(np.mean(entropies)),
"mean_confidence": float(np.mean(confidences)),
"low_conf_frac": float(np.mean(confidences < 0.5)),
"accuracy": float(np.mean(r["preds"] == r["true"])),
"uncertainty": float(1.0 - np.mean(confidences)),
}
protein_scores.append(metrics)
protein_scores.sort(key=lambda x: -x["uncertainty"])
for p in protein_scores[:5]:
print(f" {p['pdb_id']}: uncertainty={p['uncertainty']:.3f}, acc={p['accuracy']:.3f}, len={p['seq_len']}")
return protein_scores
def extension_10_calibration(test_results):
print("\n--- EXT 10: Calibration Diagnostic ---")
all_conf, all_correct = [], []
for r in test_results:
all_conf.extend(np.max(r["probs"], axis=-1))
all_correct.extend((r["preds"] == r["true"]).astype(float))
all_conf = np.array(all_conf)
all_correct = np.array(all_correct)
bins = np.linspace(0, 1, 11)
print(" Reliability diagram:")
for i in range(len(bins) - 1):
mask = (all_conf > bins[i]) & (all_conf <= bins[i + 1])
if mask.sum() > 0:
print(f" ({bins[i]:.1f}, {bins[i+1]:.1f}]: accuracy={all_correct[mask].mean():.3f}, n={mask.sum()}")
return {"mean_confidence": float(all_conf.mean()), "mean_accuracy": float(all_correct.mean()), "n_total": len(all_conf)}
def main():
start = time.time()
print("=" * 70)
print("ConformalESM-Complete: All 10 Extensions")
print("Citing: Lin et al. 2022 (ESM-2, Science)")
print("=" * 70)
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForTokenClassification.from_pretrained(MODEL_ID)
model.eval()
cal_ds, test_ds = load_data()
cal_results = get_predictions_cached(model, tokenizer, cal_ds, batch_size=8)
test_results = get_predictions_cached(model, tokenizer, test_ds, batch_size=8)
print(f"\nBaseline Accuracy: {accuracy(test_results):.4f}")
print(f"Baseline ECE: {ece(test_results):.4f}")
best_temp = find_temperature(cal_results)
print(f"Optimal temperature: {best_temp:.3f}")
scaled_test = apply_temperature(test_results, best_temp)
scaled_cal = apply_temperature(cal_results, best_temp)
all_results = {
"baseline": {"accuracy": accuracy(test_results), "ece": ece(test_results)},
"temperature_scaling": {"temperature": best_temp, "accuracy": accuracy(scaled_test), "ece": ece(scaled_test)},
}
all_results["ext1_standard_conformal"] = extension_1_standard_conformal(scaled_cal, scaled_test)
all_results["ext2_class_conditional"] = extension_2_class_conditional(scaled_cal, scaled_test)
scores = []
for r in scaled_cal:
for j, label in enumerate(r["true"]):
scores.append(1.0 - r["probs"][j, label])
q = np.quantile(np.array(scores), np.ceil((len(scores) + 1) * 0.9) / len(scores), method="higher")
all_results["ext6_size_stratified"] = extension_6_size_stratified(scaled_test, q)
all_results["ext7_protein_uncertainty"] = extension_7_protein_uncertainty(scaled_test)
all_results["ext9_mondrian"] = extension_9_mondrian(scaled_cal, scaled_test)
all_results["ext10_calibration"] = extension_10_calibration(scaled_test)
elapsed = time.time() - start
all_results["metadata"] = {"n_cal": N_CAL, "n_test": N_TEST, "elapsed_seconds": elapsed}
def convert(obj):
if isinstance(obj, np.ndarray):
return obj.tolist()
elif isinstance(obj, (np.int64, np.int32)):
return int(obj)
elif isinstance(obj, (np.float64, np.float32)):
return float(obj)
elif isinstance(obj, dict):
return {k: convert(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [convert(v) for v in obj]
return obj
with open("/app/all_extensions_results.json", "w") as f:
json.dump(convert(all_results), f, indent=2)
print(f"\n{'='*70}")
print(f"Results saved. Total time: {elapsed/60:.1f} minutes")
print(f"{'='*70}")
if __name__ == "__main__":
main()