rnaseek-full / ribozymegen-rl-figure7 /ribozyme_efficiency_api.py
schen647's picture
included pretraining from hpcc and exported dataset from ipynb; zipped all safetensors weights
83ddd7e
Raw
History Blame Contribute Delete
17.8 kB
#!/usr/bin/env python
# ribozyme_efficiency_api.py (drop-in replacement for your current API file)
#
# ALWAYS prints β€œwhat the model sees” to stdout for every /score call:
# - cleaned sequence length
# - annotated sequence length
# - prompt length
# - token count, attention length
# - head/tail token ids + token strings
# - decoded-from-ids preview
#
# Start:
# CUDA_VISIBLE_DEVICES=0 python ribozyme_efficiency_api.py
#
# Call:
# curl -X POST http://localhost:8003/score \
# -H "Content-Type: application/json" \
# -d '{"sequence":"AGC...RandomEnglishWordsCTCTA"}'
#
from __future__ import annotations
import os
import math
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
import RNA # ViennaRNA
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from safetensors.torch import load_file
from transformers import (
AutoConfig,
AutoModel,
AutoTokenizer,
PreTrainedModel,
)
# ──────────────────────────────────────────────────────────────────────────────
# Inline folding helper
# ──────────────────────────────────────────────────────────────────────────────
def annotate_decorated_sequence(decorated_seq: str, structure: str) -> str:
"""
Insert '(' before the left partner base and ')' after the right partner base,
preserving any non-base decorations in decorated_seq.
For plain ACGT input (no decorators), yields a nested-paren string like:
(A(G(C... )...)...)
"""
base_positions = []
pure_seq = ""
i = 0
while i < len(decorated_seq):
if decorated_seq[i] == "|":
j = decorated_seq.find("|", i + 1)
if j == -1:
base_positions.append(i)
pure_seq += decorated_seq[i]
i += 1
else:
i = j + 1
else:
base_positions.append(i)
pure_seq += decorated_seq[i]
i += 1
pos_to_base_idx = {pos: bi for bi, pos in enumerate(base_positions)}
stack = []
pairs = {}
for idx, s in enumerate(structure):
if s == "(":
stack.append(idx)
elif s == ")":
if not stack:
continue
j = stack.pop()
pairs[j] = idx
pairs[idx] = j
annotated = []
for pos, char in enumerate(decorated_seq):
if pos in pos_to_base_idx:
base_idx = pos_to_base_idx[pos]
if base_idx in pairs and base_idx < pairs[base_idx]:
annotated.append("(")
annotated.append(char)
if base_idx in pairs and base_idx > pairs[base_idx]:
annotated.append(")")
else:
annotated.append(char)
return "".join(annotated)
def sanitize_to_acgt(text: str) -> str:
text = (text or "").upper()
return "".join(ch for ch in text if ch in "ACGT")
def inline_fold(seq_acgt: str) -> str:
if not seq_acgt:
return ""
seq_rna = seq_acgt.replace("T", "U")
structure, _energy = RNA.fold(seq_rna)
return annotate_decorated_sequence(seq_acgt, structure)
def signed_log1p(x: float) -> float:
if x == 0.0:
return 0.0
return math.copysign(math.log1p(abs(float(x))), float(x))
# ──────────────────────────────────────────────────────────────────────────────
# QwenForRegression
# ──────────────────────────────────────────────────────────────────────────────
import os
import json
import warnings
from pathlib import Path
os.environ["CUDA_VISIBLE_DEVICES"] = "1"
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.tensorboard import SummaryWriter
from transformers import AutoConfig, AutoModelForCausalLM, PreTrainedModel, AutoTokenizer, AutoModel
try:
from accelerate.utils import unwrap_model as accelerate_unwrap
except ImportError:
accelerate_unwrap = lambda x: x
device = torch.device("cuda:0")
class LastTokenPooling(nn.Module):
"""
Pool using the last non-padded token. Supports left- or right-padding.
"""
def __init__(self):
super().__init__()
def forward(self, hidden_states, attention_mask=None):
# hidden_states: [B, T, H], attention_mask: [B, T]
if attention_mask is None:
return hidden_states[:, -1, :]
B, T, H = hidden_states.size()
# detect left-padding
if attention_mask[:, -1].sum().item() == B:
return hidden_states[:, -1, :]
# right-padding / variable lengths
seq_lens = attention_mask.sum(dim=1).long() - 1 # [B]
idx = seq_lens.view(B, 1, 1).expand(-1, 1, H) # [B,1,H]
return hidden_states.gather(1, idx).squeeze(1) # [B,H]
class OneLayerRegressionHead(nn.Module):
"""
Exactly one layernorm+linear, no residual-GELU block.
"""
def __init__(self, hidden_size):
super().__init__()
self.net = nn.Sequential(
nn.LayerNorm(hidden_size),
nn.Linear(hidden_size, 1),
)
def forward(self, x):
# x: [B, hidden_size]
return self.net(x).squeeze(-1) # [B]
class QwenForRegression(PreTrainedModel):
"""
A regression model that uses only the base transformer (no LM head) and last-token pooling.
"""
config_class = AutoConfig
base_model_prefix = "backbone"
def __init__(self, config, writer: SummaryWriter = None):
super().__init__(config)
# use base model without LM head to avoid unused lm_head parameters
self.backbone = AutoModel.from_config(config)
if getattr(config, "gradient_checkpointing", False):
self.backbone.gradient_checkpointing_enable()
hidden_size = config.hidden_size
self.pooler = LastTokenPooling()
self.regression_head = OneLayerRegressionHead(hidden_size)
self.writer = writer
self.step = 0
def supports_gradient_checkpointing(self) -> bool:
return True
def gradient_checkpointing_enable(self, **kwargs):
self.backbone.gradient_checkpointing_enable(**kwargs)
def gradient_checkpointing_disable(self, **kwargs):
self.backbone.gradient_checkpointing_disable(**kwargs)
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
outputs = self.backbone(
input_ids=input_ids,
attention_mask=attention_mask,
return_dict=True,
output_hidden_states=False,
)
hidden_states = outputs.last_hidden_state # [B,T,H]
pooled = self.pooler(hidden_states, attention_mask) # [B,H]
if self.writer is not None:
self.writer.add_scalar("pooled/mean", pooled.mean().item(), self.step)
self.writer.add_scalar("pooled/std", pooled.std().item(), self.step)
self.step += 1
logits = self.regression_head(pooled)
if labels is not None:
loss = F.mse_loss(logits, labels)
return {"loss": loss, "logits": logits}
return {"logits": logits}
def save_pretrained(
self,
save_directory: str,
state_dict=None,
accelerator=None,
**kwargs
):
model_to_save = self
if accelerator is not None:
model_to_save = accelerator.unwrap_model(self)
os.makedirs(save_directory, exist_ok=True)
model_to_save.config.save_pretrained(save_directory)
model_to_save.backbone.save_pretrained(
save_directory, state_dict=state_dict, **kwargs
)
head_sd = model_to_save.regression_head.state_dict()
for idx, layer in enumerate(model_to_save.regression_head.net):
if isinstance(layer, nn.Linear):
w_key = f"net.{idx}.weight"
if w_key in head_sd:
w = head_sd[w_key]
out_f, in_f = layer.out_features, layer.in_features
if w.dim() == 1 and w.numel() == out_f * in_f:
head_sd[w_key] = w.view(out_f, in_f)
torch.save(head_sd, os.path.join(save_directory, "regression_head.pt"))
@classmethod
def from_pretrained(cls, model_path, device="cpu", config=None, writer=None):
model_dir = Path(model_path)
config = config or AutoConfig.from_pretrained(model_dir)
fsdp_file = model_dir / "pytorch_model_fsdp.bin"
backbone_sd, head_sd = {}, {}
if fsdp_file.exists():
fsdp_sd = torch.load(fsdp_file, map_location=device)
# β€” strip the exact "backbone." prefix β€”
for k, v in fsdp_sd.items():
if k.startswith("backbone."):
new_k = k[len("backbone."):]
backbone_sd[new_k] = v
elif k.startswith("regression_head."):
new_k = k[len("regression_head."):]
head_sd[new_k] = v
# instantiate backbone from config
backbone = AutoModel.from_config(config)
# strict load: will now match
missing_b, unexpected_b = backbone.load_state_dict(backbone_sd, strict=True)
if missing_b or unexpected_b:
raise RuntimeError(
f"Backbone load mismatch.\n missing: {missing_b}\n unexpected: {unexpected_b}"
)
else:
# fallback to HF sharded .safetensors
backbone = AutoModel.from_pretrained(model_dir, device_map=None, config=config)
packed_head = model_dir / "regression_head.safetensors"
if packed_head.exists():
packed_sd = load_file(str(packed_head), device=str(device))
head_sd = {
k.removeprefix("regression_head."): v
for k, v in packed_sd.items()
if k.startswith("regression_head.")
}
else:
head_sd = torch.load(model_dir / "regression_head.pt", map_location=device)
if "net.2.weight" in head_sd and "net.1.weight" not in head_sd:
head_sd["net.1.weight"] = head_sd.pop("net.2.weight")
head_sd["net.1.bias"] = head_sd.pop("net.2.bias")
# build your full model
model = cls(config, writer=writer)
model.backbone = backbone.to(device)
if hasattr(model.config, "use_cache"):
model.config.use_cache = False
if hasattr(model.backbone, "config") and hasattr(model.backbone.config, "use_cache"):
model.backbone.config.use_cache = False
# load regression head strictly, too
missing_h, unexpected_h = model.regression_head.load_state_dict(head_sd, strict=True)
if missing_h or unexpected_h:
raise RuntimeError(
f"Head load mismatch.\n missing: {missing_h}\n unexpected: {unexpected_h}"
)
return model.to(device).eval()
# ──────────────────────────────────────────────────────────────────────────────
# ALWAYS-PRINT debug helper
# ──────────────────────────────────────────────────────────────────────────────
def print_model_view(tokenizer, prompt: str, enc: dict, cleaned: str, annotated: str, n: int = 120):
ids = enc["input_ids"][0].detach().cpu().tolist()
am = enc.get("attention_mask", None)
attn_len = int(am[0].sum().item()) if am is not None else None
toks = tokenizer.convert_ids_to_tokens(ids)
print("\n================ [MODEL VIEW] ================")
print(f"[LEN] prompt_chars={len(prompt)} cleaned_bases={len(cleaned)} annotated_chars={len(annotated)}")
print(f"[TOKENS] n_tokens={len(ids)} attn_len={attn_len}")
n = max(1, int(n))
head_ids = ids[:n]
head_toks = toks[:n]
tail_ids = ids[-n:] if len(ids) > n else []
tail_toks = toks[-n:] if len(toks) > n else []
print(f"\n[HEAD ids[:{n}]]")
print(head_ids)
print(f"[HEAD toks[:{n}]]")
print(head_toks)
if tail_ids:
print(f"\n[TAIL ids[-{n}:]]")
print(tail_ids)
print(f"[TAIL toks[-{n}:]]")
print(tail_toks)
decoded_full = tokenizer.decode(ids, skip_special_tokens=False)
print("\n[DECODED FROM IDS] (first 500 chars; newlines escaped)")
print(decoded_full[:500].replace("\n", "\\n"))
if len(decoded_full) > 500:
print("... (truncated)")
if len(decoded_full) < max(1, len(prompt) - 10):
print(f"\n[WARN] decoded_from_ids shorter than prompt; possible truncation/mismatch.")
print(f" prompt_len={len(prompt)} decoded_len={len(decoded_full)}")
print("==============================================\n", flush=True)
# ──────────────────────────────────────────────────────────────────────────────
# Model + tokenizer init
# ──────────────────────────────────────────────────────────────────────────────
DEFAULT_MODEL_PATH = (
Path(__file__).resolve().parents[1]
/ "efficiency_figure2"
/ "qwen_regression_ckpt"
/ "clean_cosine_restart_besthp_preview_fixed-wd-0.9_reproduce"
/ "checkpoint-304419"
)
MODEL_PATH = os.environ.get("RM_MODEL_PATH", str(DEFAULT_MODEL_PATH))
if not MODEL_PATH:
raise RuntimeError("Set RM_MODEL_PATH to your trained ribozyme regression checkpoint directory (checkpoint-XXXX).")
TOKENIZER_PATH = os.environ.get("RM_TOKENIZER_PATH", MODEL_PATH)
DEVICE = torch.device(os.environ.get("RM_DEVICE", "cuda:0" if torch.cuda.is_available() else "cpu"))
TASK = os.environ.get("RM_TASK", "predict_ribozyme_efficiency")
# how many tokens to print head/tail per request (can override via env)
DEBUG_N = int(os.environ.get("RM_DEBUG_N", "120"))
print(f"[INIT] Loading RM from: {MODEL_PATH} on device {DEVICE}")
rm_config = AutoConfig.from_pretrained(MODEL_PATH)
rm_model = QwenForRegression.from_pretrained(MODEL_PATH, device=DEVICE, config=rm_config)
print(f"[INIT] Loading tokenizer from: {TOKENIZER_PATH}")
tokenizer = AutoTokenizer.from_pretrained(
TOKENIZER_PATH,
use_fast=True,
trust_remote_code=True,
)
if tokenizer.pad_token_id is None and tokenizer.eos_token_id is not None:
tokenizer.pad_token_id = tokenizer.eos_token_id
try:
tokenizer.padding_side = "left"
except Exception:
pass
# ──────────────────────────────────────────────────────────────────────────────
# FastAPI app
# ──────────────────────────────────────────────────────────────────────────────
app = FastAPI(title="Ribozyme Efficiency Regression API")
class ScoreRequest(BaseModel):
sequence: str
apply_signed_log: bool = False # still supported, but debug is always printed
class ScoreResponse(BaseModel):
cleaned_sequence: str
annotated_sequence: str
prompt: str
score: float
@app.get("/health")
def health():
return {"status": "ok"}
@app.post("/score", response_model=ScoreResponse)
def score_endpoint(req: ScoreRequest):
cleaned = sanitize_to_acgt(req.sequence)
if not cleaned:
raise HTTPException(status_code=400, detail="Input contains no A/C/G/T bases after sanitization.")
annotated = inline_fold(cleaned)
if not annotated:
raise HTTPException(status_code=400, detail="Folding failed or produced empty annotated sequence.")
prompt = f"<s>{annotated}~${TASK}\n</s>"
enc = tokenizer(
prompt,
return_tensors="pt",
padding=False,
truncation=True,
add_special_tokens=False,
return_token_type_ids=False,
)
if "token_type_ids" in enc:
del enc["token_type_ids"]
enc = {k: v.to(DEVICE) for k, v in enc.items()}
# ALWAYS PRINT what the model sees
print_model_view(tokenizer, prompt, enc, cleaned, annotated, n=DEBUG_N)
with torch.inference_mode():
out = rm_model(
input_ids=enc["input_ids"],
attention_mask=enc.get("attention_mask", None),
)
val = float(out["logits"][0].item())
if req.apply_signed_log:
val = signed_log1p(val)
return ScoreResponse(
cleaned_sequence=cleaned,
annotated_sequence=annotated,
prompt=prompt,
score=val,
)
if __name__ == "__main__":
import uvicorn
uvicorn.run("ribozyme_efficiency_api:app", host="0.0.0.0", port=8013, reload=False)