# 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 # --------------------------------------------------------------------------- @torch.no_grad() 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 # --------------------------------------------------------------------------- @torch.no_grad() 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: " "_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: _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()