pCoMole / cas9 /generate.py
Maximilian Holsman
Claude Opus 5
Add Cas9 task
12fea4a
Raw History Blame Contribute Delete
2.93 kB
import torch
from transformers import EsmTokenizer
# Cas9 stack uses model/base_models.py (not reparam_models.py, which is the GFP/root variant).
from cas9.model.base_models import EditFlow, ProteinEditFlowModel, ReparameterizedProteinEditFlowModel
from cas9.model.utils import generate_from_x0, generate_from_x0_multi_edit
from cas9.logic import flow
def build_model_and_stuff(cfg, device):
"""
Rebuild exactly what train.py builds, but we won't set up lightning Trainer.
Returns:
editflow, source_distribution, tokenizer, pad_id, bos_id, eos_id, eps_id
"""
tokenizer = EsmTokenizer.from_pretrained("facebook/esm2_t33_650M_UR50D")
vocab_size = 24
source_distribution = flow.get_source_distribution(
source_distribution=cfg.flow.source_distribution,
vocab_size=vocab_size,
special_token_ids=[0, 1, 2, 3],
)
pad_id = 1
bos_id = 0
eos_id = 2
# Match train.py: reparameterized (8-value) vs plain (5-value) model per config flag.
if getattr(cfg.training, "reparameterize", False):
model = ReparameterizedProteinEditFlowModel(vocab_size=vocab_size, pad_id=pad_id, config=cfg.model)
else:
model = ProteinEditFlowModel(vocab_size=vocab_size, pad_id=pad_id, config=cfg.model)
eps_id = getattr(cfg.flow, "eps_id", -1)
path = flow.get_path(
scheduler_type=cfg.flow.scheduler_type,
exponent=cfg.flow.exponent,
eps_id=eps_id,
)
loss_fn = flow.get_loss_function(
loss_function=cfg.flow.loss_function,
path=path,
)
editflow = EditFlow(
model,
loss_fn,
path,
source_distribution,
pad_id,
bos_id,
eos_id,
cfg,
).to(device)
return editflow, source_distribution, tokenizer, pad_id, bos_id, eos_id, eps_id
def tokenize_input_str(input_str, cfg, tokenizer, bos_id, eos_id, pad_id, device):
# Signature matches the Cas9 pcomol_cas.py call site (cfg/pad_id accepted for
# compatibility; the operation is tokenize input + ensure BOS/EOS -> (1, L)).
toks = tokenizer(input_str, return_tensors='pt')
ids = toks["input_ids"][0].to(device)
if ids[0].item() != bos_id:
ids = torch.cat([torch.tensor([bos_id], device=device), ids], dim=0)
if ids[-1].item() != eos_id:
ids = torch.cat([ids, torch.tensor([eos_id], device=device)], dim=0)
x0 = ids.unsqueeze(0) # (1, L)
return x0
def detokenize_output(x, tokenizer, bos_id, eos_id, pad_id):
"""Convert a single generated sequence (1, L) back to string."""
seq = x[0].tolist()
seq = [tok for tok in seq if tok != pad_id] # strip padding
if len(seq) > 0 and seq[0] == bos_id: # strip BOS
seq = seq[1:]
if len(seq) > 0 and seq[-1] == eos_id: # strip EOS
seq = seq[:-1]
return tokenizer.batch_decode([seq], skip_special_tokens=True)[0]