Download cas9/generate.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 2.93 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/generate.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/cas9/generate.py
-
curl -L -o generate.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/generate.py
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] | |