rnaseek-full / regression_stability_functionalviral /validate_packed_regression_model.py
schen647's picture
included pretraining from hpcc and exported dataset from ipynb; zipped all safetensors weights
83ddd7e
Raw
History Blame Contribute Delete
7.05 kB
#!/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"<s>{seq}{PROMPT_SUFFIX}</s>"
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()