""" 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()