Transformers
English
conformal-prediction
protein-language-models
uncertainty-quantification
esm-2
temperature-scaling
cpu
protein-structure
protein-engineering
Instructions to use knoxel/conformalesm-paper-starter with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use knoxel/conformalesm-paper-starter with Transformers:
# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("knoxel/conformalesm-paper-starter", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """ | |
| 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() | |