Download cas9/pcomole.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 140 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/pcomole.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/cas9/pcomole.py
-
curl -L -o pcomole.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/pcomole.py
140 kB
| # NOTE: imported first, before the flow_matching package lands on sys.path during the imports | |
| # below (its flow_matching/utils subpackage otherwise shadows this local editflows/utils package). | |
| from cas9.pam_detector_hmm import HMMSCAN, Cas9PIMasker # noqa: E402 | |
| import argparse | |
| from typing import List, Callable, Optional, Tuple, Dict, Any | |
| import torch | |
| import torch.nn.functional as F | |
| import yaml | |
| from easydict import EasyDict as edict | |
| from tqdm import tqdm | |
| import time | |
| import math | |
| from cas9.generate import build_model_and_stuff, tokenize_input_str, detokenize_output | |
| from cas9.objectives import Cas9Classification, DeletionCount, PAMMatching | |
| from cas9.constraints import Cas9DomainCompleteness, ProteinLength, TargetLength, MinTargetLength, MaxTargetLength, PAMMatchingConstraint, Cas9ScoreThreshold, PAMMatchingProbabilityThreshold | |
| # PAM/PI-domain detector + mask builder (moved to top of file to avoid flow_matching/utils shadowing) | |
| # from utils.pam_domain_detector.cas9_pam_detector_hmm import HMMSCAN, Cas9PIMasker | |
| # ---- Opt-in instrumentation (runtime breakdown + rollout-failure counts) ---- | |
| # Inert unless main() enables it via --instrument (INSTR stays None otherwise), so default | |
| # generation behavior is unchanged. Timers use cuda.synchronize for GPU-accurate wall time. | |
| from collections import defaultdict | |
| INSTR = None | |
| def _instr_reset(): | |
| global INSTR | |
| INSTR = {"timers": defaultdict(float), "counts": defaultdict(int)} | |
| def _tic(): | |
| if INSTR is None: | |
| return None | |
| if torch.cuda.is_available(): | |
| torch.cuda.synchronize() | |
| return time.time() | |
| def _toc(name, t0): | |
| if t0 is None: | |
| return | |
| if torch.cuda.is_available(): | |
| torch.cuda.synchronize() | |
| INSTR["timers"][name] += time.time() - t0 | |
| def _count(name, k=1): | |
| if INSTR is not None: | |
| INSTR["counts"][name] += int(k) | |
| import pdb | |
| import warnings | |
| warnings.filterwarnings("ignore", category=FutureWarning) | |
| from transformers.utils import logging | |
| logging.set_verbosity_error() | |
| # --------------------------------------------------------------------------- | |
| # small utilities | |
| # --------------------------------------------------------------------------- | |
| class PAMDomainWrapper: | |
| """ | |
| Wrapper for PAMMatching that extracts PAM domain region before PAM prediction. | |
| When PID_PAM_prediction is enabled, this wrapper extracts only the detected | |
| PAM domain region from sequences before passing them to the underlying PAMMatching object. | |
| """ | |
| def __init__(self, pam_matching_obj, pam_domain_interval): | |
| """ | |
| Args: | |
| pam_matching_obj: PAMMatching object to wrap | |
| pam_domain_interval: Tuple (start, end) of 1-based coordinates for PAM domain region | |
| """ | |
| self._pam_matching_obj = pam_matching_obj | |
| self._pam_domain_interval = pam_domain_interval | |
| start_1based, end_1based = pam_domain_interval | |
| self._start_0based = start_1based - 1 # Convert to 0-based for slicing | |
| self._end_0based = end_1based | |
| def _extract_domain_region(self, protein_seqs): | |
| """Extract PAM domain region from each sequence.""" | |
| domain_seqs = [] | |
| for seq in protein_seqs: | |
| seq_clean = seq.replace(" ", "").strip() | |
| # Extract domain region (handle cases where sequence might be shorter) | |
| if len(seq_clean) < self._end_0based: | |
| # If sequence is shorter than expected domain end, use what's available | |
| # This can happen during generation when sequences are being edited | |
| domain_seq = seq_clean[self._start_0based:] if len(seq_clean) > self._start_0based else seq_clean | |
| else: | |
| domain_seq = seq_clean[self._start_0based:self._end_0based] | |
| domain_seqs.append(domain_seq) | |
| return domain_seqs | |
| def predict_pam(self, protein_seqs, min_confidence=None): | |
| """Extract domain region and predict PAM.""" | |
| domain_seqs = self._extract_domain_region(protein_seqs) | |
| return self._pam_matching_obj.predict_pam(domain_seqs, min_confidence=min_confidence) | |
| def get_score_for_pam(self, protein_seqs, pam_sequence, use_temperature_scaling=False): | |
| """Extract domain region and get score for PAM.""" | |
| domain_seqs = self._extract_domain_region(protein_seqs) | |
| return self._pam_matching_obj.get_score_for_pam(domain_seqs, pam_sequence, use_temperature_scaling=use_temperature_scaling) | |
| def __call__(self, protein_tokens, protein_seqs): | |
| """Extract domain region and compute objective score.""" | |
| domain_seqs = self._extract_domain_region(protein_seqs) | |
| return self._pam_matching_obj(protein_tokens, domain_seqs) | |
| def __getattr__(self, name): | |
| """Delegate all other attributes to the wrapped object.""" | |
| return getattr(self._pam_matching_obj, name) | |
| def extract_objective_vector(protein_seqs, objective_models, device, return_names=False): | |
| """ | |
| Extract objective vector from protein sequences. | |
| Args: | |
| protein_seqs: List of protein sequence strings | |
| objective_models: List of objective model callables | |
| device: torch device | |
| return_names: If True, also return list of objective names | |
| Returns: | |
| torch.Tensor of shape (B, m) where m is number of objectives | |
| If return_names=True, also returns list of objective names | |
| """ | |
| values = [] | |
| names = [] | |
| # Allow objective-free runs (base EditFlows sampling mode). | |
| if len(objective_models) == 0: | |
| empty = torch.empty((len(protein_seqs), 0), device=device, dtype=torch.float32) | |
| if return_names: | |
| return empty, names | |
| return empty | |
| for obj in objective_models: | |
| name, scores = obj(protein_tokens=None, protein_seqs=protein_seqs) # list of length B | |
| values.append(torch.tensor(scores, device=device, dtype=torch.float32)) | |
| if return_names: | |
| names.append(name) | |
| result = torch.stack(values, dim=1) # (B,m) | |
| if return_names: | |
| return result, names | |
| return result | |
| def compute_scores_print(protein_seqs, objective_models, constraint_models, device, return_scores=False): | |
| """ | |
| Compute and print scores for protein sequences. | |
| Args: | |
| protein_seqs: List of protein sequence strings | |
| objective_models: List of objective model callables | |
| constraint_models: List of constraint model callables | |
| device: torch device | |
| return_scores: Whether to return scores | |
| Returns: | |
| scores if return_scores=True, else None | |
| """ | |
| objective_scores = extract_objective_vector(protein_seqs, objective_models, device) # (B,m) | |
| constraint_scores = [] | |
| for constraint in constraint_models: | |
| # Handle constraints that take single sequences vs batches | |
| if hasattr(constraint, 'predict_batch'): | |
| # Use batch prediction if available (more efficient) | |
| constraint_score = constraint.predict_batch(protein_seqs) | |
| constraint_score = [int(c) for c in constraint_score] # Convert bool to int | |
| else: | |
| # Call for each sequence individually | |
| constraint_score = [constraint(None, seq) for seq in protein_seqs] | |
| constraint_scores.append(torch.tensor(constraint_score, device=device, dtype=torch.float32)) | |
| constraint_scores = torch.stack(constraint_scores, dim=1) # (B, n) | |
| scores = torch.concat([objective_scores, constraint_scores], dim=1) # (B, m+n) | |
| print(scores) | |
| # Print predicted PAM and PAM probability score if PAM matching objective is present | |
| for obj in objective_models: | |
| if hasattr(obj, 'predict_pam'): | |
| predicted_pams = obj.predict_pam(protein_seqs) | |
| # Get PAM probability scores for each predicted PAM (not the target PAM) | |
| # Use raw probabilities (no temperature scaling) for more interpretable results | |
| pam_prob_scores = [] | |
| for i, pam in enumerate(predicted_pams): | |
| # Compute score for this specific predicted PAM using raw probabilities | |
| score = obj.get_score_for_pam([protein_seqs[i]], pam, use_temperature_scaling=False)[0] | |
| pam_prob_scores.append(score) | |
| for i, (pam, prob_score) in enumerate(zip(predicted_pams, pam_prob_scores)): | |
| print(f"Predicted PAM: {pam} (target: {obj.target_pam}) | Predicted PAM probability (raw): {prob_score:.8f}") | |
| break # Only print once if multiple PAM objectives exist | |
| if return_scores: | |
| return scores | |
| # --------------------------------------------------------------------------- | |
| # edit utilities | |
| # --------------------------------------------------------------------------- | |
| def _sample_multiple_edits_batch( | |
| x: torch.Tensor, # (B, Lmax) padded | |
| lam_ins: torch.Tensor, # (B, Lmax) | |
| logits_ins: torch.Tensor, # (B, Lmax, V) | |
| lam_del: torch.Tensor, # (B, Lmax) | |
| lam_sub: torch.Tensor, # (B, Lmax) | |
| logits_sub: torch.Tensor, # (B, Lmax, V) | |
| pad_id: int, | |
| bos_id: int, | |
| eos_id: int, | |
| allowed_tokens: Optional[torch.Tensor] = None, # 1D LongTensor of vocab ids | |
| delta: float = 1.0, | |
| max_len_cap: Optional[int] = None, | |
| protected_mask: Optional[torch.Tensor] = None, # (B, Lmax) bool, True=no edits allowed (PAM-protected) | |
| pam_scale_edits: bool = False, # If True, scale up edit rates in PAM region instead of masking | |
| pam_edit_scale_factor: float = 10.0, # Scaling factor for ins/sub rates in PAM region | |
| deletion_rate_scale: float = 1500.0, # Scaling factor for deletion rate | |
| debug_context: Optional[str] = None, # Context label for debug output (e.g., "CANDIDATE", "ROLLOUT") | |
| zero_lam_ins: bool = False, # If True, set lam_ins to 0.0 before processing | |
| ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], dict]: | |
| """ | |
| Multi-edit small-step proposal: | |
| - per position i: total rate λ_i = λ_ins + λ_del + λ_sub (after masking invalid ops) | |
| - fire with p_i = 1 - exp(-delta * λ_i) (independently per position) | |
| - if fired: pick op ~ proportional to (λ_ins, λ_del, λ_sub) | |
| - if op is ins/sub: draw token from softmax(logits_{ins/sub}[i]) (with allowed_tokens masking) | |
| - apply all fired edits "simultaneously" using a left-to-right scan on the original tokens: | |
| del: skip token | |
| sub: replace token | |
| ins: insert *after* the token | |
| Returns: | |
| x_out: (B, Lout) padded | |
| base_rate: (B,) relative proposal weight (safe vs underflow): exp(sum_fired log_ratio) | |
| protected_mask_out: (B, Lout) bool or None, updated mask with same edits applied | |
| edit_stats: dict with keys 'num_ins', 'num_del', 'num_sub', 'total_edits' (per batch item) | |
| """ | |
| assert x.dim() == 2, f"x must be (B,Lmax), got {tuple(x.shape)}" | |
| device = x.device | |
| B, Lmax = x.shape | |
| V = logits_ins.shape[-1] | |
| eps = 1e-30 | |
| if allowed_tokens is not None: | |
| if not torch.is_tensor(allowed_tokens): | |
| allowed_tokens = torch.tensor(allowed_tokens, device=device, dtype=torch.long) | |
| else: | |
| allowed_tokens = allowed_tokens.to(device=device, dtype=torch.long) | |
| # masks | |
| nonpad = (x != pad_id) | |
| lengths = nonpad.sum(dim=1) # (B,) | |
| is_bos = (x == bos_id) | |
| is_eos = (x == eos_id) | |
| # Debug: print raw model outputs before any masking or modification | |
| context_prefix = f"[{debug_context}] " if debug_context else "" | |
| # Skip rollout debugging, skip extra info for candidates (only show edit stats) | |
| is_rollout = debug_context is not None and "ROLLOUT" in debug_context | |
| is_candidate = debug_context is not None and "CANDIDATE" in debug_context | |
| # Only print raw model outputs for non-rollout, non-candidate contexts (e.g., SELECTED) | |
| # if not is_rollout and not is_candidate and nonpad.any(): | |
| # avg_lam_del_raw = lam_del[nonpad].mean().item() | |
| # avg_lam_ins_raw = lam_ins[nonpad].mean().item() | |
| # avg_lam_sub_raw = lam_sub[nonpad].mean().item() | |
| # if avg_lam_ins_raw > 0: | |
| # ratio_raw = (avg_lam_del_raw / avg_lam_ins_raw) - 1.0 # Positive = lam_del larger, negative = lam_ins larger | |
| # print(f"{context_prefix}[DEBUG] Raw model outputs: avg_lam_del={avg_lam_del_raw:.6f}, avg_lam_ins={avg_lam_ins_raw:.6f}, avg_lam_sub={avg_lam_sub_raw:.6f}, ratio={ratio_raw:.4f} (lam_del is {ratio_raw*100:.1f}% {'larger' if ratio_raw > 0 else 'smaller'})") | |
| # else: | |
| # print(f"{context_prefix}[DEBUG] Raw model outputs: avg_lam_del={avg_lam_del_raw:.6f}, avg_lam_ins={avg_lam_ins_raw:.6f}, avg_lam_sub={avg_lam_sub_raw:.6f}") | |
| # mask rates on invalid positions (match your single-edit masking rules) | |
| if zero_lam_ins: | |
| lam_ins[:] = 0.0 | |
| ins_rate = lam_ins.clone() | |
| ins_rate = ins_rate.masked_fill(~nonpad, 0.0) | |
| ins_rate = ins_rate.masked_fill(is_eos, 0.0) # no insertion at eos | |
| del_rate = lam_del.clone() | |
| del_rate = del_rate.masked_fill(~nonpad, 0.0) | |
| del_rate = del_rate.masked_fill(is_bos | is_eos, 0.0) # no delete bos/eos | |
| sub_rate = lam_sub.clone() | |
| sub_rate = sub_rate.masked_fill(~nonpad, 0.0) | |
| sub_rate = sub_rate.masked_fill(is_bos | is_eos, 0.0) # no sub bos/eos | |
| # Only print shapes/averages during candidate generation, not during rollouts | |
| # if not is_rollout: | |
| # print(f"Shapes -- ins_rate: {ins_rate.shape}, del_rate: {del_rate.shape}, sub_rate: {sub_rate.shape}") | |
| # print(f"Averages -- ins_rate: {ins_rate.mean(dim=-1)}, del_rate: {del_rate.mean(dim=-1)}, sub_rate: {sub_rate.mean(dim=-1)}") | |
| # Apply protected mask: either block edits OR scale up ins/sub rates in PAM region | |
| if protected_mask is not None: | |
| if pam_scale_edits: | |
| # Scale up insertion and substitution rates in PAM region (but not deletion) | |
| ins_rate = torch.where(protected_mask, ins_rate * pam_edit_scale_factor, ins_rate) | |
| sub_rate = torch.where(protected_mask, sub_rate * pam_edit_scale_factor, sub_rate) | |
| # Deletion rate remains unchanged (not scaled) | |
| else: | |
| # Original behavior: block ALL edits (ins/del/sub) in PAM-protected regions | |
| ins_rate = ins_rate.masked_fill(protected_mask, 0.0) | |
| del_rate = del_rate.masked_fill(protected_mask, 0.0) | |
| sub_rate = sub_rate.masked_fill(protected_mask, 0.0) | |
| # if at cap, disallow insertions | |
| if max_len_cap is not None: | |
| at_cap = lengths >= max_len_cap | |
| if at_cap.any(): | |
| ins_rate = ins_rate.masked_fill(at_cap.unsqueeze(1), 0.0) | |
| lam_total = ins_rate + del_rate + sub_rate # (B, Lmax) | |
| # Debug: print average rates before amplification | |
| valid_mask = nonpad # (B, Lmax) | |
| # Skip rollout debugging, skip extra info for candidates (only show edit stats) | |
| # Only print before amplification for non-rollout, non-candidate contexts (e.g., SELECTED) | |
| # if not is_rollout and not is_candidate and valid_mask.any(): | |
| # avg_lam_del_before = del_rate[valid_mask].mean().item() | |
| # avg_lam_ins_before = ins_rate[valid_mask].mean().item() | |
| # avg_lam_sub_before = sub_rate[valid_mask].mean().item() | |
| # if avg_lam_ins_before > 0: | |
| # ratio = (avg_lam_del_before / avg_lam_ins_before) - 1.0 # Positive = lam_del larger, negative = lam_ins larger | |
| # print(f"{context_prefix}[DEBUG] Before amplification: avg_lam_del={avg_lam_del_before:.6f}, avg_lam_ins={avg_lam_ins_before:.6f}, avg_lam_sub={avg_lam_sub_before:.6f}, ratio={ratio:.4f}") | |
| # else: | |
| # print(f"{context_prefix}[DEBUG] Before amplification: avg_lam_del={avg_lam_del_before:.6f}, avg_lam_ins={avg_lam_ins_before:.6f}, avg_lam_sub={avg_lam_sub_before:.6f}") | |
| # note: you had this amplification; kept unchanged | |
| del_rate *= deletion_rate_scale | |
| # Debug: print average deletion rate after amplification | |
| # Skip rollout debugging, skip extra info for candidates (only show edit stats) | |
| # Only print after amplification for non-rollout, non-candidate contexts (e.g., SELECTED) | |
| # if not is_rollout and not is_candidate and valid_mask.any(): | |
| # avg_lam_del_after = del_rate[valid_mask].mean().item() | |
| # avg_lam_ins_after = ins_rate[valid_mask].mean().item() | |
| # avg_lam_sub_after = sub_rate[valid_mask].mean().item() | |
| # print(f"{context_prefix}[DEBUG] After amplification (1000x): avg_lam_del={avg_lam_del_after:.6f}, avg_lam_ins={avg_lam_ins_after:.6f}, avg_lam_sub={avg_lam_sub_after:.6f}") | |
| # fire prob: p = 1 - exp(-delta*lam_total) (use expm1 for stability) | |
| a = (delta * lam_total).clamp_min(0.0) | |
| p_fire = (-torch.expm1(-a)).masked_fill(~nonpad, 0.0) # (B, Lmax) | |
| fired = (torch.rand_like(p_fire) < p_fire) & (lam_total > 1e-12) & nonpad | |
| # op probs per fired position: proportional to rates | |
| rates3 = torch.stack([ins_rate, del_rate, sub_rate], dim=-1) # (B,Lmax,3) | |
| denom = lam_total.unsqueeze(-1).clamp_min(1e-12) | |
| op_probs = rates3 / denom # (B,Lmax,3) | |
| # sample op only where fired | |
| fired_flat = fired.view(-1) | |
| idx_fired = fired_flat.nonzero(as_tuple=True)[0] # (K,) | |
| op_idx_flat = torch.zeros((B * Lmax,), device=device, dtype=torch.long) # default 0 | |
| if idx_fired.numel() > 0: | |
| op_p = op_probs.view(-1, 3)[idx_fired] # (K,3) | |
| op_p = op_p / op_p.sum(dim=1, keepdim=True).clamp_min(1e-12) | |
| op_idx_flat[idx_fired] = torch.multinomial(op_p, 1).squeeze(1) # (K,) | |
| op_idx = op_idx_flat.view(B, Lmax) # 0=ins,1=del,2=sub | |
| ins_mask = fired & (op_idx == 0) | |
| del_mask = fired & (op_idx == 1) | |
| sub_mask = fired & (op_idx == 2) | |
| # Count edit types for statistics | |
| num_ins = ins_mask.sum().item() | |
| num_del = del_mask.sum().item() | |
| num_sub = sub_mask.sum().item() | |
| total_edits = num_ins + num_del + num_sub | |
| # Compute average rates post-amplification for each batch item | |
| avg_rates_per_batch = [] | |
| for b in range(B): | |
| valid_positions = nonpad[b] # (Lmax,) bool | |
| if valid_positions.any(): | |
| avg_ins = ins_rate[b][valid_positions].mean().item() | |
| avg_del = del_rate[b][valid_positions].mean().item() | |
| avg_sub = sub_rate[b][valid_positions].mean().item() | |
| avg_rates_per_batch.append((avg_ins, avg_del, avg_sub)) | |
| else: | |
| avg_rates_per_batch.append((0.0, 0.0, 0.0)) | |
| # Skip rollout debugging, but keep candidate debugging | |
| if not is_rollout and total_edits > 0: | |
| pct_ins = 100.0 * num_ins / total_edits | |
| pct_del = 100.0 * num_del / total_edits | |
| pct_sub = 100.0 * num_sub / total_edits | |
| # For candidates, we typically have B=1, so use first batch item | |
| avg_ins_rate, avg_del_rate, avg_sub_rate = avg_rates_per_batch[0] if avg_rates_per_batch else (0.0, 0.0, 0.0) | |
| print(f"{context_prefix}[EDIT STATS] ins={num_ins} ({pct_ins:.1f}%), del={num_del} ({pct_del:.1f}%), sub={num_sub} ({pct_sub:.1f}%) | avg_rates: ins={avg_ins_rate:.6f}, del={avg_del_rate:.6f}, sub={avg_sub_rate:.6f}") | |
| # helper: mask logits to allowed_tokens | |
| def _mask_logits_full(logits_2d: torch.Tensor) -> torch.Tensor: | |
| # logits_2d: (K, V) | |
| if allowed_tokens is None: | |
| return logits_2d | |
| add = torch.full_like(logits_2d, -1e9) | |
| add[:, allowed_tokens] = 0.0 | |
| return logits_2d + add | |
| # sample tokens for ins/sub at masked positions | |
| ins_tok = torch.full((B, Lmax), pad_id, device=device, dtype=torch.long) | |
| sub_tok = torch.full((B, Lmax), pad_id, device=device, dtype=torch.long) | |
| if ins_mask.any(): | |
| idx_ins = ins_mask.view(-1).nonzero(as_tuple=True)[0] | |
| logits_sel = logits_ins.view(-1, V)[idx_ins] | |
| logits_sel = _mask_logits_full(logits_sel) | |
| q = F.softmax(logits_sel, dim=-1) | |
| samp = torch.multinomial(q, 1).squeeze(1) | |
| ins_tok.view(-1)[idx_ins] = samp | |
| if sub_mask.any(): | |
| idx_sub = sub_mask.view(-1).nonzero(as_tuple=True)[0] | |
| logits_sel = logits_sub.view(-1, V)[idx_sub] | |
| logits_sel = _mask_logits_full(logits_sel) | |
| q = F.softmax(logits_sel, dim=-1) | |
| samp = torch.multinomial(q, 1).squeeze(1) | |
| sub_tok.view(-1)[idx_sub] = samp | |
| # ------------------------- | |
| # base_rate: (B,) relative weight to avoid underflow | |
| # ------------------------- | |
| base_log = torch.zeros((B,), device=device, dtype=torch.float32) | |
| if idx_fired.numel() > 0: | |
| b_idx = (idx_fired // Lmax).to(torch.long) # (K,) | |
| op_choice = op_idx_flat[idx_fired].to(torch.long) # (K,) | |
| a_sel = a.view(-1)[idx_fired].to(torch.float32) # (K,) | |
| log_expm1 = torch.log(torch.expm1(a_sel).clamp_min(eps)) # (K,) | |
| op_p_sel = op_probs.view(-1, 3)[idx_fired].to(torch.float32) | |
| op_p_sel = op_p_sel / op_p_sel.sum(dim=1, keepdim=True).clamp_min(1e-12) | |
| op_prob_sel = op_p_sel.gather(1, op_choice.view(-1, 1)).squeeze(1).clamp_min(eps) | |
| log_op = torch.log(op_prob_sel) | |
| log_tok = torch.zeros_like(log_op) | |
| # token prob for ins | |
| ins_k = (op_choice == 0) | |
| if ins_k.any(): | |
| idx_ins_k = idx_fired[ins_k] | |
| tok_sel = ins_tok.view(-1)[idx_ins_k] | |
| logits_sel = logits_ins.view(-1, V)[idx_ins_k] | |
| logits_sel = _mask_logits_full(logits_sel) | |
| logq = F.log_softmax(logits_sel, dim=-1) | |
| log_tok[ins_k] = logq.gather(1, tok_sel.view(-1, 1)).squeeze(1) | |
| # token prob for sub | |
| sub_k = (op_choice == 2) | |
| if sub_k.any(): | |
| idx_sub_k = idx_fired[sub_k] | |
| tok_sel = sub_tok.view(-1)[idx_sub_k] | |
| logits_sel = logits_sub.view(-1, V)[idx_sub_k] | |
| logits_sel = _mask_logits_full(logits_sel) | |
| logq = F.log_softmax(logits_sel, dim=-1) | |
| log_tok[sub_k] = logq.gather(1, tok_sel.view(-1, 1)).squeeze(1) | |
| log_ratio = log_expm1 + log_op + log_tok | |
| base_log.scatter_add_(0, b_idx, log_ratio) | |
| base_rate = torch.exp(base_log).clamp_min(0.0) # (B,) | |
| # ------------------------- | |
| # apply edits to build new padded batch | |
| # ------------------------- | |
| new_seqs = [] | |
| new_lens = [] | |
| new_masks = [] # Track masks if provided | |
| for b in range(B): | |
| seq = x[b] | |
| valid = (seq != pad_id) | |
| tokens = seq[valid].tolist() | |
| Lb = len(tokens) | |
| # Extract mask for this batch item if provided | |
| mask_vals = None | |
| if protected_mask is not None: | |
| mask_vals = protected_mask[b, :Lb].tolist() # (Lb,) bool list | |
| if Lb == 0: | |
| out_tokens = [eos_id] | |
| out_mask = [False] if mask_vals is not None else None | |
| else: | |
| out_tokens = [] | |
| out_mask = [] if mask_vals is not None else None | |
| for i in range(Lb): | |
| t_i = tokens[i] | |
| m_i = mask_vals[i] if mask_vals is not None else None | |
| if i < Lmax and bool(del_mask[b, i].item()): | |
| # Delete: skip token and mask value | |
| continue | |
| if i < Lmax and bool(sub_mask[b, i].item()): | |
| out_tokens.append(int(sub_tok[b, i].item())) | |
| else: | |
| out_tokens.append(int(t_i)) | |
| # Keep mask value for this position (substitution doesn't change position) | |
| if out_mask is not None: | |
| out_mask.append(m_i) | |
| if i < Lmax and bool(ins_mask[b, i].item()): | |
| out_tokens.append(int(ins_tok[b, i].item())) | |
| # Insert: new position, not in PAM domain, so False | |
| if out_mask is not None: | |
| out_mask.append(False) | |
| if len(out_tokens) == 0 or out_tokens[-1] != eos_id: | |
| out_tokens.append(eos_id) | |
| if out_mask is not None: | |
| out_mask.append(False) # EOS can't be deleted anyway | |
| if max_len_cap is not None and len(out_tokens) > max_len_cap: | |
| out_tokens = out_tokens[:max_len_cap] | |
| if out_mask is not None: | |
| out_mask = out_mask[:max_len_cap] | |
| if out_tokens[-1] != eos_id: | |
| out_tokens[-1] = eos_id | |
| new_seqs.append(torch.tensor(out_tokens, device=device, dtype=torch.long)) | |
| new_lens.append(len(out_tokens)) | |
| if out_mask is not None: | |
| new_masks.append(torch.tensor(out_mask, device=device, dtype=torch.bool)) | |
| Lout = max(1, max(new_lens) if new_lens else 1) | |
| x_out = torch.full((B, Lout), pad_id, device=device, dtype=x.dtype) | |
| for b, s in enumerate(new_seqs): | |
| x_out[b, : s.numel()] = s | |
| # Reconstruct mask tensor if masks were provided | |
| protected_mask_out = None | |
| if new_masks: | |
| protected_mask_out = torch.full((B, Lout), False, device=device, dtype=torch.bool) | |
| for b, m in enumerate(new_masks): | |
| protected_mask_out[b, : m.numel()] = m | |
| # Collect edit stats per batch item | |
| edit_stats = { | |
| 'num_ins': [ins_mask[b].sum().item() for b in range(B)], | |
| 'num_del': [del_mask[b].sum().item() for b in range(B)], | |
| 'num_sub': [sub_mask[b].sum().item() for b in range(B)], | |
| 'total_edits': [ins_mask[b].sum().item() + del_mask[b].sum().item() + sub_mask[b].sum().item() for b in range(B)], | |
| 'avg_rates': avg_rates_per_batch # Store average rates post-amplification | |
| } | |
| return x_out, base_rate, protected_mask_out, edit_stats | |
| # --------------------------------------------------------------------------- | |
| # ATC + G_T | |
| # --------------------------------------------------------------------------- | |
| def _augmented_tchebycheff( | |
| f_vals: torch.Tensor, | |
| w: torch.Tensor, | |
| rho: float, | |
| z: torch.Tensor, | |
| ) -> torch.Tensor: | |
| diff = f_vals - z | |
| term1 = torch.min(w * diff, dim=1).values | |
| term2 = rho * torch.sum(w * diff, dim=1) | |
| return term1 + term2 | |
| def _G_T( | |
| protein_tokens: torch.Tensor, | |
| objective_models, | |
| constraint_models, | |
| w: torch.Tensor, | |
| rho: float, | |
| z: torch.Tensor, | |
| beta: float, | |
| tokenizer, | |
| ws_for_invalid: bool = False, | |
| debug_context=None, | |
| count_terminal: bool = False, | |
| return_details: bool = False, | |
| ): | |
| """ | |
| Matches behavior of: | |
| - cope_batch_multi_edits_log_length (1).py | |
| - pcomol (1).py | |
| Key semantics: | |
| - Constraints are evaluated for all sequences. | |
| - If ws_for_invalid=True: | |
| * weighted_sum_full is computed for ALL sequences (valid or invalid) | |
| * G_full is ONLY assigned for constraint-valid sequences (invalid remain -inf) | |
| - If ws_for_invalid=False: | |
| * both weighted_sum_full and G_full are ONLY assigned for constraint-valid sequences | |
| (invalid remain -inf) | |
| """ | |
| device = protein_tokens.device | |
| # Decode sequences | |
| protein_seqs = [ | |
| seq.replace(" ", "") | |
| for seq in tokenizer.batch_decode(protein_tokens, skip_special_tokens=True) | |
| ] | |
| # ------------------------- | |
| # constraints (evaluate on ALL) | |
| # ------------------------- | |
| constraint_results = [] | |
| for constraint in constraint_models: | |
| if hasattr(constraint, "predict_batch"): | |
| res = constraint.predict_batch(protein_seqs) | |
| res = [int(r) for r in res] | |
| else: | |
| res = [ | |
| constraint(protein_tokens[i] if protein_tokens is not None else None, seq) | |
| for i, seq in enumerate(protein_seqs) | |
| ] | |
| constraint_results.append(res) | |
| constraint_results = torch.tensor(constraint_results, device=device) | |
| survived_seq_indices = (constraint_results == 1).all(dim=0).nonzero(as_tuple=True)[0] | |
| survived_seqs = [protein_seqs[idx] for idx in survived_seq_indices.tolist()] | |
| # outputs | |
| B = len(protein_seqs) | |
| weighted_sum_full = torch.full((B,), float("-inf"), device=device) | |
| G_full = torch.full((B,), float("-inf"), device=device) | |
| # ---- oracle-call accounting (inert unless --instrument) -------------------- | |
| # One "oracle call" == one sequence submitted to _G_T (constraints + objectives). | |
| # count_terminal=True marks calls that score genuine t=1 terminals, which is the | |
| # only class of call the unguided+rejection baseline ever makes, so | |
| # first_feasible_terminal_at is comparable between the two arms. | |
| if INSTR is not None: | |
| n_feas = int(survived_seq_indices.numel()) | |
| _count("oracle_evals", B) | |
| _count("oracle_evals_feasible", n_feas) | |
| if count_terminal: | |
| _count("oracle_terminal_evals", B) | |
| _count("oracle_terminal_feasible", n_feas) | |
| if n_feas > 0 and INSTR.get("first_feasible_terminal_at") is None: | |
| # Charged at batch granularity: the whole batch counts, even if the | |
| # feasible terminal was not the last element. Conservative for pCoMole. | |
| INSTR["first_feasible_terminal_at"] = INSTR["counts"]["oracle_evals"] | |
| # ------------------------- | |
| # objectives | |
| # ------------------------- | |
| f_vals = None | |
| if ws_for_invalid: | |
| # Compute objective vector for ALL sequences | |
| f_vals = extract_objective_vector(protein_seqs, objective_models, device) # (B, m) | |
| # Compute weighted sum for ALL sequences (matching pcomol.py) | |
| weighted_sum_full = torch.sum(w * f_vals, dim=1) # (B,) | |
| # Compute G for ALL sequences (matching pcomol.py behavior) | |
| u_atc = _augmented_tchebycheff(f_vals, w, rho, z) # (B,) | |
| G = beta * u_atc # (B,) | |
| # Only assign valid sequences to G_full (invalid remain -inf) | |
| G_full[survived_seq_indices] = G[survived_seq_indices] | |
| else: | |
| # Terminal scoring mode: both ws and G only for constraint-valid sequences | |
| if survived_seq_indices.numel() > 0: | |
| f_vals = extract_objective_vector(survived_seqs, objective_models, device) # (B', m) | |
| u_atc = _augmented_tchebycheff(f_vals, w, rho, z) # (B',) | |
| G = beta * u_atc # (B',) | |
| weighted_sum = torch.sum(w * f_vals, dim=1) # (B',) | |
| G_full[survived_seq_indices] = G | |
| weighted_sum_full[survived_seq_indices] = weighted_sum | |
| if return_details: | |
| details = { | |
| "protein_seqs": protein_seqs, | |
| "constraint_results": constraint_results, # (n_constraints, B) int | |
| "survived_indices": survived_seq_indices, | |
| # f_vals spans all B rows only when ws_for_invalid=True; otherwise it is | |
| # restricted to the survivors (and is None when there are none). | |
| "f_vals": f_vals, | |
| } | |
| return G_full, weighted_sum_full, details | |
| return G_full, weighted_sum_full | |
| # def _G_T( | |
| # protein_tokens: torch.Tensor, | |
| # objective_models: List[Callable[[torch.Tensor], Tuple[str, Any]]], | |
| # constraint_models: List[Callable[[torch.Tensor], torch.Tensor]], | |
| # w: torch.Tensor, | |
| # rho: float, | |
| # z: torch.Tensor, | |
| # beta: float, | |
| # tokenizer, | |
| # ws_for_invalid=False, | |
| # debug_context: Optional[str] = None | |
| # ): | |
| # device = protein_tokens.device | |
| # # Decode protein sequences from tokens | |
| # protein_seqs = [seq.replace(' ', '') for seq in tokenizer.batch_decode(protein_tokens, skip_special_tokens=True)] | |
| # constraint_results = [] | |
| # for constraint in constraint_models: | |
| # # Handle constraints that take single sequences vs batches | |
| # if hasattr(constraint, 'predict_batch'): | |
| # # Use batch prediction if available (more efficient) | |
| # res = constraint.predict_batch(protein_seqs) | |
| # res = [int(r) for r in res] # Convert bool to int | |
| # else: | |
| # # Call for each sequence individually | |
| # res = [constraint(protein_tokens[i] if protein_tokens is not None else None, seq) | |
| # for i, seq in enumerate(protein_seqs)] | |
| # constraint_results.append(res) | |
| # constraint_results = torch.tensor(constraint_results, device=device) | |
| # survived_seq_indices = (constraint_results == 1).all(dim=0).nonzero(as_tuple=True)[0] | |
| # survived_seqs = [protein_seqs[idx] for idx in survived_seq_indices.tolist()] # (B') | |
| # weighted_sum_full = torch.full((len(protein_seqs),), float("-inf"), device=device) | |
| # G_full = torch.full((len(protein_seqs),), float("-inf"), device=device) | |
| # # Get objective names and find DeletionCount objective for absolute count | |
| # obj_names = [] | |
| # deletion_obj = None | |
| # deletion_obj_idx = None | |
| # if protein_seqs: | |
| # # Get names by calling with first sequence | |
| # for obj_idx, obj in enumerate(objective_models): | |
| # name, _ = obj(protein_tokens=None, protein_seqs=[protein_seqs[0]]) | |
| # obj_names.append(name) | |
| # # Check if this is DeletionCount objective | |
| # if hasattr(obj, 'original_length') and hasattr(obj, 'max_deletion'): | |
| # deletion_obj = obj | |
| # deletion_obj_idx = obj_idx | |
| # else: | |
| # # Fallback: use generic names | |
| # obj_names = [f"obj_{i}" for i in range(len(objective_models))] | |
| # # Helper function to format objective score with absolute deletion count if applicable | |
| # def format_obj_score(obj_name, raw_score, obj_idx, seq): | |
| # if obj_name == 'deletion_count' and deletion_obj is not None: | |
| # current_length = len(seq.replace(' ', '')) | |
| # abs_deletion = deletion_obj.original_length - current_length | |
| # return f"{obj_name}: {raw_score:.4f} (abs: {abs_deletion})" | |
| # return f"{obj_name}: {raw_score:.4f}" | |
| # # objectives | |
| # if ws_for_invalid: | |
| # # Compute objective scores for all sequences (including invalid ones) | |
| # f_vals = extract_objective_vector(protein_seqs, objective_models, device) | |
| # # Compute weighted sum for all sequences | |
| # weighted_scores = w.unsqueeze(0) * f_vals # (B, m) - element-wise multiplication | |
| # weighted_sum_all = torch.sum(weighted_scores, dim=1) # (B,) | |
| # # Only set weighted_sum_full for valid sequences (invalid ones stay -inf) | |
| # weighted_sum_full[survived_seq_indices] = weighted_sum_all[survived_seq_indices] | |
| # # Compute G only for valid sequences | |
| # if survived_seq_indices.numel() > 0: | |
| # f_vals_valid = f_vals[survived_seq_indices] | |
| # u_atc = _augmented_tchebycheff(f_vals_valid, w, rho, z) | |
| # G = beta * u_atc | |
| # G_full[survived_seq_indices] = G | |
| # # DEBUG: Print logG calculation (ATC-based) | |
| # if debug_context is not None: | |
| # seq_idx = 0 # Show first sequence only for brevity | |
| # if seq_idx < len(f_vals_valid): | |
| # f_seq = f_vals_valid[seq_idx] # (m,) | |
| # diff = f_seq - z # (m,) - distance from reference point | |
| # w_diff = w * diff # (m,) - weighted differences | |
| # term1 = torch.min(w_diff).item() # min(w * diff) | |
| # term2 = (rho * torch.sum(w_diff)).item() # rho * sum(w * diff) | |
| # u_atc_val = u_atc[seq_idx].item() | |
| # logG_val = G[seq_idx].item() | |
| # # Format output | |
| # diff_parts = [] | |
| # w_diff_parts = [] | |
| # for obj_idx, obj_name in enumerate(obj_names): | |
| # raw_score = f_seq[obj_idx].item() | |
| # diff_val = diff[obj_idx].item() | |
| # w_diff_val = w_diff[obj_idx].item() | |
| # weight = w[obj_idx].item() | |
| # ref_val = z[obj_idx].item() | |
| # # Add absolute deletion count if applicable | |
| # if obj_name == 'deletion_count' and deletion_obj is not None: | |
| # # Get the actual sequence index in the original protein_seqs | |
| # actual_seq_idx = survived_seq_indices[seq_idx].item() | |
| # seq_str = protein_seqs[actual_seq_idx] | |
| # abs_deletion = deletion_obj.original_length - len(seq_str.replace(' ', '')) | |
| # diff_parts.append(f"{obj_name}: {raw_score:.4f} (abs: {abs_deletion}) - {ref_val:.4f} = {diff_val:.4f}") | |
| # else: | |
| # diff_parts.append(f"{obj_name}: {raw_score:.4f} - {ref_val:.4f} = {diff_val:.4f}") | |
| # w_diff_parts.append(f"{obj_name}: {diff_val:.4f} × {weight:.3f} = {w_diff_val:.4f}") | |
| # print(f"[{debug_context}] logG calc: diffs=[{', '.join(diff_parts)}] | " | |
| # f"w×diffs=[{', '.join(w_diff_parts)}] | " | |
| # f"min={term1:.4f}, rho×sum={term2:.4f} (rho={rho:.3f}) | " | |
| # f"u_atc={u_atc_val:.4f}, logG={logG_val:.4f} (β={beta:.3f})") | |
| # else: | |
| # if survived_seq_indices.numel() > 0: | |
| # f_vals = extract_objective_vector(survived_seqs, objective_models, device) # (B', m) | |
| # # Compute weighted scores | |
| # weighted_scores = w.unsqueeze(0) * f_vals # (B', m) | |
| # weighted_sum = torch.sum(weighted_scores, dim=1) # (B',) | |
| # u_atc = _augmented_tchebycheff(f_vals, w, rho, z) # (B',) | |
| # G = beta * u_atc # (B',) | |
| # G_full[survived_seq_indices] = G | |
| # weighted_sum_full[survived_seq_indices] = weighted_sum | |
| # # DEBUG: Print logG calculation (ATC-based) | |
| # if debug_context is not None: | |
| # seq_idx = 0 # Show first sequence only for brevity | |
| # if seq_idx < len(f_vals): | |
| # f_seq = f_vals[seq_idx] # (m,) | |
| # diff = f_seq - z # (m,) - distance from reference point | |
| # w_diff = w * diff # (m,) - weighted differences | |
| # term1 = torch.min(w_diff).item() # min(w * diff) | |
| # term2 = (rho * torch.sum(w_diff)).item() # rho * sum(w * diff) | |
| # u_atc_val = u_atc[seq_idx].item() | |
| # logG_val = G[seq_idx].item() | |
| # # Format output | |
| # diff_parts = [] | |
| # w_diff_parts = [] | |
| # for obj_idx, obj_name in enumerate(obj_names): | |
| # raw_score = f_seq[obj_idx].item() | |
| # diff_val = diff[obj_idx].item() | |
| # w_diff_val = w_diff[obj_idx].item() | |
| # weight = w[obj_idx].item() | |
| # ref_val = z[obj_idx].item() | |
| # # Add absolute deletion count if applicable | |
| # if obj_name == 'deletion_count' and deletion_obj is not None: | |
| # seq_str = survived_seqs[seq_idx] | |
| # abs_deletion = deletion_obj.original_length - len(seq_str.replace(' ', '')) | |
| # diff_parts.append(f"{obj_name}: {raw_score:.4f} (abs: {abs_deletion}) - {ref_val:.4f} = {diff_val:.4f}") | |
| # else: | |
| # diff_parts.append(f"{obj_name}: {raw_score:.4f} - {ref_val:.4f} = {diff_val:.4f}") | |
| # w_diff_parts.append(f"{obj_name}: {diff_val:.4f} × {weight:.3f} = {w_diff_val:.4f}") | |
| # print(f"[{debug_context}] logG calc: diffs=[{', '.join(diff_parts)}] | " | |
| # f"w×diffs=[{', '.join(w_diff_parts)}] | " | |
| # f"min={term1:.4f}, rho×sum={term2:.4f} (rho={rho:.3f}) | " | |
| # f"u_atc={u_atc_val:.4f}, logG={logG_val:.4f} (β={beta:.3f})") | |
| # # return full-size tensors (B,) | |
| # return G_full, weighted_sum_full | |
| # --------------------------------------------------------------------------- | |
| # rollout | |
| # --------------------------------------------------------------------------- | |
| def short_rollout_batch( | |
| model, | |
| x0: torch.Tensor, # (B, Lmax) padded | |
| time_grid: torch.Tensor, | |
| start_idx: int, | |
| pad_id: int, | |
| bos_id: int, | |
| eos_id: int, | |
| allowed_tokens: Optional[torch.Tensor], | |
| max_len_cap: Optional[int], | |
| num_rollouts: int = 1, | |
| num_steps: int =32, | |
| protected_mask: Optional[torch.Tensor] = None, # (B, Lmax) bool, True=no edits allowed | |
| deletion_rate_scale: float = 1500.0, # Scaling factor for deletion rate | |
| zero_lam_ins: bool = False, # If True, set lam_ins to 0.0 before processing | |
| ) -> torch.Tensor: | |
| """ | |
| Returns: | |
| xT: (B*num_rollouts, Lmax) | |
| Grouping: | |
| xT[i*num_rollouts:(i+1)*num_rollouts] corresponds to candidate i. | |
| """ | |
| device = x0.device | |
| B, Lmax = x0.shape | |
| # repeat each candidate num_rollouts times (grouped) | |
| x = x0.repeat_interleave(num_rollouts, dim=0) # (B*num_rollouts, Lmax) | |
| # repeat protected_mask if provided | |
| if protected_mask is not None: | |
| protected_mask_repeated = protected_mask.repeat_interleave(num_rollouts, dim=0) # (B*num_rollouts, Lmax) | |
| else: | |
| protected_mask_repeated = None | |
| # rollout in batch | |
| for j in range(start_idx + 1, time_grid.numel()): | |
| t_j = time_grid[j].view(1).to(device) | |
| mask = (x != pad_id) | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub, *_ = model(x_t=x, mask=mask, t=t_j) | |
| x, _, protected_mask_repeated, _ = _sample_multiple_edits_batch( | |
| x, | |
| lam_ins, logits_ins, | |
| lam_del, lam_sub, logits_sub, | |
| pad_id, bos_id, eos_id, | |
| allowed_tokens, | |
| delta=float(1/(num_steps-1)), | |
| max_len_cap=max_len_cap, | |
| protected_mask=protected_mask_repeated, | |
| pam_scale_edits=False, # Rollouts use masking mode (scale_edits only applies to candidate generation) | |
| pam_edit_scale_factor=10.0, # Not used in rollouts | |
| deletion_rate_scale=deletion_rate_scale, | |
| debug_context=f"ROLLOUT t={j}/{time_grid.numel()-1}", | |
| zero_lam_ins=zero_lam_ins, | |
| ) | |
| return x | |
| # --------------------------------------------------------------------------- | |
| # finalizer | |
| # --------------------------------------------------------------------------- | |
| def _finalize_from_last( | |
| model, | |
| x_last: torch.Tensor, | |
| time_grid: torch.Tensor, | |
| last_step: int, | |
| pad_id: int, | |
| bos_id: int, | |
| eos_id: int, | |
| allowed_tokens: Optional[torch.Tensor], | |
| objective_models: List[Callable[[torch.Tensor], Tuple[str, Any]]], | |
| constraint_models: List[Callable[[torch.Tensor], torch.Tensor]], | |
| w: torch.Tensor, | |
| rho: float, | |
| ref_z: torch.Tensor, | |
| beta_final: float, | |
| max_len_cap: Optional[int] = None, | |
| num_final_rollouts: int = 50, | |
| num_steps: int = 32, | |
| tokenizer=None, | |
| # NEW: recompute PI mask for finalization step | |
| pam_masker: Optional[Cas9PIMasker] = None, | |
| deletion_rate_scale: float = 1500.0, # Scaling factor for deletion rate | |
| zero_lam_ins: bool = False, # If True, set lam_ins to 0.0 before processing | |
| ) -> torch.Tensor: | |
| logG_last, _ = _G_T(x_last, objective_models, constraint_models, w, rho, ref_z, beta_final, tokenizer, ws_for_invalid=False, debug_context="FINALIZATION_INITIAL") | |
| # start_idx = min(last_step, time_grid.numel() - 2) if time_grid.numel() >= 2 else 0 | |
| # Recompute PI mask for finalization step (protects against all edits) | |
| protected_mask = None | |
| if pam_masker is not None: | |
| seq_last = tokenizer.batch_decode(x_last, skip_special_tokens=True)[0].replace(" ", "") | |
| protected_mask = pam_masker.build_no_del_mask(x_last, seq_last, pad_id=pad_id, bos_at_index0=True) | |
| x_Ts = short_rollout_batch(model, x_last, time_grid, last_step, pad_id, bos_id, eos_id, allowed_tokens, max_len_cap, num_final_rollouts, num_steps, protected_mask=protected_mask, deletion_rate_scale=deletion_rate_scale, zero_lam_ins=zero_lam_ins) | |
| logG, _ = _G_T(x_Ts, objective_models, constraint_models, w, rho, ref_z, beta_final, tokenizer, ws_for_invalid=False, debug_context="FINALIZATION_ROLLOUTS", count_terminal=True) | |
| idx = torch.isfinite(logG).nonzero(as_tuple=True)[0].tolist() | |
| if len(idx) == 0 or torch.max(logG) < logG_last: | |
| return x_last, logG_last | |
| else: | |
| best_idx = torch.argmax(logG).item() | |
| best_seq = x_Ts[best_idx].unsqueeze(0) | |
| return best_seq, logG[best_idx] | |
| def pCoMol( | |
| model, | |
| x0: torch.Tensor, | |
| *, | |
| pad_id: int, | |
| bos_id: int, | |
| eos_id: int, | |
| allowed_tokens: Optional[torch.Tensor], | |
| objective_models: List[Callable[[torch.Tensor], Tuple[str, Any]]], | |
| constraint_models: List[Callable[[torch.Tensor], torch.Tensor]], | |
| w: torch.Tensor, | |
| rho: float, | |
| ref_z: torch.Tensor, | |
| beta_start: float = 1.0, | |
| beta_end: float = 3.0, | |
| num_steps: int = 32, | |
| num_candidates: int = 8, | |
| num_rollouts: int = 4, | |
| max_len_cap: Optional[int] = None, | |
| device: Optional[torch.device] = None, | |
| num_final_rollouts: int = 16, | |
| cfg, | |
| tokenizer, | |
| # PAM masking parameters | |
| pam_masker: Optional[Cas9PIMasker] = None, | |
| pam_mask_refresh_every: int = 1, | |
| pam_debug: bool = False, | |
| pam_scale_edits: bool = False, # If True, scale up edit rates in PAM region instead of masking | |
| pam_edit_scale_factor: float = 10.0, # Scaling factor for ins/sub rates in PAM region | |
| deletion_rate_scale: float = 1500.0, # Scaling factor for deletion rate | |
| zero_lam_ins: bool = False, # If True, set lam_ins to 0.0 before processing | |
| legacy_beta_incumbent: bool = False, # True restores the pre-fix incumbent rule | |
| ) -> torch.Tensor: | |
| if device is None: | |
| device = x0.device | |
| x = x0.clone().to(device) | |
| time_grid = torch.linspace(0.0, 1.0, steps=num_steps, device=device) | |
| last_timestep = 0 | |
| best_terminal = None | |
| best_terminal_logG = float("-inf") | |
| # --- incumbent scoring ------------------------------------------------- | |
| # _G_T returns logG = beta * U, and beta_t is annealed beta_start -> beta_end | |
| # across the run. Comparing raw logG across steps therefore ranks terminals by | |
| # WHEN they were found rather than by utility (beta spans 3x; feasible U spans | |
| # ~1.4x), so the incumbent drifts toward late-trajectory terminals. Eq. (9) | |
| # defines G with a single fixed beta and Prop. C.6 requires the returned design | |
| # to maximise it over all evaluated feasible terminals, so we normalise logG by | |
| # the beta it was computed with before comparing. Set legacy_beta_incumbent=True | |
| # to restore the previous (beta-weighted) behaviour. | |
| def _incumbent_score(logG_val, beta): | |
| if legacy_beta_incumbent: | |
| return logG_val | |
| return logG_val / beta | |
| best_terminal_score = float("-inf") | |
| if legacy_beta_incumbent: | |
| print("[pCoMol] legacy_beta_incumbent=True: incumbent ranked by beta_t*U (pre-fix behaviour).") | |
| # Track cumulative edit statistics for selected steps only | |
| total_ins = 0 | |
| total_del = 0 | |
| total_sub = 0 | |
| protected_mask = None # (1, Lmax) bool, True=no edits allowed in PAM-protected regions | |
| def _refresh_protected_mask(curr_x: torch.Tensor, step_num: int) -> Optional[torch.Tensor]: | |
| if pam_masker is None: | |
| return None | |
| seq = tokenizer.batch_decode(curr_x, skip_special_tokens=True)[0].replace(" ", "") | |
| m = pam_masker.build_no_del_mask(curr_x, seq, pad_id=pad_id, bos_at_index0=True) | |
| if pam_debug: | |
| # Get masked interval for debug output | |
| interval = pam_masker.pi_core_interval(seq) | |
| masked_count = int(m.sum().item()) | |
| seq_len = len(seq) | |
| if interval is not None: | |
| s, t = interval # 1-based AA positions | |
| mode_str = f"edit rates scaled by {pam_edit_scale_factor}x" if pam_scale_edits else "all edits blocked" | |
| print(f"[PAM mask] Step {step_num}: protected_range=[{s}-{t}] (1-based AA), protected_positions={masked_count}, seq_len={seq_len} ({mode_str})") | |
| else: | |
| print(f"[PAM mask] Step {step_num}: no PI hit found, protected_positions={masked_count}, seq_len={seq_len}") | |
| return m | |
| with torch.no_grad(): | |
| for step in tqdm(range(num_steps - 1)): | |
| t = time_grid[step].view(1) | |
| frac = step / max(1, (num_steps - 1)) | |
| beta_t = beta_start + (beta_end - beta_start) * frac | |
| # DEBUG: Print current sequence logG calculation at start of step (compact) | |
| if step == 0 or step % max(1, num_steps // 5) == 0: # Print at start and every ~20% of steps | |
| curr_seq_str = tokenizer.batch_decode(x, skip_special_tokens=True)[0].replace(" ", "") | |
| curr_obj_scores, obj_names = extract_objective_vector([curr_seq_str], objective_models, device, return_names=True) | |
| curr_obj_scores = curr_obj_scores.squeeze(0) | |
| # Find DeletionCount objective for absolute count | |
| deletion_obj = None | |
| for obj in objective_models: | |
| if hasattr(obj, 'original_length') and hasattr(obj, 'max_deletion'): | |
| deletion_obj = obj | |
| break | |
| # Compute logG components | |
| diff = curr_obj_scores - ref_z | |
| w_diff = w * diff | |
| term1 = torch.min(w_diff).item() | |
| term2 = (rho * torch.sum(w_diff)).item() | |
| u_atc_val = term1 + term2 | |
| logG_val = beta_t * u_atc_val | |
| # Format output | |
| diff_parts = [] | |
| w_diff_parts = [] | |
| for obj_idx, obj_name in enumerate(obj_names): | |
| raw_score = curr_obj_scores[obj_idx].item() | |
| diff_val = diff[obj_idx].item() | |
| w_diff_val = w_diff[obj_idx].item() | |
| weight = w[obj_idx].item() | |
| ref_val = ref_z[obj_idx].item() | |
| if obj_name == 'deletion_count' and deletion_obj is not None: | |
| abs_deletion = deletion_obj.original_length - len(curr_seq_str) | |
| diff_parts.append(f"{obj_name}: {raw_score:.4f} (abs: {abs_deletion}) - {ref_val:.4f} = {diff_val:.4f}") | |
| else: | |
| diff_parts.append(f"{obj_name}: {raw_score:.4f} - {ref_val:.4f} = {diff_val:.4f}") | |
| w_diff_parts.append(f"{obj_name}: {diff_val:.4f} × {weight:.3f} = {w_diff_val:.4f}") | |
| print(f"[STEP {step} START] len={len(curr_seq_str)} | diffs=[{', '.join(diff_parts)}] | " | |
| f"w×diffs=[{', '.join(w_diff_parts)}] | min={term1:.4f}, rho×sum={term2:.4f} | " | |
| f"u_atc={u_atc_val:.4f}, logG={logG_val:.4f} (β={beta_t:.3f})") | |
| # Refresh PI protection mask for current accepted sequence (blocks all edits) | |
| if pam_masker is not None and (step % max(1, pam_mask_refresh_every) == 0): | |
| protected_mask = _refresh_protected_mask(x, step) | |
| _count("steps"); _t_cand = _tic() | |
| # model forward | |
| mask = (x != pad_id) | |
| # ReparameterizedProteinEditFlowModel returns 8 values, ProteinEditFlowModel returns 5 | |
| model_output = model(x_t=x, mask=mask, t=t) | |
| if len(model_output) == 8: | |
| # ReparameterizedProteinEditFlowModel: (lam_ins, logits_ins, lam_del, lam_sub, logits_sub, lam_total, logits_type, pi_type) | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub, lam_total, logits_type, pi_type = model_output | |
| elif len(model_output) == 5: | |
| # ProteinEditFlowModel: (lam_ins, logits_ins, lam_del, lam_sub, logits_sub) | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub = model_output | |
| lam_total = lam_ins + lam_del + lam_sub | |
| pi_type = torch.stack([lam_ins, lam_del, lam_sub], dim=-1) / lam_total.clamp_min(1e-12) | |
| else: | |
| raise ValueError(f"Unexpected model output length: {len(model_output)}") | |
| candidates = [x.squeeze(0)] # compute the scores of current sequence with the candidates | |
| base_rates = [] | |
| candidate_edit_stats = [] # Store edit stats for each candidate | |
| candidate_to_stats = {} # Map candidate tensor to edit stats (for deduplication) | |
| for cand_idx in range(num_candidates): | |
| cand_seq, base_rate, _, edit_stats = _sample_multiple_edits_batch( | |
| x, | |
| lam_ins, logits_ins, | |
| lam_del, lam_sub, logits_sub, | |
| pad_id, bos_id, eos_id, | |
| allowed_tokens, | |
| delta=float(1/(num_steps-1)), | |
| max_len_cap=max_len_cap, | |
| protected_mask=protected_mask, | |
| pam_scale_edits=pam_scale_edits, | |
| pam_edit_scale_factor=pam_edit_scale_factor, | |
| deletion_rate_scale=deletion_rate_scale, | |
| debug_context=f"CANDIDATE step={step} cand={cand_idx}", | |
| zero_lam_ins=zero_lam_ins, | |
| ) | |
| if not torch.equal(cand_seq, x): | |
| cand_seq_squeezed = cand_seq.squeeze(0) | |
| # Use a hash of the tensor as key (simple approach) | |
| cand_key = tuple(cand_seq_squeezed.cpu().tolist()) | |
| if cand_key not in candidate_to_stats: | |
| candidates.append(cand_seq_squeezed) | |
| base_rates.append(base_rate) | |
| candidate_edit_stats.append(edit_stats) | |
| candidate_to_stats[cand_key] = len(candidate_edit_stats) - 1 | |
| else: | |
| # Duplicate candidate, keep the stats from first occurrence | |
| pass | |
| batch_candidates = torch.nn.utils.rnn.pad_sequence(candidates, batch_first=True, padding_value=pad_id) | |
| num_generated_candidates = len(candidates) - 1 # Exclude the current sequence | |
| _toc("candidate_proposal", _t_cand); _count("candidates", num_generated_candidates) | |
| # print("Initial Candidates: ", len(candidates) - 1) | |
| # pdb.set_trace() | |
| # We only want the survived candidates to improve the objective weights | |
| start = time.time() | |
| _t_scr = _tic() | |
| cand_logG, cand_ws = _G_T(batch_candidates, objective_models, constraint_models, w, rho, ref_z, beta_t, tokenizer, ws_for_invalid=True, debug_context=f"CANDIDATE_EVAL step={step}") | |
| _toc("screening", _t_scr) | |
| # print("Candidate Time: ", time.time() - start) | |
| curr_logG = cand_logG[0] | |
| curr_ws = cand_ws[0] | |
| cand_logG = cand_logG[1:] | |
| cand_ws = cand_ws[1:] | |
| batch_candidates = batch_candidates[1:, :] | |
| # DEBUG: Print final scores used for candidate selection (compact) | |
| if len(cand_ws) > 0: | |
| # valid_mask = torch.isfinite(cand_ws) | |
| valid_mask = torch.isfinite(cand_logG) # valid = passed constraints | |
| if valid_mask.any(): | |
| valid_indices = valid_mask.nonzero(as_tuple=True)[0] | |
| print(f"[CANDIDATE_SELECTION step={step}] Current: logG={curr_logG.item():.4f}, WS={curr_ws.item():.4f} | " | |
| f"Valid: {valid_mask.sum().item()}/{len(cand_ws)} | " | |
| f"Top 3: {', '.join([f'logG={cand_logG[valid_indices[i]].item():.4f},WS={cand_ws[valid_indices[i]].item():.4f}' for i in range(min(3, len(valid_indices)))])}") | |
| else: | |
| print(f"[CANDIDATE_SELECTION step={step}] WARNING: No valid candidates (all failed constraints)") | |
| # Debug: Print which constraints each candidate failed | |
| if len(batch_candidates) > 0: | |
| candidate_seqs = tokenizer.batch_decode(batch_candidates, skip_special_tokens=True) | |
| candidate_seqs_clean = [seq.replace(" ", "").replace("\n", "") for seq in candidate_seqs] | |
| print(f"[CONSTRAINT_FAILURE_DEBUG step={step}] Analyzing constraint failures for {len(batch_candidates)} candidates:") | |
| for cand_idx in range(len(batch_candidates)): | |
| failed_constraints = [] | |
| seq_len = len(candidate_seqs_clean[cand_idx]) | |
| for constraint in constraint_models: | |
| constraint_name = constraint.__class__.__name__ | |
| if hasattr(constraint, "predict_batch"): | |
| result = constraint.predict_batch([candidate_seqs[cand_idx]])[0] | |
| else: | |
| result = constraint(None, candidate_seqs[cand_idx]) | |
| if not result: | |
| # Get failure reason | |
| if constraint_name == "MinTargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} < {constraint.min_target_length})") | |
| elif constraint_name == "MaxTargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} > {constraint.max_target_length})") | |
| elif constraint_name == "TargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} != {constraint.target_length})") | |
| elif constraint_name == "ProteinLength": | |
| upper = f"{constraint.L0}]" if constraint.allow_equal_length else f"{constraint.L0})" | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} not in [{constraint.L0 // 2}, {upper}") | |
| elif constraint_name == "Cas9DomainCompleteness": | |
| failed_constraints.append(f"{constraint_name}(domain incomplete)") | |
| elif constraint_name == "PAMMatchingConstraint": | |
| failed_constraints.append(f"{constraint_name}(predicted PAM does not match target)") | |
| elif constraint_name == "Cas9ScoreThreshold": | |
| failed_constraints.append(f"{constraint_name}(Cas9 score <= {constraint.threshold})") | |
| elif constraint_name == "PAMMatchingProbabilityThreshold": | |
| failed_constraints.append(f"{constraint_name}(PAM matching probability <= {constraint.threshold})") | |
| else: | |
| failed_constraints.append(f"{constraint_name}") | |
| if failed_constraints: | |
| print(f" Candidate {cand_idx}: length={seq_len}, failed: {', '.join(failed_constraints)}") | |
| else: | |
| print(f" Candidate {cand_idx}: length={seq_len}, passed all constraints (unexpected!)") | |
| if len(batch_candidates) == 0: | |
| _count("step_no_candidates") | |
| if pam_masker is not None: | |
| print(f"[PAM DEBUG] Step {step}: No candidates generated (all identical to current sequence). " | |
| f"This may indicate PAM mask is blocking all edits.") | |
| continue | |
| improve_idx = (cand_ws > curr_ws).nonzero(as_tuple=True)[0] | |
| survived_candidates = batch_candidates[improve_idx, :] | |
| base_rates = [base_rates[i] for i in improve_idx] # (num_survived_candidates,) | |
| survived_edit_stats = [candidate_edit_stats[i] for i in improve_idx] # Store edit stats for survived candidates | |
| # print([len(seq.replace(' ' ,'')) for seq in tokenizer.batch_decode(survived_candidates, skip_special_tokens=True)]) | |
| # print("Num Candidates Survived: ", len(improve_idx)) | |
| if len(improve_idx) == 0: | |
| _count("step_no_improve") | |
| # Debug: Print why no candidates improved | |
| if num_generated_candidates > 0: | |
| print(f"[PAM DEBUG] Step {step}: {num_generated_candidates} candidates generated, but none improved weighted sum.") | |
| print(f" Current weighted sum: {curr_ws.item():.6f}") | |
| if len(cand_ws) > 0: | |
| valid_cand_mask = torch.isfinite(cand_ws) | |
| num_valid = valid_cand_mask.sum().item() | |
| print(f" Valid candidates: {num_valid}/{len(cand_ws)}") | |
| if num_valid > 0: | |
| print(f" Valid candidate weighted sums: min={cand_ws[valid_cand_mask].min().item():.6f}, max={cand_ws[valid_cand_mask].max().item():.6f}, mean={cand_ws[valid_cand_mask].mean().item():.6f}") | |
| else: | |
| print(f" WARNING: All candidates are invalid (don't pass constraints)!") | |
| # Decode and show objective scores for current and a few candidates | |
| curr_seq_str = tokenizer.batch_decode(x, skip_special_tokens=True)[0].replace(" ", "") | |
| curr_obj_scores = extract_objective_vector([curr_seq_str], objective_models, device) | |
| print(f" Current sequence: ws={curr_ws.item():.6f}, obj_scores={curr_obj_scores.squeeze().tolist()}, weights={w.tolist()}") | |
| # Show a few valid candidates if any | |
| valid_indices = valid_cand_mask.nonzero(as_tuple=True)[0][:3] | |
| for idx in valid_indices: | |
| cand_seq_str = tokenizer.batch_decode(batch_candidates[idx:idx+1], skip_special_tokens=True)[0].replace(" ", "") | |
| cand_obj_scores = extract_objective_vector([cand_seq_str], objective_models, device) | |
| print(f" Valid candidate {idx.item()}: ws={cand_ws[idx].item():.6f}, obj_scores={cand_obj_scores.squeeze().tolist()}") | |
| else: | |
| print(f"[PAM DEBUG] Step {step}: No candidates generated (all identical to current sequence).") | |
| if pam_masker is not None: | |
| curr_seq_str = tokenizer.batch_decode(x, skip_special_tokens=True)[0].replace(" ", "") | |
| if protected_mask is not None: | |
| protected_count = protected_mask.sum().item() | |
| print(f" Protected positions: {protected_count}/{len(curr_seq_str)} (all edits blocked in these positions)") | |
| continue | |
| # Expand mask to survived batch if provided | |
| if protected_mask is not None: | |
| B_mask, L_mask = protected_mask.shape | |
| B_survived, L_survived = survived_candidates.shape | |
| if L_survived != L_mask: | |
| # Length mismatch: mask was computed for a different sequence length | |
| # Recompute mask for each candidate sequence to ensure correct positions | |
| if pam_masker is not None: | |
| print(f"[PAM mask] Warning: Length mismatch detected (mask_len={L_mask}, candidate_len={L_survived}). " | |
| f"Recomputing mask for each candidate sequence.") | |
| protected_masks = [] | |
| for i in range(B_survived): | |
| cand_seq = survived_candidates[i:i+1] # Keep batch dim | |
| seq_str = tokenizer.batch_decode(cand_seq, skip_special_tokens=True)[0].replace(" ", "") | |
| cand_mask = pam_masker.build_no_del_mask(cand_seq, seq_str, pad_id=pad_id, bos_at_index0=True) | |
| protected_masks.append(cand_mask) | |
| protected_mask_expanded = torch.cat(protected_masks, dim=0) # (B_survived, L_survived) | |
| else: | |
| # No masker available, can't recompute - skip mask | |
| print(f"[PAM mask] Warning: Length mismatch (mask_len={L_mask}, candidate_len={L_survived}) " | |
| f"but no pam_masker available. Skipping mask for rollouts.") | |
| protected_mask_expanded = None | |
| else: | |
| # Lengths match, safe to expand | |
| protected_mask_expanded = protected_mask.expand(B_survived, -1) | |
| else: | |
| protected_mask_expanded = None | |
| # Keep all the rollout terminal sequences in one batch | |
| # Note: Rollouts use masking mode (not scaling) to preserve stability | |
| start = time.time() | |
| _t_roll = _tic() | |
| x_Ts = short_rollout_batch(model, survived_candidates, time_grid, step, pad_id, bos_id, eos_id, allowed_tokens, max_len_cap, num_rollouts, num_steps, protected_mask=protected_mask_expanded, deletion_rate_scale=deletion_rate_scale, zero_lam_ins=zero_lam_ins) | |
| _toc("rollout", _t_roll); _count("rollouts", int(survived_candidates.shape[0]) * num_rollouts) | |
| # print("Rollout Time: ", time.time() - start) | |
| # Debug: Print minimum terminal sequence length per candidate | |
| if len(survived_candidates) > 0: | |
| terminal_seqs = tokenizer.batch_decode(x_Ts, skip_special_tokens=True) | |
| terminal_lengths = [len(seq.replace(" ", "").replace("\n", "")) for seq in terminal_seqs] | |
| # Reshape: (num_candidates, num_rollouts) - each candidate has num_rollouts terminal sequences | |
| num_survived = len(survived_candidates) | |
| terminal_lengths_reshaped = [terminal_lengths[i*num_rollouts:(i+1)*num_rollouts] for i in range(num_survived)] | |
| min_lengths_per_candidate = [min(lengths) for lengths in terminal_lengths_reshaped] | |
| print(f"[TERMINAL_LENGTHS step={step}] Min length per candidate (across {num_rollouts} rollouts): {min_lengths_per_candidate}") | |
| # Debug: Print objective values, Cas9 scores, lengths, and predicted PAMs for all terminal sequences per candidate | |
| terminal_seqs_clean = [seq.replace(" ", "").replace("\n", "") for seq in terminal_seqs] | |
| # Get objective values for all terminal sequences | |
| obj_vals, obj_names = extract_objective_vector(terminal_seqs_clean, objective_models, device, return_names=True) | |
| obj_vals = obj_vals.cpu().tolist() # list of lists: (num_total_terminals, num_objectives) | |
| # Find Cas9 classifier and PAM matching objects | |
| cas9_classifier_obj = None | |
| pam_matching_obj = None | |
| for obj in objective_models: | |
| if isinstance(obj, Cas9Classification): | |
| cas9_classifier_obj = obj | |
| # Check for PAMMatching directly or wrapped in PAMDomainWrapper | |
| if isinstance(obj, PAMMatching) or (hasattr(obj, 'predict_pam') and hasattr(obj, 'target_pam')): | |
| pam_matching_obj = obj | |
| # Get Cas9 scores for all terminal sequences | |
| cas9_scores = None | |
| if cas9_classifier_obj is not None: | |
| cas9_scores = cas9_classifier_obj.get_scores(terminal_seqs_clean) # list of length num_total_terminals | |
| # Get predicted PAMs for all terminal sequences | |
| predicted_pams_all = None | |
| if pam_matching_obj is not None: | |
| predicted_pams_all = pam_matching_obj.predict_pam(terminal_seqs_clean) # list of length num_total_terminals | |
| # Print per candidate | |
| print(f"[TERMINAL_DEBUG step={step}] Per-candidate terminal sequence details:") | |
| for cand_idx in range(num_survived): | |
| print(f" Candidate {cand_idx}:") | |
| start_idx = cand_idx * num_rollouts | |
| end_idx = start_idx + num_rollouts | |
| for rollout_idx in range(num_rollouts): | |
| term_idx = start_idx + rollout_idx | |
| seq_clean = terminal_seqs_clean[term_idx] | |
| seq_len = terminal_lengths[term_idx] | |
| # Objective values | |
| obj_str = ", ".join([f"{obj_names[i]}={obj_vals[term_idx][i]:.4f}" for i in range(len(obj_names))]) | |
| # Cas9 score | |
| cas9_str = f"cas9_score={cas9_scores[term_idx]:.4f}" if cas9_scores is not None else "cas9_score=N/A" | |
| # Predicted PAM | |
| pam_str = f"predicted_pam={predicted_pams_all[term_idx]}" if predicted_pams_all is not None else "predicted_pam=N/A" | |
| print(f" Rollout {rollout_idx}: len={seq_len}, {obj_str}, {cas9_str}, {pam_str}") | |
| # pdb.set_trace() | |
| # Constraints are taken into account for the terminal sequences | |
| start = time.time() | |
| _t_term = _tic() | |
| logG, _, = _G_T(x_Ts, objective_models, constraint_models, w, rho, ref_z, beta_t, tokenizer, ws_for_invalid=False, debug_context=f"ROLLOUT_TERMINAL step={step}", count_terminal=True) | |
| _toc("terminal_oracle", _t_term) | |
| # pdb.set_trace() | |
| # Save the best teminal sequence | |
| curr_best_terminal_logG = torch.max(logG) | |
| curr_best_terminal_score = _incumbent_score(curr_best_terminal_logG, beta_t) | |
| if curr_best_terminal_logG != float('-inf') and best_terminal_score <= curr_best_terminal_score: | |
| best_terminal_idx = torch.argmax(logG) | |
| best_terminal = x_Ts[best_terminal_idx].unsqueeze(0) | |
| best_terminal_logG = curr_best_terminal_logG | |
| best_terminal_score = curr_best_terminal_score | |
| best_terminal_seq = tokenizer.batch_decode(best_terminal.tolist(), skip_special_tokens=True)[0] | |
| print("\nSaved Best Terminal: ", best_terminal_seq) | |
| print("Saved Best Terminal Length: ", len(best_terminal_seq.replace(' ', ''))) | |
| print("Saved Best Terminal logG: ", best_terminal_logG) | |
| # If logG is -inf, check which constraints failed | |
| # Handle both tensor and float values | |
| logG_value = best_terminal_logG.item() if isinstance(best_terminal_logG, torch.Tensor) else best_terminal_logG | |
| if math.isinf(logG_value) and logG_value < 0: | |
| print("Saved Best Terminal FAILED constraints. Checking which constraints failed:") | |
| failed_constraints = [] | |
| seq_clean = best_terminal_seq.replace(' ', '').replace('\n', '') | |
| seq_len = len(seq_clean) | |
| # Debug: Print sequence being evaluated | |
| print(f" [DEBUG] Sequence being evaluated (length={seq_len}):") | |
| print(f" First 100 chars: {seq_clean[:100]}") | |
| print(f" Last 100 chars: {seq_clean[-100:]}") | |
| print(f" Sequence contains spaces: {' ' in seq_clean}") | |
| newline_char = '\n' | |
| print(f" Sequence contains newlines: {newline_char in seq_clean}") | |
| for constraint in constraint_models: | |
| constraint_name = constraint.__class__.__name__ | |
| if hasattr(constraint, "predict_batch"): | |
| result = constraint.predict_batch([seq_clean])[0] | |
| else: | |
| result = constraint(None, seq_clean) | |
| if not result: | |
| # Get failure reason | |
| if constraint_name == "MinTargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} < {constraint.min_target_length})") | |
| elif constraint_name == "MaxTargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} > {constraint.max_target_length})") | |
| elif constraint_name == "TargetLength": | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} != {constraint.target_length})") | |
| elif constraint_name == "ProteinLength": | |
| upper = f"{constraint.L0}]" if constraint.allow_equal_length else f"{constraint.L0})" | |
| failed_constraints.append(f"{constraint_name}(length {seq_len} not in [{constraint.L0 // 2}, {upper}") | |
| elif constraint_name == "Cas9DomainCompleteness": | |
| failed_constraints.append(f"{constraint_name}(domain incomplete)") | |
| elif constraint_name == "PAMMatchingConstraint": | |
| # Get predicted PAM to show what it was vs target | |
| predicted_pams = constraint.pam_matching_obj.predict_pam([seq_clean]) | |
| predicted_pam = predicted_pams[0] if predicted_pams else "N/A" | |
| failed_constraints.append(f"{constraint_name}(predicted PAM '{predicted_pam}' does not match target '{constraint.target_pam}')") | |
| elif constraint_name == "Cas9ScoreThreshold": | |
| # Get actual Cas9 score to show what it was vs threshold | |
| scores = constraint.cas9_classifier_obj.get_scores([seq_clean]) | |
| cas9_score = scores[0] if scores else 0.0 | |
| failed_constraints.append(f"{constraint_name}(Cas9 score {cas9_score:.4f} <= {constraint.threshold})") | |
| elif constraint_name == "PAMMatchingProbabilityThreshold": | |
| # Get actual PAM matching probability to show what it was vs threshold | |
| pam_scores = constraint.pam_matching_obj.get_score_for_pam( | |
| [seq_clean], constraint.target_pam, use_temperature_scaling=False | |
| ) | |
| pam_prob = pam_scores[0] if pam_scores else 0.0 | |
| failed_constraints.append( | |
| f"{constraint_name}(PAM matching probability {pam_prob:.4f} <= {constraint.threshold})" | |
| ) | |
| else: | |
| failed_constraints.append(f"{constraint_name}") | |
| if failed_constraints: | |
| print(f" Failed constraints: {', '.join(failed_constraints)}") | |
| else: | |
| print(f" WARNING: logG is -inf but no constraints failed (unexpected!)") | |
| # print("Terminal Time: ", time.time() - start) | |
| _t_sel = _tic() | |
| logG = logG.reshape(survived_candidates.shape[0], num_rollouts) | |
| log_h_hat = torch.logsumexp(logG, dim=1) - math.log(num_rollouts) # (num_survived_candidates,) | |
| _count("cand_all_infeasible", int((log_h_hat == float('-inf')).sum().item())); _count("cand_evaluated", int(log_h_hat.numel())) | |
| idx = (logG.max(dim=1).values > curr_logG).nonzero(as_tuple=True)[0] | |
| final_survived_candidates = survived_candidates[idx, :] | |
| if len(final_survived_candidates) == 0: | |
| _count("step_no_feasible"); _toc("selection", _t_sel) | |
| continue | |
| # DEBUG: Print final selection scores (compact) | |
| if len(idx) > 0: | |
| top_logG_vals = [logG.max(dim=1).values[cand_idx].item() for cand_idx in idx[:3]] | |
| print(f"[FINAL_SELECTION step={step}] Improved: {len(idx)}/{len(survived_candidates)} | " | |
| f"Top logG: {', '.join([f'{v:.4f}' for v in top_logG_vals])}") | |
| # Doob-like transform | |
| log_h_hat = log_h_hat[idx] | |
| base_rates_t = torch.tensor([base_rates[i] for i in idx.tolist()], device=device, dtype=torch.float32) | |
| log_base = 0.5 * torch.log(base_rates_t.clamp_min(1e-30)) | |
| log_weights = log_base + log_h_hat | |
| probs = torch.softmax(log_weights, dim=0) | |
| # if torch.isnan(probs).any(): | |
| # pdb.set_trace() | |
| selected_idx = torch.multinomial(probs, 1).item() | |
| print(f"[FINAL_SELECTION step={step}] Selected: cand {idx[selected_idx].item()} (prob={probs[selected_idx].item():.4f})") | |
| x = final_survived_candidates[selected_idx].unsqueeze(0) | |
| _toc("selection", _t_sel) | |
| # Reprint debug info for the selected candidate | |
| selected_edit_stats = survived_edit_stats[idx[selected_idx]] | |
| num_ins = selected_edit_stats['num_ins'][0] | |
| num_del = selected_edit_stats['num_del'][0] | |
| num_sub = selected_edit_stats['num_sub'][0] | |
| total_edits = selected_edit_stats['total_edits'][0] | |
| avg_ins_rate, avg_del_rate, avg_sub_rate = selected_edit_stats['avg_rates'][0] if 'avg_rates' in selected_edit_stats and selected_edit_stats['avg_rates'] else (0.0, 0.0, 0.0) | |
| # Accumulate edit statistics for selected steps | |
| total_ins += num_ins | |
| total_del += num_del | |
| total_sub += num_sub | |
| # DEBUG: Print final logG calculation for selected candidate (compact) | |
| selected_seq_str = tokenizer.batch_decode(x, skip_special_tokens=True)[0].replace(" ", "") | |
| selected_obj_scores, obj_names = extract_objective_vector([selected_seq_str], objective_models, device, return_names=True) | |
| selected_obj_scores = selected_obj_scores.squeeze(0) | |
| selected_logG = logG.max(dim=1).values[idx[selected_idx]].item() | |
| # Find DeletionCount objective for absolute count | |
| deletion_obj = None | |
| for obj in objective_models: | |
| if hasattr(obj, 'original_length') and hasattr(obj, 'max_deletion'): | |
| deletion_obj = obj | |
| break | |
| # Compute logG components | |
| diff = selected_obj_scores - ref_z | |
| w_diff = w * diff | |
| term1 = torch.min(w_diff).item() | |
| term2 = (rho * torch.sum(w_diff)).item() | |
| u_atc_val = term1 + term2 | |
| logG_calc = beta_t * u_atc_val | |
| # Format output | |
| diff_parts = [] | |
| w_diff_parts = [] | |
| for obj_idx, obj_name in enumerate(obj_names): | |
| raw_score = selected_obj_scores[obj_idx].item() | |
| diff_val = diff[obj_idx].item() | |
| w_diff_val = w_diff[obj_idx].item() | |
| weight = w[obj_idx].item() | |
| ref_val = ref_z[obj_idx].item() | |
| if obj_name == 'deletion_count' and deletion_obj is not None: | |
| abs_deletion = deletion_obj.original_length - len(selected_seq_str) | |
| diff_parts.append(f"{obj_name}: {raw_score:.4f} (abs: {abs_deletion}) - {ref_val:.4f} = {diff_val:.4f}") | |
| else: | |
| diff_parts.append(f"{obj_name}: {raw_score:.4f} - {ref_val:.4f} = {diff_val:.4f}") | |
| w_diff_parts.append(f"{obj_name}: {diff_val:.4f} × {weight:.3f} = {w_diff_val:.4f}") | |
| # Always print selected step info, even if no edits | |
| edit_info = "" | |
| if total_edits > 0: | |
| pct_ins = 100.0 * num_ins / total_edits | |
| pct_del = 100.0 * num_del / total_edits | |
| pct_sub = 100.0 * num_sub / total_edits | |
| edit_info = f" | Edits: ins={num_ins}({pct_ins:.0f}%), del={num_del}({pct_del:.0f}%), sub={num_sub}({pct_sub:.0f}%)" | |
| else: | |
| edit_info = " | No edits" | |
| print(f"[SELECTED step={step}] len={len(selected_seq_str)} | diffs=[{', '.join(diff_parts)}] | " | |
| f"w×diffs=[{', '.join(w_diff_parts)}] | min={term1:.4f}, rho×sum={term2:.4f} | " | |
| f"u_atc={u_atc_val:.4f}, logG={selected_logG:.4f} (calc: {logG_calc:.4f}){edit_info}") | |
| protein_seq = tokenizer.batch_decode(x.tolist(), skip_special_tokens=True)[0] | |
| # print(protein_seq) # Commented out: don't print full sequence at every step | |
| print("Current Length: ", len(protein_seq.replace(' ', ''))) | |
| compute_scores_print([protein_seq], objective_models, constraint_models, device) | |
| last_timestep = step | |
| # finalize | |
| x_final_rollout, logG_final_rollout = _finalize_from_last( | |
| model, | |
| x, | |
| time_grid, | |
| last_timestep, | |
| pad_id, | |
| bos_id, | |
| eos_id, | |
| allowed_tokens, | |
| objective_models, | |
| constraint_models, | |
| w, | |
| rho, | |
| ref_z, | |
| beta_end, | |
| max_len_cap=max_len_cap, | |
| num_final_rollouts=num_final_rollouts, | |
| num_steps=num_steps, | |
| tokenizer=tokenizer, | |
| pam_masker=pam_masker, | |
| deletion_rate_scale=deletion_rate_scale, | |
| zero_lam_ins=zero_lam_ins, | |
| ) | |
| # Finalization terminals are scored at beta_end (the largest beta), so the | |
| # same normalisation is required here or the finalization rollout wins on | |
| # its beta rather than on its utility. | |
| final_score = _incumbent_score(logG_final_rollout, beta_end) | |
| if torch.isfinite(logG_final_rollout) and (best_terminal_score == float('-inf') or final_score >= best_terminal_score): | |
| best_terminal = x_final_rollout | |
| best_terminal_score = final_score | |
| if best_terminal is None: | |
| _count("run_no_feasible") | |
| print("[pCoMol] No constraint-satisfying terminal; returning original sequence.") | |
| best_terminal = x0.clone().to(device) | |
| # Print cumulative edit statistics for all selected steps | |
| total_all_edits = total_ins + total_del + total_sub | |
| # Compute actual net deletion count from final sequence | |
| final_seq_str = tokenizer.batch_decode(best_terminal, skip_special_tokens=True)[0].replace(" ", "") | |
| final_seq_len = len(final_seq_str) | |
| # Find DeletionCount objective to get original length | |
| deletion_obj = None | |
| for obj in objective_models: | |
| if hasattr(obj, 'original_length') and hasattr(obj, 'max_deletion'): | |
| deletion_obj = obj | |
| break | |
| net_deletion_count = None | |
| if deletion_obj is not None: | |
| net_deletion_count = deletion_obj.original_length - final_seq_len | |
| print(f"\n{'='*80}") | |
| print(f"[FINAL SUMMARY] Cumulative edit statistics for all selected steps:") | |
| print(f" NOTE: These statistics only count edits in selected candidate steps during generation,") | |
| print(f" not including any edits made during finalization rollouts.") | |
| if total_all_edits > 0: | |
| pct_ins_total = 100.0 * total_ins / total_all_edits | |
| pct_del_total = 100.0 * total_del / total_all_edits | |
| pct_sub_total = 100.0 * total_sub / total_all_edits | |
| print(f" Total insertions: {total_ins} ({pct_ins_total:.1f}%)") | |
| print(f" Total deletions: {total_del} ({pct_del_total:.1f}%)") | |
| print(f" Total substitutions: {total_sub} ({pct_sub_total:.1f}%)") | |
| print(f" Total edits: {total_all_edits}") | |
| else: | |
| print(f" No edits were made across all selected steps") | |
| if net_deletion_count is not None: | |
| print(f"\n Net deletion count (from final sequence length): {net_deletion_count}") | |
| print(f" Original length: {deletion_obj.original_length}") | |
| print(f" Final length: {final_seq_len}") | |
| print(f" Difference: {net_deletion_count} (may differ from edit statistics due to") | |
| print(f" insertions reducing net deletions and finalization edits)") | |
| print(f"{'='*80}\n") | |
| return best_terminal | |
| # --------------------------------------------------------------------------- | |
| # main | |
| # --------------------------------------------------------------------------- | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--pcomol_config", type=str, required=True) | |
| parser.add_argument("--ckpt", type=str, required=True) | |
| parser.add_argument("--input", type=str, required=True) | |
| parser.add_argument("--num_steps", type=int, default=32) | |
| parser.add_argument("--max_len_cap", type=int, default=None) | |
| parser.add_argument("--num_candidates", type=int, default=10) | |
| parser.add_argument("--num_rollouts", type=int, default=5) | |
| parser.add_argument("--beta_start", type=float, default=1.0) | |
| parser.add_argument("--beta_end", type=float, default=3.0) | |
| parser.add_argument("--num_final_rollouts", type=int, default=50) | |
| parser.add_argument( | |
| "--disable_guidance", | |
| action="store_true", | |
| help=( | |
| "Disable pCoMol objective-guided candidate selection and use unguided " | |
| "EditFlows rollouts for generation. Objectives/constraints are still " | |
| "initialized and used for final scoring/output." | |
| ), | |
| ) | |
| parser.add_argument("--deletion_rate_scale", type=float, default=1500.0, | |
| help="Scaling factor for deletion rate during candidate generation and rollouts (default: 1500.0)") | |
| parser.add_argument("--num_sequences", type=int, default=100, | |
| help="Number of sequences to generate overall (default: 100)") | |
| parser.add_argument("--time_limit_s", type=float, default=None, | |
| help="Optional wall-clock generation budget in seconds. If set, generate until this time expires or --num_sequences is reached.") | |
| parser.add_argument("--objective_weights", type=float, nargs='+') | |
| parser.add_argument("--ref_z", type=float, nargs='+') | |
| parser.add_argument("--rho", type=float, default=1) | |
| parser.add_argument("--output_file", type=str, default=None) | |
| # Cas9 domain completeness constraint arguments | |
| parser.add_argument("--cas9_hmm_db", type=str, default=None, | |
| help="Path to HMM database for Cas9 domain detection (e.g., cas9_bootstrap_pfam.hmm). If not provided, domain detection constraint will be disabled.") | |
| parser.add_argument("--cas9_evalue", type=float, default=1e-2, | |
| help="E-value cutoff for domain detection") | |
| parser.add_argument("--cas9_minlen", type=int, default=35, | |
| help="Minimum domain length for detection") | |
| # Objective flags (explicit opt-in) | |
| parser.add_argument("--cas9_objective", action="store_true", | |
| help="Include Cas9 classification objective.") | |
| parser.add_argument("--PAM_distr_objective", "--pam_distr_objective", | |
| dest="pam_distr_objective", action="store_true", | |
| help="Include PAM distribution matching objective.") | |
| parser.add_argument("--deletion_count_obj", "--delection_count_obj", | |
| dest="deletion_count_obj", action="store_true", | |
| help="Include deletion count objective.") | |
| # Cas9 classifier arguments | |
| parser.add_argument("--cas9_classifier_ckpt", type=str, default=None, | |
| help="Path to Cas9 classifier checkpoint (default: uses default path)") | |
| parser.add_argument("--cas9_classifier_config", type=str, default=None, | |
| help="Path to Cas9 classifier config YAML (optional)") | |
| parser.add_argument("--cas9_score_threshold", type=float, default=None, | |
| help="Minimum Cas9 classifier score threshold (0-1). If provided, adds a terminal constraint requiring sequences to have Cas9 score > threshold. Default: None (no threshold constraint)") | |
| # PAM matching objective arguments | |
| parser.add_argument("--target_pam", type=str, default=None, | |
| help="Target PAM sequence (10 nucleotides, e.g., 'NGGNNNNNNN'). If set to 'matching', will predict the PAM from the input sequence and use that as the target. If provided, enables PAM matching objective.") | |
| parser.add_argument("--pam_model_name", type=str, default="Profluent-Bio/protein2pam-cas9_full", | |
| help="HuggingFace model name for PAM prediction (default: Profluent-Bio/protein2pam-cas9_full)") | |
| parser.add_argument("--pam_no_entropy", action="store_false", dest="pam_use_entropy", default=True, | |
| help="Disable entropy-based scoring for N positions. Score only considers log likelihood of target PAM at non-N positions. By default, entropy scoring is enabled for N positions.") | |
| parser.add_argument("--use_ce_loss", action="store_true", | |
| help="Use cross-entropy loss for PAM matching objective instead of log probability approach. Score = exp(-mean_ce_loss) to convert to [0, 1] range.") | |
| parser.add_argument("--pam_min_confidence", type=float, default=0.55, | |
| help="Minimum probability threshold for PAM prediction. If max probability < pam_min_confidence, predict 'N' instead of specific nucleotide (default: 0.55)") | |
| parser.add_argument("--pam_prediction_temperature", type=float, default=1.0, | |
| help="Temperature for PAM prediction. Values < 1.0 make distributions sharper (more confident), > 1.0 make them softer. Default: 1.0 (no temperature scaling). Note: This only affects prediction, not scoring.") | |
| parser.add_argument("--pam_probability_threshold", type=float, default=None, | |
| help="Minimum PAM matching probability threshold (0-1). If provided, adds a terminal constraint requiring sequences to have PAM matching probability > threshold for the target PAM. Default: None (no threshold constraint)") | |
| # Deletion count objective arguments | |
| parser.add_argument("--max_deletion_percentage", type=float, default=None, | |
| help="Maximum deletion percentage (0-1) for normalization. The deletion count objective will return a value between 0 and 1, representing the percentage of (original_length * max_deletion_percentage) that has been deleted. If not provided, defaults to 1.0 (allowing 100%% deletion).") | |
| # Protein length constraint | |
| parser.add_argument("--allow_equal_length", action="store_true", | |
| help="Relax ProteinLength so final sequences may equal the input length (L0), not only strictly shorter. Still requires length >= L0//2.") | |
| # Target length constraint arguments | |
| parser.add_argument("--target_length", type=int, default=None, | |
| help="Target sequence length in amino acids. If provided, adds a terminal constraint requiring the final sequence length to exactly match this value.") | |
| parser.add_argument("--min_target_length", type=float, default=None, | |
| help="Minimum target sequence length in amino acids (inclusive). If provided as a decimal (0-1), interpreted as a percentage of input sequence length. If provided as an integer (>=1), used as absolute value. If provided, adds a terminal constraint requiring the final sequence length to be greater than or equal to this value.") | |
| parser.add_argument("--max_target_length", type=float, default=None, | |
| help="Maximum target sequence length in amino acids (inclusive). If provided as a decimal (0-1), interpreted as a percentage of input sequence length. If provided as an integer (>=1), used as absolute value. If provided, adds a terminal constraint requiring the final sequence length to be less than or equal to this value.") | |
| # PAM/PI-domain masking arguments | |
| parser.add_argument("--pam_hmm_db", type=str, default=None, | |
| help="Path to cas9_pi.hmm (mini Pfam DB with Cas9_PI models). Required if --pam_mask or --pam_scale_edits is set.") | |
| parser.add_argument("--pam_mask", action="store_true", | |
| help="Enable PAM/PI domain masking (blocks all edits in PI domain region). Requires --pam_hmm_db.") | |
| parser.add_argument("--pam_mask_max_len", type=int, default=200, | |
| help="Max number of AA positions to hard-mask (Option 1 core window cap).") | |
| parser.add_argument("--pam_mask_min_len", type=int, default=0, | |
| help="Minimum number of AA positions to mask (0 = no minimum). If detected domain is shorter, it will be expanded to this length (up to max_mask_len).") | |
| parser.add_argument("--pam_evalue", type=float, default=1e-5, | |
| help="Per-domain i-evalue cutoff for PI hits.") | |
| parser.add_argument("--pam_refresh_every", type=int, default=1, | |
| help="Recompute PI deletion mask every N accepted steps (>=1).") | |
| parser.add_argument("--hmmscan_bin", type=str, default="hmmscan", | |
| help="hmmscan executable (default: hmmscan).") | |
| parser.add_argument("--hmmscan_cpu", type=int, default=1, | |
| help="CPUs to give hmmscan.") | |
| parser.add_argument("--pam_debug", action="store_true", | |
| help="Print PI masking debug info during generation.") | |
| parser.add_argument("--pam_scale_edits", action="store_true", | |
| help="If set, scale up insertion/substitution rates in PAM region instead of masking edits. Requires --pam_hmm_db.") | |
| parser.add_argument("--pam_edit_scale_factor", type=float, default=10.0, | |
| help="Scaling factor for insertion/substitution rates in PAM region when --pam_scale_edits is enabled (default: 10.0).") | |
| parser.add_argument("--zero_lam_ins", action="store_true", | |
| help="If set, set lam_ins to 0.0 before processing (disables insertion operations).") | |
| parser.add_argument("--PID_PAM_prediction", action="store_true", | |
| help="If set, detect PAM domain using HMM and use only the detected PAM domain region for PAM matching objectives/constraints. Also changes model to 'Profluent-Bio/protein2pam-cas9'. Requires --pam_hmm_db.") | |
| parser.add_argument("--legacy_beta_incumbent", action="store_true", | |
| help="Restore the pre-fix incumbent rule, which ranked observed feasible " | |
| "terminals by beta_t*U instead of U. Because beta_t is annealed " | |
| "beta_start->beta_end, that ranked terminals by when they were found " | |
| "rather than by utility. Default (flag absent) normalises by beta so " | |
| "the returned design maximises the fixed G of Eq. (9), per Prop. C.6.") | |
| # --- budget-matched unguided baseline (Doob-h guidance removed entirely) ------- | |
| parser.add_argument("--rejection_baseline", action="store_true", | |
| help="Run unguided Edit Flow + terminal rejection instead of pCoMol: draw " | |
| "--oracle_budget i.i.d. terminals from the base kernel starting at the " | |
| "input, score each once with the same objectives/constraints, and keep " | |
| "the best feasible one. Writes a per-terminal pool CSV.") | |
| parser.add_argument("--oracle_budget", type=int, default=1000, | |
| help="Number of terminal sequences to draw and score in --rejection_baseline " | |
| "mode. One drawn terminal == one oracle call.") | |
| parser.add_argument("--rejection_batch", type=int, default=50, | |
| help="Terminals sampled/scored per GPU batch in --rejection_baseline mode " | |
| "(default 50, matching pCoMol's finalization rollout batch).") | |
| parser.add_argument("--pool_out", type=str, default=None, | |
| help="Per-terminal pool CSV path for --rejection_baseline (default: " | |
| "<output_file>_pool.csv).") | |
| parser.add_argument("--instrument", action="store_true", | |
| help="Record per-sequence runtime breakdown by component + rollout-failure counts.") | |
| parser.add_argument("--instr_out", type=str, default=None, | |
| help="Where to write the instrumentation JSON (default: <output_file>_instr.json).") | |
| parser.add_argument("--seed", type=int, default=None, | |
| help="RNG seed for reproducibility / seed-level robustness studies. Sets python/numpy/torch(+cuda) seeds. If unset, generation is unseeded as before.") | |
| args = parser.parse_args() | |
| # Validate that pam_hmm_db is provided if pam_mask or pam_scale_edits is enabled | |
| if args.pam_mask and args.pam_hmm_db is None: | |
| raise ValueError("--pam_mask requires --pam_hmm_db to be specified") | |
| if args.pam_scale_edits and args.pam_hmm_db is None: | |
| raise ValueError("--pam_scale_edits requires --pam_hmm_db to be specified") | |
| if args.pam_mask and args.pam_scale_edits: | |
| raise ValueError("--pam_mask and --pam_scale_edits cannot be used together. Choose one.") | |
| if args.PID_PAM_prediction and args.pam_hmm_db is None: | |
| raise ValueError("--PID_PAM_prediction requires --pam_hmm_db to be specified") | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| if args.seed is not None: | |
| import random as _random, numpy as _np | |
| _random.seed(args.seed) | |
| _np.random.seed(args.seed) | |
| torch.manual_seed(args.seed) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(args.seed) | |
| print(f"[SEED] Set python/numpy/torch RNG seeds to {args.seed}") | |
| with open(args.pcomol_config, "r") as f: | |
| cfg = edict(yaml.safe_load(f)) | |
| editflow, source_dist, tokenizer, pad_id, bos_id, eos_id, eps_id = build_model_and_stuff(cfg, device) | |
| ckpt = torch.load(args.ckpt, map_location=device) | |
| editflow.load_state_dict(ckpt["state_dict"], strict=False) | |
| model = editflow.model.to(device) | |
| model.eval() | |
| x0 = tokenize_input_str(args.input, cfg, tokenizer, bos_id, eos_id, pad_id, device) | |
| allowed_tokens = torch.tensor( | |
| [tok for tok in source_dist._allowed_tokens if tok not in (eps_id,)], | |
| device=device, | |
| dtype=torch.long, | |
| ) | |
| # Initialize PAM domain detector early if PID_PAM_prediction is enabled | |
| pam_domain_interval = None # (start, end) 1-based coordinates | |
| if args.PID_PAM_prediction: | |
| print(f"[PID_PAM_prediction] Detecting PAM domain from input sequence...") | |
| hmmscan_runner = HMMSCAN( | |
| hmm_db_path=args.pam_hmm_db, | |
| hmmscan_bin=args.hmmscan_bin, | |
| cpu=args.hmmscan_cpu, | |
| ) | |
| pam_detector = Cas9PIMasker( | |
| hmmscan=hmmscan_runner, | |
| use_env_coords=False, | |
| max_mask_len=args.pam_mask_max_len, | |
| min_mask_len=args.pam_mask_min_len, | |
| evalue_cutoff=args.pam_evalue, | |
| min_ali_len=30, | |
| fallback_last_n=None, | |
| cache_size=2048, | |
| ) | |
| input_seq_clean = args.input.replace(" ", "").strip() | |
| pam_domain_interval = pam_detector.pi_core_interval(input_seq_clean) | |
| if pam_domain_interval is None: | |
| raise ValueError(f"[PID_PAM_prediction] Failed to detect PAM domain in input sequence. Cannot proceed.") | |
| start_1based, end_1based = pam_domain_interval | |
| # Extract PAM domain region (convert 1-based to 0-based for slicing) | |
| pam_domain_seq = input_seq_clean[start_1based - 1:end_1based] | |
| print(f"[PID_PAM_prediction] Detected PAM domain: positions {start_1based}-{end_1based} (1-based), length={len(pam_domain_seq)}") | |
| print(f"[PID_PAM_prediction] PAM domain sequence: {pam_domain_seq[:50]}...{pam_domain_seq[-50:] if len(pam_domain_seq) > 100 else pam_domain_seq}") | |
| # Override model name for PAM prediction | |
| args.pam_model_name = "Profluent-Bio/protein2pam-cas9" | |
| print(f"[PID_PAM_prediction] Using model: {args.pam_model_name} (changed from default)") | |
| cas9_classifier = None | |
| if args.cas9_objective or args.cas9_score_threshold is not None: | |
| # Initialize Cas9 classifier objective with shared ESM model | |
| # This avoids loading ESM twice, saving GPU memory | |
| cas9_classifier = Cas9Classification( | |
| device=device, | |
| checkpoint_path=args.cas9_classifier_ckpt, | |
| config_path=args.cas9_classifier_config, | |
| shared_esm_model=model.esm_emb # Share ESM model from editflow | |
| ) | |
| objective_models = [] | |
| if args.cas9_objective: | |
| objective_models.append(cas9_classifier) | |
| if args.deletion_count_obj: | |
| # Initialize deletion count objective (maximizes number of deletions, normalized to [0, 1]) | |
| deletion_count_obj = DeletionCount( | |
| original_seq=args.input, | |
| max_deletion_percentage=args.max_deletion_percentage, | |
| ) | |
| objective_models.append(deletion_count_obj) | |
| # Initialize PAM matching objective if target PAM is provided | |
| pam_matching_obj_for_constraint = None # Will be set if PAM matching is enabled | |
| if args.target_pam is not None: | |
| # Handle 'matching' mode: predict PAM from input sequence | |
| if args.target_pam.lower() == 'matching': | |
| print(f"Target PAM set to 'matching' - will predict PAM from input sequence") | |
| # Create temporary PAM matching object to predict the input sequence's PAM | |
| # Use PAM domain region if PID_PAM_prediction is enabled | |
| input_seq_for_pam_prediction = args.input | |
| if args.PID_PAM_prediction and pam_domain_interval is not None: | |
| input_seq_clean = args.input.replace(" ", "").strip() | |
| start_1based, end_1based = pam_domain_interval | |
| pam_domain_seq = input_seq_clean[start_1based - 1:end_1based] | |
| input_seq_for_pam_prediction = pam_domain_seq | |
| print(f" Using PAM domain region (positions {start_1based}-{end_1based}) for initial PAM prediction") | |
| temp_pam_obj = PAMMatching( | |
| device=device, | |
| target_pam="NNNNNNNNNN", # Dummy target, we'll replace it | |
| model_name=args.pam_model_name, | |
| use_entropy_for_n_positions=args.pam_use_entropy, | |
| use_ce_loss=args.use_ce_loss, | |
| pam_min_confidence=args.pam_min_confidence, | |
| pam_prediction_temperature=args.pam_prediction_temperature | |
| ) | |
| # Predict PAM for input sequence (or PAM domain region) | |
| predicted_pams = temp_pam_obj.predict_pam([input_seq_for_pam_prediction]) | |
| if predicted_pams: | |
| actual_target_pam = predicted_pams[0] | |
| print(f" Predicted PAM for {'PAM domain region' if args.PID_PAM_prediction else 'input sequence'}: {actual_target_pam}") | |
| print(f" Using this as target PAM for optimization") | |
| else: | |
| raise ValueError("Failed to predict PAM for input sequence") | |
| # Now create the actual PAM matching objective with the predicted PAM | |
| entropy_mode = "enabled" if args.pam_use_entropy else "disabled" | |
| ce_loss_mode = "CE loss" if args.use_ce_loss else "log probability" | |
| print(f" Entropy scoring for N positions: {entropy_mode}") | |
| print(f" Scoring method: {ce_loss_mode}") | |
| pam_matching_obj_base = PAMMatching( | |
| device=device, | |
| target_pam=actual_target_pam, | |
| model_name=args.pam_model_name, | |
| use_entropy_for_n_positions=args.pam_use_entropy, | |
| use_ce_loss=args.use_ce_loss, | |
| pam_min_confidence=args.pam_min_confidence, | |
| pam_prediction_temperature=args.pam_prediction_temperature | |
| ) | |
| # Wrap with PAM domain extractor if PID_PAM_prediction is enabled | |
| if args.PID_PAM_prediction and pam_domain_interval is not None: | |
| pam_matching_obj = PAMDomainWrapper(pam_matching_obj_base, pam_domain_interval) | |
| print(f" Wrapped PAM matching objective with PAM domain extractor (interval: {pam_domain_interval})") | |
| else: | |
| pam_matching_obj = pam_matching_obj_base | |
| if args.pam_distr_objective: | |
| objective_models.append(pam_matching_obj) | |
| print(f"PAM matching objective initialized successfully with target PAM: {actual_target_pam}") | |
| # Store reference to PAM matching object for constraint creation later | |
| pam_matching_obj_for_constraint = pam_matching_obj | |
| else: | |
| # Regular mode: use provided PAM string | |
| print(f"Initializing PAM matching objective with target PAM: {args.target_pam}") | |
| entropy_mode = "enabled" if args.pam_use_entropy else "disabled" | |
| ce_loss_mode = "CE loss" if args.use_ce_loss else "log probability" | |
| print(f" Entropy scoring for N positions: {entropy_mode}") | |
| print(f" Scoring method: {ce_loss_mode}") | |
| pam_matching_obj_base = PAMMatching( | |
| device=device, | |
| target_pam=args.target_pam, | |
| model_name=args.pam_model_name, | |
| use_entropy_for_n_positions=args.pam_use_entropy, | |
| use_ce_loss=args.use_ce_loss, | |
| pam_min_confidence=args.pam_min_confidence, | |
| pam_prediction_temperature=args.pam_prediction_temperature | |
| ) | |
| # Wrap with PAM domain extractor if PID_PAM_prediction is enabled | |
| if args.PID_PAM_prediction and pam_domain_interval is not None: | |
| pam_matching_obj = PAMDomainWrapper(pam_matching_obj_base, pam_domain_interval) | |
| print(f" Wrapped PAM matching objective with PAM domain extractor (interval: {pam_domain_interval})") | |
| else: | |
| pam_matching_obj = pam_matching_obj_base | |
| if args.pam_distr_objective: | |
| objective_models.append(pam_matching_obj) | |
| print(f"PAM matching objective initialized successfully.") | |
| # Store reference to PAM matching object for constraint creation later | |
| pam_matching_obj_for_constraint = pam_matching_obj | |
| num_objectives = len(objective_models) | |
| no_objective_mode = (num_objectives == 0) | |
| if no_objective_mode: | |
| print( | |
| "[INFO] No objectives selected; running in base EditFlows sampling mode " | |
| "(unguided rollouts)." | |
| ) | |
| unguided_generation_mode = args.disable_guidance or no_objective_mode | |
| if args.disable_guidance: | |
| print( | |
| "[INFO] --disable_guidance enabled; generation uses unguided EditFlows " | |
| "rollouts while objective/constraint scores are still reported." | |
| ) | |
| if not args.objective_weights: | |
| if num_objectives > 0: | |
| objective_weights = torch.tensor([1.0 / num_objectives] * num_objectives).to(device) | |
| else: | |
| objective_weights = torch.empty((0,), device=device, dtype=torch.float32) | |
| else: | |
| if num_objectives == 0: | |
| raise ValueError( | |
| "objective_weights provided but no objectives are enabled. " | |
| "Either enable objective flags or remove --objective_weights." | |
| ) | |
| if len(args.objective_weights) != num_objectives: | |
| raise ValueError( | |
| f"objective_weights length ({len(args.objective_weights)}) does not match " | |
| f"number of objectives ({num_objectives})." | |
| ) | |
| objective_weights = torch.tensor(args.objective_weights).to(device) | |
| if not args.ref_z: | |
| ref_z = torch.zeros(num_objectives).to(device) | |
| else: | |
| if num_objectives == 0: | |
| raise ValueError( | |
| "ref_z provided but no objectives are enabled. " | |
| "Either enable objective flags or remove --ref_z." | |
| ) | |
| ref_z = torch.tensor(args.ref_z).to(device) | |
| # Initialize protein length constraint | |
| protein_length_constraint = ProteinLength( | |
| orig_protein_seq=args.input, | |
| cfg=cfg, | |
| tokenizer=tokenizer, | |
| allow_equal_length=args.allow_equal_length, | |
| ) | |
| constraint_models = [protein_length_constraint] | |
| # Initialize Cas9 domain completeness constraint (optional) | |
| if args.cas9_hmm_db is not None: | |
| cas9_domain_constraint = Cas9DomainCompleteness( | |
| hmm_path=args.cas9_hmm_db, | |
| ev1=args.cas9_evalue, | |
| minlen1=args.cas9_minlen, | |
| cpu=1, # Can be made configurable if needed | |
| ) | |
| constraint_models.append(cas9_domain_constraint) | |
| print(f"Cas9 domain completeness constraint enabled with HMM database: {args.cas9_hmm_db}") | |
| else: | |
| print("Cas9 domain completeness constraint disabled (--cas9_hmm_db not provided)") | |
| # Initialize target length constraint if specified | |
| if args.target_length is not None: | |
| target_length_constraint = TargetLength(target_length=args.target_length) | |
| constraint_models.append(target_length_constraint) | |
| print(f"Target length constraint enabled: sequences must have exactly {args.target_length} amino acids") | |
| # Initialize min target length constraint if specified | |
| computed_min_target_length = None | |
| if args.min_target_length is not None: | |
| # If min_target_length is a decimal (0 < value < 1), interpret as percentage of input length | |
| # If it's >= 1, use as absolute value | |
| input_length = len(args.input.replace(" ", "").replace("\n", "")) | |
| if 0 < args.min_target_length < 1: | |
| # Decimal: interpret as percentage | |
| computed_min_target_length = int(args.min_target_length * input_length) | |
| print(f"Min target length constraint: {args.min_target_length} (decimal) interpreted as {args.min_target_length*100:.1f}% of input length ({input_length}) = {computed_min_target_length}") | |
| else: | |
| # Integer: use as absolute value | |
| computed_min_target_length = int(args.min_target_length) | |
| print(f"Min target length constraint: {args.min_target_length} (integer) used as absolute value = {computed_min_target_length}") | |
| min_target_length_constraint = MinTargetLength(min_target_length=computed_min_target_length) | |
| constraint_models.append(min_target_length_constraint) | |
| print(f"Min target length constraint enabled: sequences must have length >= {computed_min_target_length} amino acids") | |
| # Initialize max target length constraint if specified | |
| computed_max_target_length = None | |
| if args.max_target_length is not None: | |
| # If max_target_length is a decimal (0 < value < 1), interpret as percentage of input length | |
| # If it's >= 1, use as absolute value | |
| input_length = len(args.input.replace(" ", "").replace("\n", "")) | |
| if 0 < args.max_target_length < 1: | |
| # Decimal: interpret as percentage | |
| computed_max_target_length = int(args.max_target_length * input_length) | |
| print(f"Max target length constraint: {args.max_target_length} (decimal) interpreted as {args.max_target_length*100:.1f}% of input length ({input_length}) = {computed_max_target_length}") | |
| else: | |
| # Integer: use as absolute value | |
| computed_max_target_length = int(args.max_target_length) | |
| print(f"Max target length constraint: {args.max_target_length} (integer) used as absolute value = {computed_max_target_length}") | |
| max_target_length_constraint = MaxTargetLength(max_target_length=computed_max_target_length) | |
| constraint_models.append(max_target_length_constraint) | |
| print(f"Max target length constraint enabled: sequences must have length <= {computed_max_target_length} amino acids") | |
| # Validate explicit length range if both bounds are provided | |
| if computed_min_target_length is not None and computed_max_target_length is not None: | |
| if computed_min_target_length > computed_max_target_length: | |
| raise ValueError( | |
| f"Incompatible length range: min_target_length ({computed_min_target_length}) " | |
| f"must be <= max_target_length ({computed_max_target_length})." | |
| ) | |
| # Initialize PAM matching constraint if target_pam is specified | |
| if args.target_pam is not None and pam_matching_obj_for_constraint is not None: | |
| pam_matching_constraint = PAMMatchingConstraint(pam_matching_obj_for_constraint) | |
| constraint_models.append(pam_matching_constraint) | |
| print(f"PAM matching constraint enabled: sequences must have predicted PAM matching target PAM") | |
| # Initialize PAM matching probability threshold constraint if threshold is specified | |
| if args.pam_probability_threshold is not None: | |
| if pam_matching_obj_for_constraint is None: | |
| raise ValueError("--pam_probability_threshold requires --target_pam to enable PAM matching evaluation.") | |
| pam_probability_threshold_constraint = PAMMatchingProbabilityThreshold( | |
| pam_matching_obj_for_constraint, args.pam_probability_threshold | |
| ) | |
| constraint_models.append(pam_probability_threshold_constraint) | |
| print( | |
| "PAM matching probability threshold constraint enabled: sequences must have " | |
| f"PAM matching probability > {args.pam_probability_threshold}" | |
| ) | |
| # Initialize Cas9 score threshold constraint if threshold is specified | |
| if args.cas9_score_threshold is not None: | |
| if cas9_classifier is None: | |
| raise ValueError("--cas9_score_threshold requires --cas9_objective or a valid Cas9 classifier.") | |
| cas9_score_threshold_constraint = Cas9ScoreThreshold(cas9_classifier, args.cas9_score_threshold) | |
| constraint_models.append(cas9_score_threshold_constraint) | |
| print(f"Cas9 score threshold constraint enabled: sequences must have Cas9 score > {args.cas9_score_threshold}") | |
| # Initialize PAM/PI masker if --pam_mask flag is set | |
| pam_masker = None | |
| if args.pam_mask: | |
| if args.pam_hmm_db is None: | |
| raise ValueError("--pam_mask requires --pam_hmm_db to be specified") | |
| hmmscan_runner = HMMSCAN( | |
| hmm_db_path=args.pam_hmm_db, | |
| hmmscan_bin=args.hmmscan_bin, | |
| cpu=args.hmmscan_cpu, | |
| ) | |
| pam_masker = Cas9PIMasker( | |
| hmmscan=hmmscan_runner, | |
| use_env_coords=False, # ali coords (smaller) | |
| max_mask_len=args.pam_mask_max_len, | |
| min_mask_len=args.pam_mask_min_len, | |
| evalue_cutoff=args.pam_evalue, | |
| min_ali_len=30, | |
| fallback_last_n=None, # avoid masking arbitrary tail when no hit | |
| cache_size=2048, | |
| ) | |
| print(f"[PI mask] enabled: db={args.pam_hmm_db}, max_len={args.pam_mask_max_len}, min_len={args.pam_mask_min_len}, evalue<={args.pam_evalue} (blocks all edits: ins/del/sub)") | |
| elif args.pam_scale_edits: | |
| if args.pam_hmm_db is None: | |
| raise ValueError("--pam_scale_edits requires --pam_hmm_db to be specified") | |
| # For scaling mode, we still need the masker to identify the PAM region | |
| hmmscan_runner = HMMSCAN( | |
| hmm_db_path=args.pam_hmm_db, | |
| hmmscan_bin=args.hmmscan_bin, | |
| cpu=args.hmmscan_cpu, | |
| ) | |
| pam_masker = Cas9PIMasker( | |
| hmmscan=hmmscan_runner, | |
| use_env_coords=False, # ali coords (smaller) | |
| max_mask_len=args.pam_mask_max_len, | |
| min_mask_len=args.pam_mask_min_len, | |
| evalue_cutoff=args.pam_evalue, | |
| min_ali_len=30, | |
| fallback_last_n=None, # avoid masking arbitrary tail when no hit | |
| cache_size=2048, | |
| ) | |
| print(f"[PI scaling] enabled: db={args.pam_hmm_db}, max_len={args.pam_mask_max_len}, min_len={args.pam_mask_min_len}, evalue<={args.pam_evalue} (scales ins/sub rates by {args.pam_edit_scale_factor}x)") | |
| print("Initial Scores:") | |
| input_scores = compute_scores_print([args.input], objective_models, constraint_models, device, return_scores=True).squeeze(0) | |
| # ----------------------------------------------------------------------- | |
| # Budget-matched baseline: unguided Edit Flow + terminal rejection. | |
| # No candidate proposal, no rollouts, no Doob-h tilt -- just i.i.d. terminals | |
| # from the base kernel started at the input, each scored once by the same | |
| # oracle, keeping the best feasible one. Because x0 is fixed and the base | |
| # kernel carries no run-level state, the drawn terminals are i.i.d., so this | |
| # pool characterises the sampler at every budget <= --oracle_budget. | |
| # ----------------------------------------------------------------------- | |
| if args.rejection_baseline: | |
| import os, csv, json as _json | |
| _instr_reset() | |
| rho_pool = 0.5 # same value pCoMol() is invoked with below | |
| obj_names_pool = [] | |
| for obj in objective_models: | |
| _nm, _ = obj(protein_tokens=None, protein_seqs=[args.input]) | |
| obj_names_pool.append(_nm) | |
| cons_names = [c.__class__.__name__ for c in constraint_models] | |
| orig_len = len(args.input.replace(' ', '')) | |
| pool_path = args.pool_out or ( | |
| (args.output_file.rsplit('.', 1)[0] + '_pool.csv') if args.output_file else 'rejection_pool.csv' | |
| ) | |
| os.makedirs(os.path.dirname(os.path.abspath(pool_path)) or '.', exist_ok=True) | |
| header = (["idx", "final_len", "length_diff"] | |
| + [f"{n}_score" for n in obj_names_pool] | |
| + ["utility", "feasible"] | |
| + [f"pass_{n}" for n in cons_names] | |
| + ["sequence"]) | |
| with open(pool_path, 'w', newline='') as _f: | |
| csv.writer(_f).writerow(header) | |
| print(f"[REJECTION] unguided Edit Flow + terminal rejection | budget={args.oracle_budget} " | |
| f"terminals | batch={args.rejection_batch} | steps={args.num_steps}") | |
| print(f"[REJECTION] objectives={obj_names_pool} constraints={cons_names}") | |
| time_grid = torch.linspace(0.0, 1.0, steps=args.num_steps, device=device) | |
| n_done = 0 | |
| n_feasible = 0 | |
| first_feasible_idx = None | |
| best_u = float('-inf') | |
| best_seq_str = None | |
| t_sample_tot = 0.0 | |
| t_oracle_tot = 0.0 | |
| pool_t0 = time.time() | |
| while n_done < args.oracle_budget: | |
| k = min(args.rejection_batch, args.oracle_budget - n_done) | |
| _t0 = time.time() | |
| x_T = short_rollout_batch( | |
| model, x0, time_grid, 0, pad_id, bos_id, eos_id, allowed_tokens, | |
| args.max_len_cap, num_rollouts=k, num_steps=args.num_steps, | |
| protected_mask=None, deletion_rate_scale=args.deletion_rate_scale, | |
| zero_lam_ins=args.zero_lam_ins, | |
| ) | |
| if torch.cuda.is_available(): | |
| torch.cuda.synchronize() | |
| t_sample_tot += time.time() - _t0 | |
| _t1 = time.time() | |
| _, _, det = _G_T( | |
| x_T, objective_models, constraint_models, objective_weights, rho_pool, ref_z, | |
| args.beta_end, tokenizer, ws_for_invalid=True, | |
| debug_context="REJECTION_TERMINAL", count_terminal=True, return_details=True, | |
| ) | |
| if torch.cuda.is_available(): | |
| torch.cuda.synchronize() | |
| t_oracle_tot += time.time() - _t1 | |
| f_vals = det["f_vals"] | |
| u_all = _augmented_tchebycheff(f_vals, objective_weights, rho_pool, ref_z) | |
| cons = det["constraint_results"] | |
| feas = (cons == 1).all(dim=0) | |
| seqs = det["protein_seqs"] | |
| rows = [] | |
| for i in range(len(seqs)): | |
| gidx = n_done + i | |
| s = seqs[i] | |
| is_feas = bool(feas[i].item()) | |
| u_i = float(u_all[i].item()) | |
| if is_feas: | |
| n_feasible += 1 | |
| if first_feasible_idx is None: | |
| first_feasible_idx = gidx + 1 # 1-based oracle-call index, exact | |
| if u_i > best_u: | |
| best_u = u_i | |
| best_seq_str = s | |
| rows.append([gidx, len(s), orig_len - len(s)] | |
| + [f"{float(f_vals[i, j].item()):.6f}" for j in range(f_vals.shape[1])] | |
| + [f"{u_i:.6f}", int(is_feas)] | |
| + [int(cons[j, i].item()) for j in range(cons.shape[0])] | |
| + [s]) | |
| with open(pool_path, 'a', newline='') as _f: | |
| csv.writer(_f).writerows(rows) | |
| n_done += k | |
| _best_disp = best_u if best_u > float('-inf') else float('nan') | |
| print(f"[REJECTION] {n_done}/{args.oracle_budget} terminals | " | |
| f"feasible={n_feasible} ({100.0 * n_feasible / n_done:.2f}%) | " | |
| f"best_utility={_best_disp:.4f} | " | |
| f"sample={t_sample_tot:.0f}s oracle={t_oracle_tot:.0f}s", flush=True) | |
| summary = { | |
| "mode": "rejection_baseline", | |
| "input_length": orig_len, | |
| "oracle_budget": int(args.oracle_budget), | |
| "rejection_batch": int(args.rejection_batch), | |
| "num_steps": int(args.num_steps), | |
| "deletion_rate_scale": float(args.deletion_rate_scale), | |
| "zero_lam_ins": bool(args.zero_lam_ins), | |
| "seed": args.seed, | |
| "n_terminals": int(n_done), | |
| "n_feasible": int(n_feasible), | |
| "first_feasible_oracle_call": first_feasible_idx, | |
| "best_utility": (best_u if best_u > float('-inf') else None), | |
| "best_sequence": best_seq_str, | |
| "objective_names": obj_names_pool, | |
| "constraint_names": cons_names, | |
| "objective_weights": [float(v) for v in objective_weights.tolist()], | |
| "ref_z": [float(v) for v in ref_z.tolist()], | |
| "rho": rho_pool, | |
| "time_sample_s": t_sample_tot, | |
| "time_oracle_s": t_oracle_tot, | |
| "wall_s": time.time() - pool_t0, | |
| "instr_counts": dict(INSTR["counts"]), | |
| "pool_csv": pool_path, | |
| } | |
| sum_path = pool_path.rsplit('.', 1)[0] + '_summary.json' | |
| with open(sum_path, 'w') as _f: | |
| _json.dump(summary, _f, indent=2) | |
| print(f"[REJECTION] wrote pool -> {pool_path}") | |
| print(f"[REJECTION] wrote summary -> {sum_path}") | |
| return | |
| # CSV bookkeeping for optional output file | |
| csv_header_parts = None | |
| obj_names = [] | |
| pam_obj = None | |
| generated_csv_rows = [] # Generated rows only (excludes the initial input row) | |
| run_start_time = time.time() | |
| # Write CSV header and input sequence row if file doesn't exist | |
| if args.output_file is not None: | |
| import os | |
| import csv | |
| for obj in objective_models: | |
| name, _ = obj(protein_tokens=None, protein_seqs=[args.input]) | |
| obj_names.append(name) | |
| if hasattr(obj, 'predict_pam'): | |
| pam_obj = obj | |
| csv_header_parts = ["final_len", "length_diff"] | |
| csv_header_parts.extend([f"{name}_score" for name in obj_names]) | |
| csv_header_parts.extend([f"{name}_weight" for name in obj_names]) | |
| if pam_obj is not None: | |
| csv_header_parts.append("predicted_pam") | |
| if hasattr(pam_obj, 'use_ce_loss') and pam_obj.use_ce_loss: | |
| csv_header_parts.append("pam_matching_ce_loss_score") | |
| csv_header_parts.append("pam_matching_probability") | |
| csv_header_parts.append("pam_target_match_binary") | |
| csv_header_parts.extend([ | |
| "length_pass", | |
| "cas9_score_pass", | |
| "pam_prob_pass", | |
| "usable_pass", | |
| "total_sequences_generated", | |
| "usable_sequences_generated", | |
| "avg_usable_generation_time", | |
| "wall_elapsed_s", | |
| ]) | |
| csv_header_parts.append("deletion_rate_scale") | |
| if pam_obj is not None: | |
| csv_header_parts.append("pam_min_confidence") | |
| csv_header_parts.append("generation_time") | |
| csv_header_parts.append("sequence") | |
| def _constraint_passes(sequence): | |
| passes = [] | |
| for constraint in constraint_models: | |
| if hasattr(constraint, 'predict_batch'): | |
| result = constraint.predict_batch([sequence])[0] | |
| else: | |
| result = constraint(None, sequence) | |
| passes.append(bool(result)) | |
| return passes | |
| def _build_csv_row(sequence, final_len, length_diff, objective_scores_only, generation_time): | |
| row = {k: "" for k in csv_header_parts} | |
| row["final_len"] = final_len | |
| row["length_diff"] = length_diff | |
| row["deletion_rate_scale"] = float(args.deletion_rate_scale) | |
| row["generation_time"] = float(generation_time) | |
| row["wall_elapsed_s"] = float(time.time() - run_start_time) | |
| row["sequence"] = sequence.replace(' ', '') | |
| for name, score in zip(obj_names, objective_scores_only): | |
| row[f"{name}_score"] = float(score.item() if torch.is_tensor(score) else score) | |
| for name, weight in zip(obj_names, objective_weights): | |
| row[f"{name}_weight"] = float(weight.item() if torch.is_tensor(weight) else weight) | |
| if pam_obj is not None: | |
| predicted_pams = pam_obj.predict_pam([sequence]) | |
| predicted_pam = predicted_pams[0] if predicted_pams else "" | |
| row["predicted_pam"] = predicted_pam | |
| if hasattr(pam_obj, 'use_ce_loss') and pam_obj.use_ce_loss: | |
| pam_score_idx = None | |
| for i, obj in enumerate(objective_models): | |
| if obj == pam_obj: | |
| pam_score_idx = i | |
| break | |
| if pam_score_idx is not None: | |
| row["pam_matching_ce_loss_score"] = float(objective_scores_only[pam_score_idx].item()) | |
| target_pam = pam_obj.target_pam | |
| prob_score = pam_obj.get_score_for_pam([sequence], target_pam, use_temperature_scaling=False)[0] | |
| row["pam_matching_probability"] = float(prob_score) | |
| row["pam_target_match_binary"] = 1.0 if predicted_pam == target_pam else 0.0 | |
| row["pam_min_confidence"] = float(args.pam_min_confidence) | |
| constraint_passes = _constraint_passes(row["sequence"]) | |
| row["length_pass"] = int(all(constraint_passes[:1])) | |
| if args.target_length is not None: | |
| row["length_pass"] = int(final_len == int(args.target_length)) | |
| if computed_min_target_length is not None and final_len < computed_min_target_length: | |
| row["length_pass"] = 0 | |
| if computed_max_target_length is not None and final_len > computed_max_target_length: | |
| row["length_pass"] = 0 | |
| if args.cas9_score_threshold is not None and row.get("cas9_score", "") != "": | |
| row["cas9_score_pass"] = int(float(row["cas9_score"]) > float(args.cas9_score_threshold)) | |
| if args.pam_probability_threshold is not None and row.get("pam_matching_probability", "") != "": | |
| row["pam_prob_pass"] = int(float(row["pam_matching_probability"]) > float(args.pam_probability_threshold)) | |
| row["usable_pass"] = int(all(constraint_passes)) | |
| return row | |
| if not os.path.exists(args.output_file): | |
| with open(args.output_file, 'w', newline='') as f: | |
| writer = csv.DictWriter(f, fieldnames=csv_header_parts) | |
| writer.writeheader() | |
| input_len = len(args.input.replace(' ', '')) | |
| input_objective_scores_only = input_scores[:len(objective_models)] | |
| input_row = _build_csv_row( | |
| sequence=args.input, | |
| final_len=input_len, | |
| length_diff=0, | |
| objective_scores_only=input_objective_scores_only, | |
| generation_time=0.0, | |
| ) | |
| writer.writerow(input_row) | |
| valid = 0 | |
| target_valid = args.num_sequences # how many successful designs you want | |
| attempt = 0 | |
| max_attempts = max(500, target_valid) # safety cap so you don't infinite-loop if it keeps OOM'ing | |
| instr_records = [] | |
| while valid < target_valid and attempt < max_attempts: | |
| if args.time_limit_s is not None and (time.time() - run_start_time) >= args.time_limit_s: | |
| print(f"[TIME_LIMIT] Reached {args.time_limit_s}s generation budget.") | |
| break | |
| attempt += 1 | |
| try: | |
| # Start timing for this generation | |
| generation_start_time = time.time() | |
| if args.instrument: | |
| _instr_reset() | |
| if unguided_generation_mode: | |
| # Objective-free mode: pure model rollout without pCoMol guidance. | |
| time_grid = torch.linspace(0.0, 1.0, steps=args.num_steps, device=device) | |
| x_T = short_rollout_batch( | |
| model, | |
| x0, | |
| time_grid, | |
| start_idx=0, | |
| pad_id=pad_id, | |
| bos_id=bos_id, | |
| eos_id=eos_id, | |
| allowed_tokens=allowed_tokens, | |
| max_len_cap=args.max_len_cap, | |
| num_rollouts=1, | |
| num_steps=args.num_steps, | |
| protected_mask=None, | |
| deletion_rate_scale=args.deletion_rate_scale, | |
| zero_lam_ins=args.zero_lam_ins, | |
| ) | |
| else: | |
| x_T = pCoMol( | |
| model=model, | |
| x0=x0, | |
| pad_id=pad_id, | |
| bos_id=bos_id, | |
| eos_id=eos_id, | |
| allowed_tokens=allowed_tokens, | |
| objective_models=objective_models, | |
| constraint_models=constraint_models, | |
| w=objective_weights, | |
| rho=0.5, | |
| ref_z=ref_z, | |
| beta_start=args.beta_start, | |
| beta_end=args.beta_end, | |
| num_steps=args.num_steps, | |
| num_candidates=args.num_candidates, | |
| num_rollouts=args.num_rollouts, | |
| max_len_cap=args.max_len_cap, | |
| num_final_rollouts=args.num_final_rollouts, | |
| cfg=cfg, | |
| tokenizer=tokenizer, | |
| pam_masker=pam_masker, | |
| pam_mask_refresh_every=max(1, args.pam_refresh_every), | |
| pam_debug=args.pam_debug, | |
| pam_scale_edits=args.pam_scale_edits, | |
| pam_edit_scale_factor=args.pam_edit_scale_factor, | |
| deletion_rate_scale=args.deletion_rate_scale, | |
| zero_lam_ins=args.zero_lam_ins, | |
| legacy_beta_incumbent=args.legacy_beta_incumbent, | |
| ) | |
| # End timing for this generation | |
| generation_time = time.time() - generation_start_time | |
| out_str = tokenizer.batch_decode(x_T.tolist(), skip_special_tokens=True)[0] | |
| if args.instrument and INSTR is not None: | |
| _rec = {"generation_time": generation_time, "final_len": len(out_str.replace(' ', ''))} | |
| _rec.update({f"time_{k}": v for k, v in INSTR["timers"].items()}) | |
| _rec.update({f"n_{k}": v for k, v in INSTR["counts"].items()}) | |
| _rec["first_feasible_terminal_at"] = INSTR.get("first_feasible_terminal_at") | |
| _rec["legacy_beta_incumbent"] = bool(args.legacy_beta_incumbent) | |
| instr_records.append(_rec) | |
| print("----------------------------") | |
| print(f"\nDesigned Sequence: {out_str}\n") | |
| print("Final scores:") | |
| scores = compute_scores_print( | |
| [out_str], | |
| objective_models, constraint_models, | |
| device, return_scores=True | |
| ).squeeze(0) | |
| # Only count + save on success | |
| valid += 1 | |
| orig_len = len(args.input.replace(' ', '')) | |
| final_len = len(out_str.replace(' ', '')) | |
| length_diff = orig_len - final_len | |
| if args.output_file is not None: | |
| import csv | |
| objective_scores_only = scores[:len(objective_models)] | |
| row = _build_csv_row( | |
| sequence=out_str, | |
| final_len=final_len, | |
| length_diff=length_diff, | |
| objective_scores_only=objective_scores_only, | |
| generation_time=generation_time, | |
| ) | |
| generated_csv_rows.append(row) | |
| with open(args.output_file, 'a', newline='') as f: | |
| writer = csv.DictWriter(f, fieldnames=csv_header_parts) | |
| writer.writerow(row) | |
| except torch.cuda.OutOfMemoryError: | |
| # OOM error occurred - discard this generation and restart a new one | |
| print(f"[WARN] CUDA OOM during pCoMol generation (attempt {attempt}/{max_attempts}). Discarding this generation and restarting.") | |
| print(f" Current progress: {valid}/{target_valid} valid sequences generated") | |
| # Clear CUDA cache to free memory | |
| torch.cuda.empty_cache() | |
| torch.cuda.ipc_collect() # optional, can help in some cases | |
| # Continue to next iteration - this will start a new generation attempt | |
| # Note: attempt counter already incremented, so this doesn't count as a valid sequence | |
| continue | |
| if args.instrument and instr_records: | |
| import json | |
| _iout = args.instr_out or ((args.output_file.rsplit('.', 1)[0] + '_instr.json') if args.output_file else 'pcomol_instr.json') | |
| with open(_iout, 'w') as _f: | |
| json.dump(instr_records, _f, indent=2) | |
| print(f"[INSTRUMENT] wrote {len(instr_records)} per-sequence records to {_iout}") | |
| usable_csv_rows = [row for row in generated_csv_rows if int(row.get("usable_pass", 0)) == 1] | |
| total_generated = len(generated_csv_rows) | |
| total_usable = len(usable_csv_rows) | |
| final_wall_elapsed_s = float(time.time() - run_start_time) | |
| avg_usable_generation_time = final_wall_elapsed_s / total_usable if total_usable else "" | |
| print( | |
| "[TIME_BUDGET_SUMMARY] " | |
| f"usable_sequences={total_usable}, total_sequences={total_generated}, " | |
| f"avg_usable_generation_time={avg_usable_generation_time if avg_usable_generation_time != '' else 'nan'}, " | |
| f"wall_elapsed_s={final_wall_elapsed_s:.2f}" | |
| ) | |
| # Append an average row over generated outputs only (excludes the initial input row). | |
| if args.output_file is not None and generated_csv_rows: | |
| import csv | |
| avg_row = {k: "" for k in csv_header_parts} | |
| for col in csv_header_parts: | |
| vals = [] | |
| for row in generated_csv_rows: | |
| v = row.get(col, "") | |
| try: | |
| vals.append(float(v)) | |
| except (TypeError, ValueError): | |
| continue | |
| if vals: | |
| avg_row[col] = sum(vals) / len(vals) | |
| avg_row["sequence"] = "AVERAGE_EXCL_INPUT" | |
| if "predicted_pam" in avg_row: | |
| avg_row["predicted_pam"] = "" | |
| if "pam_target_match_binary" in avg_row and avg_row["pam_target_match_binary"] != "": | |
| avg_row["pam_target_match_binary"] = avg_row["pam_target_match_binary"] * 100.0 | |
| avg_row["total_sequences_generated"] = total_generated | |
| avg_row["usable_sequences_generated"] = total_usable | |
| avg_row["avg_usable_generation_time"] = avg_usable_generation_time | |
| with open(args.output_file, 'a', newline='') as f: | |
| writer = csv.DictWriter(f, fieldnames=csv_header_parts) | |
| writer.writerow(avg_row) | |
| print(f"[CSV] Appended average row over {len(generated_csv_rows)} generated sequences (excluding input row).") | |
| summary_row = {k: "" for k in csv_header_parts} | |
| summary_row["sequence"] = "SUMMARY_TIME_BUDGET" | |
| summary_row["total_sequences_generated"] = total_generated | |
| summary_row["usable_sequences_generated"] = total_usable | |
| summary_row["avg_usable_generation_time"] = avg_usable_generation_time | |
| summary_row["wall_elapsed_s"] = final_wall_elapsed_s | |
| with open(args.output_file, 'a', newline='') as f: | |
| writer = csv.DictWriter(f, fieldnames=csv_header_parts) | |
| writer.writerow(summary_row) | |
| print("[CSV] Appended time-budget summary row.") | |
| # for _ in range(100): | |
| # x_T = pCoMol( | |
| # model=model, | |
| # x0=x0, | |
| # pad_id=pad_id, | |
| # bos_id=bos_id, | |
| # eos_id=eos_id, | |
| # allowed_tokens=allowed_tokens, | |
| # objective_models=objective_models, | |
| # constraint_models=constraint_models, | |
| # w=objective_weights, | |
| # rho=0.5, | |
| # ref_z=ref_z, | |
| # beta_start=args.beta_start, | |
| # beta_end=args.beta_end, | |
| # num_steps=args.num_steps, | |
| # num_candidates=args.num_candidates, | |
| # num_rollouts=args.num_rollouts, | |
| # max_len_cap=args.max_len_cap, | |
| # num_final_rollouts=args.num_final_rollouts, | |
| # cfg=cfg, | |
| # selfies_tokenizer=selfies_tokenizer, | |
| # smiles_tokenizer=smiles_tokenizer | |
| # ) | |
| # out_str = selfies_tokenizer.batch_decode(x_T.tolist())[0] | |
| # smiles_token = smiles_tokenizer(out_str, return_tensors='pt')['input_ids'].to(device) | |
| # print("----------------------------") | |
| # # print(f"Initial Sequence: {args.input}\n") | |
| # # print(f"Initial Scores:") | |
| # # compute_scores_print([args.input], objective_models, constraint_models, device) | |
| # print(f"\nDesigned Sequence: {out_str}\n") | |
| # print("Final scores:") | |
| # scores = compute_scores_print(smiles_token, [out_str], objective_models, constraint_models, device, return_scores=True).squeeze(0) | |
| # with open(args.output_file, 'a') as f: | |
| # f.write(f"{smiles_token.shape[1]}") | |
| # for score in scores: | |
| # f.write(f",{score.item()}") | |
| # f.write(f",{out_str}\n") | |
| if __name__ == "__main__": | |
| main() | |