rnaseek-full / validate_release.py
schen647's picture
Update release documentation and model entry points
db88907 verified
Raw History Blame Contribute Delete
9.11 kB
"""Validate published regression checkpoints against their fixed evaluation data.
Paths supplied on the command line are relative to the repository root. Output
directories must be new, so validation cannot overwrite archived evidence.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
ROOT = Path(__file__).resolve().parent
REGRESSION_BASE = "RNASeek-regression-pretrained"
DEFAULTS = {
"efficiency": {
"checkpoint": "efficiency_figure2/qwen_regression_ckpt/clean_cosine_restart_besthp_preview_fixed-wd-0.9_reproduce/checkpoint-304419",
"data": "efficiency_figure2/evenBetterDataFolded-vl.json",
"metric": "r2",
"target": 0.389,
},
"stability": {
"checkpoint": "regression_stability_functionalviral/rnaseek_abs20_tuning/hpcc_best_wave29_copyback/best_model",
"data": "regression_stability_functionalviral/rnaseek_abs20_tuning/hpcc_best_wave29_copyback/valid_predictions.tsv",
"metric": "spearman_r",
"target": 0.342,
},
}
def local_path(value):
path = Path(value).expanduser()
return path.resolve() if path.is_absolute() else (ROOT / path).resolve()
def load_regression_checkpoint(checkpoint, tokenizer_path=None, *, initialize_head=False):
"""Strictly load a predictor, or explicitly initialize a training-only head."""
import torch
from torch import nn
from safetensors.torch import load_file
from transformers import AutoConfig, AutoModel, AutoTokenizer
checkpoint = local_path(checkpoint)
config = AutoConfig.from_pretrained(checkpoint, local_files_only=True)
tokenizer = AutoTokenizer.from_pretrained(
local_path(tokenizer_path) if tokenizer_path else checkpoint,
local_files_only=True,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
# Strict tensor loading supports both archived layouts: a standalone
# backbone plus head, or a single state dict prefixed with backbone.
index = checkpoint / "model.safetensors.index.json"
shards = sorted(set(json.loads(index.read_text())["weight_map"].values())) if index.exists() else ["model.safetensors"]
state = {}
for shard in shards:
state.update(load_file(str(checkpoint / shard)))
if any(key.startswith("backbone.") for key in state):
head_state = {key.removeprefix("regression_head."): value for key, value in state.items() if key.startswith("regression_head.")}
backbone_state = {key.removeprefix("backbone."): value for key, value in state.items() if key.startswith("backbone.")}
unexpected = [key for key in state if not key.startswith(("backbone.", "regression_head."))]
if unexpected:
raise ValueError("Unexpected tensors in combined checkpoint")
elif initialize_head and not (checkpoint / "regression_head.pt").exists() and any(
name.endswith("ForCausalLM") for name in (config.architectures or [])
):
unexpected = [key for key in state if not key.startswith("model.") and key != "lm_head.weight"]
if unexpected:
raise ValueError("Unexpected tensors in pretrained causal-language-model checkpoint")
backbone_state = {key.removeprefix("model."): value for key, value in state.items() if key.startswith("model.")}
head_state = None
else:
backbone_state = state
if not (checkpoint / "regression_head.pt").exists():
raise ValueError("This checkpoint has no trained regression head. Use train_regression.py to fine-tune the pretrained base first.")
head_state = torch.load(checkpoint / "regression_head.pt", map_location="cpu", weights_only=True)
linear_index = 2 if head_state is None or "net.2.weight" in head_state else 1
head = nn.Module()
layers = [nn.LayerNorm(config.hidden_size)]
if linear_index == 2:
layers.append(nn.Dropout(0.1))
layers.append(nn.Linear(config.hidden_size, 1))
head.net = nn.Sequential(*layers)
if head_state is not None:
head.load_state_dict(head_state, strict=True)
backbone = AutoModel.from_config(config, attn_implementation="sdpa")
backbone.load_state_dict(backbone_state, strict=True)
del state, backbone_state, head_state
backbone.config.use_cache = False
return backbone, head, tokenizer
def evaluate(args):
import numpy as np
import pandas as pd
import torch
from safetensors.torch import load_file
from scipy.stats import pearsonr, spearmanr
from sklearn.metrics import r2_score
from torch import nn
from transformers import AutoConfig, AutoModel, AutoTokenizer
spec = DEFAULTS[args.task]
checkpoint = local_path(args.checkpoint or spec["checkpoint"])
data = local_path(args.data or spec["data"])
output = local_path(args.output)
if output.exists():
raise FileExistsError("Validation requires a new output directory.")
backbone, head, tokenizer = load_regression_checkpoint(checkpoint, args.tokenizer)
device = torch.device(args.device)
dtype = torch.bfloat16 if args.precision == "bfloat16" else torch.float32
backbone.to(device=device, dtype=dtype).eval()
head.to(device=device, dtype=dtype).eval()
if args.task == "efficiency":
examples = json.loads(data.read_text())
texts, labels = list(examples), list(examples.values())
add_special_tokens = True
else:
frame = pd.read_csv(data, sep="\t")
texts, labels = frame["text"].tolist(), frame["label"].tolist()
add_special_tokens = False
if args.limit:
texts, labels = texts[:args.limit], labels[:args.limit]
predictions = []
with torch.inference_mode(), torch.autocast(device_type=device.type, dtype=torch.bfloat16, enabled=args.precision == "autocast"):
for start in range(0, len(texts), args.batch_size):
batch = tokenizer(texts[start:start + args.batch_size], padding=True, pad_to_multiple_of=8, truncation=True, max_length=512, return_tensors="pt", add_special_tokens=add_special_tokens)
batch = {key: batch[key].to(device) for key in ("input_ids", "attention_mask")}
hidden = backbone(**batch).last_hidden_state
positions = torch.arange(hidden.shape[1], device=device)[None, :].expand(hidden.shape[:2])
last = positions.masked_fill(~batch["attention_mask"].bool(), -1).max(dim=1).values
if (last < 0).any():
raise ValueError("Empty tokenized sequence")
pooled = hidden[torch.arange(hidden.shape[0], device=device), last]
predictions.extend(head.net(pooled).flatten().float().cpu().tolist())
if start % (args.batch_size * 20) == 0:
print(f"{args.task}: {min(start + args.batch_size, len(texts))}/{len(texts)} evaluated", flush=True)
labels, predictions = np.asarray(labels), np.asarray(predictions)
if not np.isfinite(predictions).all():
raise ValueError("Non-finite model predictions")
metrics = {
"task": args.task, "n": len(labels), "precision": args.precision,
"r2": float(r2_score(labels, predictions)),
"pearson_r": float(pearsonr(labels, predictions).statistic),
"spearman_r": float(spearmanr(labels, predictions).statistic),
"full_validation": not bool(args.limit),
"target_metric": spec["metric"], "target": spec["target"],
}
metrics["allowed_decrease"] = args.allowed_decrease
metrics["target_passed"] = not args.limit and metrics[spec["metric"]] >= spec["target"] - args.allowed_decrease
output.mkdir(parents=True, exist_ok=False)
pd.DataFrame({"label": labels, "prediction": predictions}).to_csv(output / "predictions.tsv", sep="\t", index=False)
(output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n")
print(json.dumps(metrics, indent=2), flush=True)
if not args.limit and not metrics["target_passed"]:
raise SystemExit("Metric fell below the reproduction tolerance; inspect results before publishing.")
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("task", choices=DEFAULTS)
parser.add_argument("--checkpoint")
parser.add_argument("--tokenizer")
parser.add_argument("--data")
parser.add_argument("--output", required=True)
parser.add_argument("--device", default="cuda:0")
parser.add_argument("--precision", choices=["float32", "bfloat16", "autocast"], default="float32")
parser.add_argument("--batch-size", type=int, default=16)
parser.add_argument("--limit", type=int, default=0)
parser.add_argument("--allowed-decrease", type=float, default=0.01,
help="Maximum absolute decrease from the reference metric (default: 0.01).")
args = parser.parse_args()
if args.batch_size < 1 or args.limit < 0 or args.allowed_decrease < 0:
parser.error("batch-size must be positive; limit and allowed-decrease must be nonnegative")
evaluate(args)
if __name__ == "__main__":
main()