rnaseek-full / efficiency_figure2 /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
6.42 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 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 / (
"qwen_regression_ckpt/"
"clean_cosine_restart_besthp_preview_fixed-wd-0.9_reproduce/"
"checkpoint-304419"
)
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.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))
head_state = {k.removeprefix("regression_head."): v for k, v in state.items()}
else:
head_state = torch.load(model_dir / "regression_head.pt", map_location="cpu")
if "net.2.weight" in head_state and "net.1.weight" not in head_state:
head_state["net.1.weight"] = head_state.pop("net.2.weight")
head_state["net.1.bias"] = head_state.pop("net.2.bias")
return head_state
def metrics(labels, preds):
labels = np.asarray(labels, dtype=float)
preds = np.asarray(preds, dtype=float)
out = {
"mse": float(mean_squared_error(labels, preds)),
"mae": float(mean_absolute_error(labels, preds)),
"r2": float(r2_score(labels, preds)),
}
if labels.std() > 1e-8 and preds.std() > 1e-8:
out["pearson_r"] = float(pearsonr(labels, preds)[0])
out["spearman_r"] = float(spearmanr(labels, preds)[0])
else:
out["pearson_r"] = 0.0
out["spearman_r"] = 0.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=None)
parser.add_argument("--valid-json", type=Path, default=SCRIPT_DIR / "evenBetterDataFolded-vl.json")
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)
parser.add_argument("--limit", type=int, default=None)
args = parser.parse_args()
if args.device == "cuda" and not torch.cuda.is_available():
raise RuntimeError("CUDA requested but not available.")
if args.tokenizer_dir is None:
args.tokenizer_dir = args.model_dir
args.out_dir.mkdir(parents=True, exist_ok=True)
device = torch.device(args.device)
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_dir, use_fast=True)
tokenizer.pad_token = tokenizer.eos_token
collator = DataCollatorWithPadding(tokenizer=tokenizer, pad_to_multiple_of=8, return_tensors="pt")
model = RegressionModel(args.model_dir, device)
with args.valid_json.open("r", encoding="utf-8") as handle:
valid = json.load(handle)
texts = list(valid.keys())
labels = np.asarray([float(valid[text]) for text in texts], dtype=float)
if args.limit is not None:
texts = texts[: args.limit]
labels = labels[: args.limit]
preds = []
starts = range(0, len(texts), args.batch_size)
iterator = tqdm(starts, desc="Evaluating efficiency valid") 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, add_special_tokens=True) 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)
np.savetxt(
args.out_dir / "valid_predictions.tsv",
np.column_stack([labels, preds]),
delimiter="\t",
header="label\tprediction",
comments="",
)
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"Efficiency validation r={report['pearson_r']:.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()