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]