#!/usr/bin/env python3 import argparse import json from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import pandas as pd import torch from scipy.stats import pearsonr, spearmanr from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score from safetensors.torch import load_file from torch import nn from transformers import AutoConfig, AutoModel, AutoTokenizer from transformers.data.data_collator import DataCollatorWithPadding try: from tqdm.auto import tqdm except Exception: tqdm = None SCRIPT_DIR = Path(__file__).resolve().parent DEFAULT_MODEL_DIR = SCRIPT_DIR / "checkpoint-44040_best" PROMPT_SUFFIX = "~$predict_stability\n" class LastTokenPooling(nn.Module): def forward(self, hidden_states, attention_mask=None): if attention_mask is None: return hidden_states[:, -1, :] batch_size, _, hidden_size = hidden_states.shape if attention_mask[:, -1].sum().item() == batch_size: return hidden_states[:, -1, :] seq_lens = attention_mask.sum(dim=1).long() - 1 idx = seq_lens.view(batch_size, 1, 1).expand(-1, 1, hidden_size) return hidden_states.gather(1, idx).squeeze(1) class RegressionHead(nn.Module): def __init__(self, hidden_size): super().__init__() self.net = nn.Sequential( nn.LayerNorm(hidden_size), nn.Dropout(0.0), nn.Linear(hidden_size, 1), ) def forward(self, x): return self.net(x).squeeze(-1) class RegressionModel(nn.Module): def __init__(self, model_dir, device): super().__init__() config = AutoConfig.from_pretrained(model_dir) self.backbone = AutoModel.from_pretrained(model_dir, config=config, device_map=None) self.pooler = LastTokenPooling() self.regression_head = RegressionHead(config.hidden_size) self.regression_head.load_state_dict(load_head_state(model_dir), strict=True) self.to(device).eval() def forward(self, input_ids, attention_mask): outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask, return_dict=True, output_hidden_states=False, ) pooled = self.pooler(outputs.last_hidden_state, attention_mask) return self.regression_head(pooled) def load_head_state(model_dir): packed = model_dir / "regression_head.safetensors" if packed.exists(): state = load_file(str(packed)) return {k.removeprefix("regression_head."): v for k, v in state.items()} return torch.load(model_dir / "regression_head.pt", map_location="cpu") def make_prompt(seq): seq = str(seq).upper().replace("U", "T") seq = "".join(base for base in seq if base in "ACGT") return f"{seq}{PROMPT_SUFFIX}" def metrics(labels, preds): labels = np.asarray(labels, dtype=float) preds = np.asarray(preds, dtype=float) out = { "n": int(len(labels)), "mse": float(mean_squared_error(labels, preds)), "mae": float(mean_absolute_error(labels, preds)), "r2": float(r2_score(labels, preds)), "gt_mean": float(np.mean(labels)), "pred_mean": float(np.mean(preds)), "gt_std": float(np.std(labels)), "pred_std": float(np.std(preds)), } if labels.std() > 1e-8 and preds.std() > 1e-8: out["pearson"] = float(pearsonr(labels, preds)[0]) out["spearman"] = float(spearmanr(labels, preds)[0]) else: out["pearson"] = 0.0 out["spearman"] = 0.0 abs_labels = np.abs(labels) threshold = float(np.quantile(abs_labels, 0.80)) mask = abs_labels >= threshold if int(mask.sum()) >= 3: out["abs20_threshold"] = threshold out["abs20_n"] = int(mask.sum()) out["abs20_r2"] = float(r2_score(labels[mask], preds[mask])) out["abs20_pearson"] = float(pearsonr(labels[mask], preds[mask])[0]) out["abs20_spearman"] = float(spearmanr(labels[mask], preds[mask])[0]) return out def main(): parser = argparse.ArgumentParser() parser.add_argument("--model-dir", type=Path, default=DEFAULT_MODEL_DIR) parser.add_argument("--tokenizer-dir", type=Path, default=DEFAULT_MODEL_DIR) parser.add_argument("--data-tsv", type=Path, default=SCRIPT_DIR / "training_seq_score_extreme_weighted.tsv") parser.add_argument("--split", default="val") parser.add_argument("--out-dir", type=Path, default=SCRIPT_DIR / "packed_validation") parser.add_argument("--device", choices=["cpu", "cuda"], required=True) parser.add_argument("--batch-size", type=int, default=1) args = parser.parse_args() if args.device == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA requested but not available.") args.out_dir.mkdir(parents=True, exist_ok=True) device = torch.device(args.device) tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_dir, use_fast=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token collator = DataCollatorWithPadding(tokenizer=tokenizer, pad_to_multiple_of=8, return_tensors="pt") model = RegressionModel(args.model_dir, device) df = pd.read_csv(args.data_tsv, sep="\t") df["split"] = df["split"].astype(str).str.lower() valid = df[df["split"] == args.split.lower()].copy() if len(valid) == 0: raise ValueError(f"No rows found for split {args.split!r}") texts = [make_prompt(seq) for seq in valid["seq"].tolist()] labels = valid["score"].astype(float).to_numpy() preds = [] starts = range(0, len(texts), args.batch_size) iterator = tqdm(starts, desc=f"Evaluating stability {args.split}") if tqdm else starts with torch.no_grad(): for start in iterator: batch_texts = texts[start : start + args.batch_size] encoded = [ tokenizer(text, truncation=True, max_length=512, add_special_tokens=False) for text in batch_texts ] batch = {k: v.to(device) for k, v in collator(encoded).items()} preds.extend(model(**batch).detach().cpu().float().numpy().tolist()) preds = np.asarray(preds, dtype=float) report = metrics(labels, preds) out = valid[["element_id", "seq", "score", "split"]].copy() out["prediction"] = preds out.to_csv(args.out_dir / "valid_predictions.tsv", sep="\t", index=False) with (args.out_dir / "valid_metrics.json").open("w", encoding="utf-8") as handle: json.dump(report, handle, indent=2) plt.figure(figsize=(5, 5), dpi=180) plt.scatter(labels, preds, s=7, alpha=0.35) plt.xlabel("Validation label") plt.ylabel("Prediction") plt.title(f"Stability validation r={report['pearson']:.3f}") plt.tight_layout() plt.savefig(args.out_dir / "valid_scatter.png") print(json.dumps(report, indent=2)) print(f"Wrote {args.out_dir / 'valid_scatter.png'}") if __name__ == "__main__": main()