pCoMole / cas9 /pcomole.py
Maximilian Holsman
Claude Opus 5
Add Cas9 task
12fea4a
Raw History Blame Contribute Delete
140 kB
# NOTE: imported first, before the flow_matching package lands on sys.path during the imports
# below (its flow_matching/utils subpackage otherwise shadows this local editflows/utils package).
from cas9.pam_detector_hmm import HMMSCAN, Cas9PIMasker # noqa: E402
import argparse
from typing import List, Callable, Optional, Tuple, Dict, Any
import torch
import torch.nn.functional as F
import yaml
from easydict import EasyDict as edict
from tqdm import tqdm
import time
import math
from cas9.generate import build_model_and_stuff, tokenize_input_str, detokenize_output
from cas9.objectives import Cas9Classification, DeletionCount, PAMMatching
from cas9.constraints import Cas9DomainCompleteness, ProteinLength, TargetLength, MinTargetLength, MaxTargetLength, PAMMatchingConstraint, Cas9ScoreThreshold, PAMMatchingProbabilityThreshold
# PAM/PI-domain detector + mask builder (moved to top of file to avoid flow_matching/utils shadowing)
# from utils.pam_domain_detector.cas9_pam_detector_hmm import HMMSCAN, Cas9PIMasker
# ---- Opt-in instrumentation (runtime breakdown + rollout-failure counts) ----
# Inert unless main() enables it via --instrument (INSTR stays None otherwise), so default
# generation behavior is unchanged. Timers use cuda.synchronize for GPU-accurate wall time.
from collections import defaultdict
INSTR = None
def _instr_reset():
global INSTR
INSTR = {"timers": defaultdict(float), "counts": defaultdict(int)}
def _tic():
if INSTR is None:
return None
if torch.cuda.is_available():
torch.cuda.synchronize()
return time.time()
def _toc(name, t0):
if t0 is None:
return
if torch.cuda.is_available():
torch.cuda.synchronize()
INSTR["timers"][name] += time.time() - t0
def _count(name, k=1):
if INSTR is not None:
INSTR["counts"][name] += int(k)
import pdb
import warnings
warnings.filterwarnings("ignore", category=FutureWarning)
from transformers.utils import logging
logging.set_verbosity_error()
# ---------------------------------------------------------------------------
# small utilities
# ---------------------------------------------------------------------------
class PAMDomainWrapper:
"""
Wrapper for PAMMatching that extracts PAM domain region before PAM prediction.
When PID_PAM_prediction is enabled, this wrapper extracts only the detected
PAM domain region from sequences before passing them to the underlying PAMMatching object.
"""
def __init__(self, pam_matching_obj, pam_domain_interval):
"""
Args:
pam_matching_obj: PAMMatching object to wrap
pam_domain_interval: Tuple (start, end) of 1-based coordinates for PAM domain region
"""
self._pam_matching_obj = pam_matching_obj
self._pam_domain_interval = pam_domain_interval
start_1based, end_1based = pam_domain_interval
self._start_0based = start_1based - 1 # Convert to 0-based for slicing
self._end_0based = end_1based
def _extract_domain_region(self, protein_seqs):
"""Extract PAM domain region from each sequence."""
domain_seqs = []
for seq in protein_seqs:
seq_clean = seq.replace(" ", "").strip()
# Extract domain region (handle cases where sequence might be shorter)
if len(seq_clean) < self._end_0based:
# If sequence is shorter than expected domain end, use what's available
# This can happen during generation when sequences are being edited
domain_seq = seq_clean[self._start_0based:] if len(seq_clean) > self._start_0based else seq_clean
else:
domain_seq = seq_clean[self._start_0based:self._end_0based]
domain_seqs.append(domain_seq)
return domain_seqs
def predict_pam(self, protein_seqs, min_confidence=None):
"""Extract domain region and predict PAM."""
domain_seqs = self._extract_domain_region(protein_seqs)
return self._pam_matching_obj.predict_pam(domain_seqs, min_confidence=min_confidence)
def get_score_for_pam(self, protein_seqs, pam_sequence, use_temperature_scaling=False):
"""Extract domain region and get score for PAM."""
domain_seqs = self._extract_domain_region(protein_seqs)
return self._pam_matching_obj.get_score_for_pam(domain_seqs, pam_sequence, use_temperature_scaling=use_temperature_scaling)
def __call__(self, protein_tokens, protein_seqs):
"""Extract domain region and compute objective score."""
domain_seqs = self._extract_domain_region(protein_seqs)
return self._pam_matching_obj(protein_tokens, domain_seqs)
def __getattr__(self, name):
"""Delegate all other attributes to the wrapped object."""
return getattr(self._pam_matching_obj, name)
def extract_objective_vector(protein_seqs, objective_models, device, return_names=False):
"""
Extract objective vector from protein sequences.
Args:
protein_seqs: List of protein sequence strings
objective_models: List of objective model callables
device: torch device
return_names: If True, also return list of objective names
Returns:
torch.Tensor of shape (B, m) where m is number of objectives
If return_names=True, also returns list of objective names
"""
values = []
names = []
# Allow objective-free runs (base EditFlows sampling mode).
if len(objective_models) == 0:
empty = torch.empty((len(protein_seqs), 0), device=device, dtype=torch.float32)
if return_names:
return empty, names
return empty
for obj in objective_models:
name, scores = obj(protein_tokens=None, protein_seqs=protein_seqs) # list of length B
values.append(torch.tensor(scores, device=device, dtype=torch.float32))
if return_names:
names.append(name)
result = torch.stack(values, dim=1) # (B,m)
if return_names:
return result, names
return result
def compute_scores_print(protein_seqs, objective_models, constraint_models, device, return_scores=False):
"""
Compute and print scores for protein sequences.
Args:
protein_seqs: List of protein sequence strings
objective_models: List of objective model callables
constraint_models: List of constraint model callables
device: torch device
return_scores: Whether to return scores
Returns:
scores if return_scores=True, else None
"""
objective_scores = extract_objective_vector(protein_seqs, objective_models, device) # (B,m)
constraint_scores = []
for constraint in constraint_models:
# Handle constraints that take single sequences vs batches
if hasattr(constraint, 'predict_batch'):
# Use batch prediction if available (more efficient)
constraint_score = constraint.predict_batch(protein_seqs)
constraint_score = [int(c) for c in constraint_score] # Convert bool to int
else:
# Call for each sequence individually
constraint_score = [constraint(None, seq) for seq in protein_seqs]
constraint_scores.append(torch.tensor(constraint_score, device=device, dtype=torch.float32))
constraint_scores = torch.stack(constraint_scores, dim=1) # (B, n)
scores = torch.concat([objective_scores, constraint_scores], dim=1) # (B, m+n)
print(scores)
# Print predicted PAM and PAM probability score if PAM matching objective is present
for obj in objective_models:
if hasattr(obj, 'predict_pam'):
predicted_pams = obj.predict_pam(protein_seqs)
# Get PAM probability scores for each predicted PAM (not the target PAM)
# Use raw probabilities (no temperature scaling) for more interpretable results
pam_prob_scores = []
for i, pam in enumerate(predicted_pams):
# Compute score for this specific predicted PAM using raw probabilities
score = obj.get_score_for_pam([protein_seqs[i]], pam, use_temperature_scaling=False)[0]
pam_prob_scores.append(score)
for i, (pam, prob_score) in enumerate(zip(predicted_pams, pam_prob_scores)):
print(f"Predicted PAM: {pam} (target: {obj.target_pam}) | Predicted PAM probability (raw): {prob_score:.8f}")
break # Only print once if multiple PAM objectives exist
if return_scores:
return scores
# ---------------------------------------------------------------------------
# edit utilities
# ---------------------------------------------------------------------------
@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: "
"<output_file>_pool.csv).")
parser.add_argument("--instrument", action="store_true",
help="Record per-sequence runtime breakdown by component + rollout-failure counts.")
parser.add_argument("--instr_out", type=str, default=None,
help="Where to write the instrumentation JSON (default: <output_file>_instr.json).")
parser.add_argument("--seed", type=int, default=None,
help="RNG seed for reproducibility / seed-level robustness studies. Sets python/numpy/torch(+cuda) seeds. If unset, generation is unseeded as before.")
args = parser.parse_args()
# Validate that pam_hmm_db is provided if pam_mask or pam_scale_edits is enabled
if args.pam_mask and args.pam_hmm_db is None:
raise ValueError("--pam_mask requires --pam_hmm_db to be specified")
if args.pam_scale_edits and args.pam_hmm_db is None:
raise ValueError("--pam_scale_edits requires --pam_hmm_db to be specified")
if args.pam_mask and args.pam_scale_edits:
raise ValueError("--pam_mask and --pam_scale_edits cannot be used together. Choose one.")
if args.PID_PAM_prediction and args.pam_hmm_db is None:
raise ValueError("--PID_PAM_prediction requires --pam_hmm_db to be specified")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if args.seed is not None:
import random as _random, numpy as _np
_random.seed(args.seed)
_np.random.seed(args.seed)
torch.manual_seed(args.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(args.seed)
print(f"[SEED] Set python/numpy/torch RNG seeds to {args.seed}")
with open(args.pcomol_config, "r") as f:
cfg = edict(yaml.safe_load(f))
editflow, source_dist, tokenizer, pad_id, bos_id, eos_id, eps_id = build_model_and_stuff(cfg, device)
ckpt = torch.load(args.ckpt, map_location=device)
editflow.load_state_dict(ckpt["state_dict"], strict=False)
model = editflow.model.to(device)
model.eval()
x0 = tokenize_input_str(args.input, cfg, tokenizer, bos_id, eos_id, pad_id, device)
allowed_tokens = torch.tensor(
[tok for tok in source_dist._allowed_tokens if tok not in (eps_id,)],
device=device,
dtype=torch.long,
)
# Initialize PAM domain detector early if PID_PAM_prediction is enabled
pam_domain_interval = None # (start, end) 1-based coordinates
if args.PID_PAM_prediction:
print(f"[PID_PAM_prediction] Detecting PAM domain from input sequence...")
hmmscan_runner = HMMSCAN(
hmm_db_path=args.pam_hmm_db,
hmmscan_bin=args.hmmscan_bin,
cpu=args.hmmscan_cpu,
)
pam_detector = Cas9PIMasker(
hmmscan=hmmscan_runner,
use_env_coords=False,
max_mask_len=args.pam_mask_max_len,
min_mask_len=args.pam_mask_min_len,
evalue_cutoff=args.pam_evalue,
min_ali_len=30,
fallback_last_n=None,
cache_size=2048,
)
input_seq_clean = args.input.replace(" ", "").strip()
pam_domain_interval = pam_detector.pi_core_interval(input_seq_clean)
if pam_domain_interval is None:
raise ValueError(f"[PID_PAM_prediction] Failed to detect PAM domain in input sequence. Cannot proceed.")
start_1based, end_1based = pam_domain_interval
# Extract PAM domain region (convert 1-based to 0-based for slicing)
pam_domain_seq = input_seq_clean[start_1based - 1:end_1based]
print(f"[PID_PAM_prediction] Detected PAM domain: positions {start_1based}-{end_1based} (1-based), length={len(pam_domain_seq)}")
print(f"[PID_PAM_prediction] PAM domain sequence: {pam_domain_seq[:50]}...{pam_domain_seq[-50:] if len(pam_domain_seq) > 100 else pam_domain_seq}")
# Override model name for PAM prediction
args.pam_model_name = "Profluent-Bio/protein2pam-cas9"
print(f"[PID_PAM_prediction] Using model: {args.pam_model_name} (changed from default)")
cas9_classifier = None
if args.cas9_objective or args.cas9_score_threshold is not None:
# Initialize Cas9 classifier objective with shared ESM model
# This avoids loading ESM twice, saving GPU memory
cas9_classifier = Cas9Classification(
device=device,
checkpoint_path=args.cas9_classifier_ckpt,
config_path=args.cas9_classifier_config,
shared_esm_model=model.esm_emb # Share ESM model from editflow
)
objective_models = []
if args.cas9_objective:
objective_models.append(cas9_classifier)
if args.deletion_count_obj:
# Initialize deletion count objective (maximizes number of deletions, normalized to [0, 1])
deletion_count_obj = DeletionCount(
original_seq=args.input,
max_deletion_percentage=args.max_deletion_percentage,
)
objective_models.append(deletion_count_obj)
# Initialize PAM matching objective if target PAM is provided
pam_matching_obj_for_constraint = None # Will be set if PAM matching is enabled
if args.target_pam is not None:
# Handle 'matching' mode: predict PAM from input sequence
if args.target_pam.lower() == 'matching':
print(f"Target PAM set to 'matching' - will predict PAM from input sequence")
# Create temporary PAM matching object to predict the input sequence's PAM
# Use PAM domain region if PID_PAM_prediction is enabled
input_seq_for_pam_prediction = args.input
if args.PID_PAM_prediction and pam_domain_interval is not None:
input_seq_clean = args.input.replace(" ", "").strip()
start_1based, end_1based = pam_domain_interval
pam_domain_seq = input_seq_clean[start_1based - 1:end_1based]
input_seq_for_pam_prediction = pam_domain_seq
print(f" Using PAM domain region (positions {start_1based}-{end_1based}) for initial PAM prediction")
temp_pam_obj = PAMMatching(
device=device,
target_pam="NNNNNNNNNN", # Dummy target, we'll replace it
model_name=args.pam_model_name,
use_entropy_for_n_positions=args.pam_use_entropy,
use_ce_loss=args.use_ce_loss,
pam_min_confidence=args.pam_min_confidence,
pam_prediction_temperature=args.pam_prediction_temperature
)
# Predict PAM for input sequence (or PAM domain region)
predicted_pams = temp_pam_obj.predict_pam([input_seq_for_pam_prediction])
if predicted_pams:
actual_target_pam = predicted_pams[0]
print(f" Predicted PAM for {'PAM domain region' if args.PID_PAM_prediction else 'input sequence'}: {actual_target_pam}")
print(f" Using this as target PAM for optimization")
else:
raise ValueError("Failed to predict PAM for input sequence")
# Now create the actual PAM matching objective with the predicted PAM
entropy_mode = "enabled" if args.pam_use_entropy else "disabled"
ce_loss_mode = "CE loss" if args.use_ce_loss else "log probability"
print(f" Entropy scoring for N positions: {entropy_mode}")
print(f" Scoring method: {ce_loss_mode}")
pam_matching_obj_base = PAMMatching(
device=device,
target_pam=actual_target_pam,
model_name=args.pam_model_name,
use_entropy_for_n_positions=args.pam_use_entropy,
use_ce_loss=args.use_ce_loss,
pam_min_confidence=args.pam_min_confidence,
pam_prediction_temperature=args.pam_prediction_temperature
)
# Wrap with PAM domain extractor if PID_PAM_prediction is enabled
if args.PID_PAM_prediction and pam_domain_interval is not None:
pam_matching_obj = PAMDomainWrapper(pam_matching_obj_base, pam_domain_interval)
print(f" Wrapped PAM matching objective with PAM domain extractor (interval: {pam_domain_interval})")
else:
pam_matching_obj = pam_matching_obj_base
if args.pam_distr_objective:
objective_models.append(pam_matching_obj)
print(f"PAM matching objective initialized successfully with target PAM: {actual_target_pam}")
# Store reference to PAM matching object for constraint creation later
pam_matching_obj_for_constraint = pam_matching_obj
else:
# Regular mode: use provided PAM string
print(f"Initializing PAM matching objective with target PAM: {args.target_pam}")
entropy_mode = "enabled" if args.pam_use_entropy else "disabled"
ce_loss_mode = "CE loss" if args.use_ce_loss else "log probability"
print(f" Entropy scoring for N positions: {entropy_mode}")
print(f" Scoring method: {ce_loss_mode}")
pam_matching_obj_base = PAMMatching(
device=device,
target_pam=args.target_pam,
model_name=args.pam_model_name,
use_entropy_for_n_positions=args.pam_use_entropy,
use_ce_loss=args.use_ce_loss,
pam_min_confidence=args.pam_min_confidence,
pam_prediction_temperature=args.pam_prediction_temperature
)
# Wrap with PAM domain extractor if PID_PAM_prediction is enabled
if args.PID_PAM_prediction and pam_domain_interval is not None:
pam_matching_obj = PAMDomainWrapper(pam_matching_obj_base, pam_domain_interval)
print(f" Wrapped PAM matching objective with PAM domain extractor (interval: {pam_domain_interval})")
else:
pam_matching_obj = pam_matching_obj_base
if args.pam_distr_objective:
objective_models.append(pam_matching_obj)
print(f"PAM matching objective initialized successfully.")
# Store reference to PAM matching object for constraint creation later
pam_matching_obj_for_constraint = pam_matching_obj
num_objectives = len(objective_models)
no_objective_mode = (num_objectives == 0)
if no_objective_mode:
print(
"[INFO] No objectives selected; running in base EditFlows sampling mode "
"(unguided rollouts)."
)
unguided_generation_mode = args.disable_guidance or no_objective_mode
if args.disable_guidance:
print(
"[INFO] --disable_guidance enabled; generation uses unguided EditFlows "
"rollouts while objective/constraint scores are still reported."
)
if not args.objective_weights:
if num_objectives > 0:
objective_weights = torch.tensor([1.0 / num_objectives] * num_objectives).to(device)
else:
objective_weights = torch.empty((0,), device=device, dtype=torch.float32)
else:
if num_objectives == 0:
raise ValueError(
"objective_weights provided but no objectives are enabled. "
"Either enable objective flags or remove --objective_weights."
)
if len(args.objective_weights) != num_objectives:
raise ValueError(
f"objective_weights length ({len(args.objective_weights)}) does not match "
f"number of objectives ({num_objectives})."
)
objective_weights = torch.tensor(args.objective_weights).to(device)
if not args.ref_z:
ref_z = torch.zeros(num_objectives).to(device)
else:
if num_objectives == 0:
raise ValueError(
"ref_z provided but no objectives are enabled. "
"Either enable objective flags or remove --ref_z."
)
ref_z = torch.tensor(args.ref_z).to(device)
# Initialize protein length constraint
protein_length_constraint = ProteinLength(
orig_protein_seq=args.input,
cfg=cfg,
tokenizer=tokenizer,
allow_equal_length=args.allow_equal_length,
)
constraint_models = [protein_length_constraint]
# Initialize Cas9 domain completeness constraint (optional)
if args.cas9_hmm_db is not None:
cas9_domain_constraint = Cas9DomainCompleteness(
hmm_path=args.cas9_hmm_db,
ev1=args.cas9_evalue,
minlen1=args.cas9_minlen,
cpu=1, # Can be made configurable if needed
)
constraint_models.append(cas9_domain_constraint)
print(f"Cas9 domain completeness constraint enabled with HMM database: {args.cas9_hmm_db}")
else:
print("Cas9 domain completeness constraint disabled (--cas9_hmm_db not provided)")
# Initialize target length constraint if specified
if args.target_length is not None:
target_length_constraint = TargetLength(target_length=args.target_length)
constraint_models.append(target_length_constraint)
print(f"Target length constraint enabled: sequences must have exactly {args.target_length} amino acids")
# Initialize min target length constraint if specified
computed_min_target_length = None
if args.min_target_length is not None:
# If min_target_length is a decimal (0 < value < 1), interpret as percentage of input length
# If it's >= 1, use as absolute value
input_length = len(args.input.replace(" ", "").replace("\n", ""))
if 0 < args.min_target_length < 1:
# Decimal: interpret as percentage
computed_min_target_length = int(args.min_target_length * input_length)
print(f"Min target length constraint: {args.min_target_length} (decimal) interpreted as {args.min_target_length*100:.1f}% of input length ({input_length}) = {computed_min_target_length}")
else:
# Integer: use as absolute value
computed_min_target_length = int(args.min_target_length)
print(f"Min target length constraint: {args.min_target_length} (integer) used as absolute value = {computed_min_target_length}")
min_target_length_constraint = MinTargetLength(min_target_length=computed_min_target_length)
constraint_models.append(min_target_length_constraint)
print(f"Min target length constraint enabled: sequences must have length >= {computed_min_target_length} amino acids")
# Initialize max target length constraint if specified
computed_max_target_length = None
if args.max_target_length is not None:
# If max_target_length is a decimal (0 < value < 1), interpret as percentage of input length
# If it's >= 1, use as absolute value
input_length = len(args.input.replace(" ", "").replace("\n", ""))
if 0 < args.max_target_length < 1:
# Decimal: interpret as percentage
computed_max_target_length = int(args.max_target_length * input_length)
print(f"Max target length constraint: {args.max_target_length} (decimal) interpreted as {args.max_target_length*100:.1f}% of input length ({input_length}) = {computed_max_target_length}")
else:
# Integer: use as absolute value
computed_max_target_length = int(args.max_target_length)
print(f"Max target length constraint: {args.max_target_length} (integer) used as absolute value = {computed_max_target_length}")
max_target_length_constraint = MaxTargetLength(max_target_length=computed_max_target_length)
constraint_models.append(max_target_length_constraint)
print(f"Max target length constraint enabled: sequences must have length <= {computed_max_target_length} amino acids")
# Validate explicit length range if both bounds are provided
if computed_min_target_length is not None and computed_max_target_length is not None:
if computed_min_target_length > computed_max_target_length:
raise ValueError(
f"Incompatible length range: min_target_length ({computed_min_target_length}) "
f"must be <= max_target_length ({computed_max_target_length})."
)
# Initialize PAM matching constraint if target_pam is specified
if args.target_pam is not None and pam_matching_obj_for_constraint is not None:
pam_matching_constraint = PAMMatchingConstraint(pam_matching_obj_for_constraint)
constraint_models.append(pam_matching_constraint)
print(f"PAM matching constraint enabled: sequences must have predicted PAM matching target PAM")
# Initialize PAM matching probability threshold constraint if threshold is specified
if args.pam_probability_threshold is not None:
if pam_matching_obj_for_constraint is None:
raise ValueError("--pam_probability_threshold requires --target_pam to enable PAM matching evaluation.")
pam_probability_threshold_constraint = PAMMatchingProbabilityThreshold(
pam_matching_obj_for_constraint, args.pam_probability_threshold
)
constraint_models.append(pam_probability_threshold_constraint)
print(
"PAM matching probability threshold constraint enabled: sequences must have "
f"PAM matching probability > {args.pam_probability_threshold}"
)
# Initialize Cas9 score threshold constraint if threshold is specified
if args.cas9_score_threshold is not None:
if cas9_classifier is None:
raise ValueError("--cas9_score_threshold requires --cas9_objective or a valid Cas9 classifier.")
cas9_score_threshold_constraint = Cas9ScoreThreshold(cas9_classifier, args.cas9_score_threshold)
constraint_models.append(cas9_score_threshold_constraint)
print(f"Cas9 score threshold constraint enabled: sequences must have Cas9 score > {args.cas9_score_threshold}")
# Initialize PAM/PI masker if --pam_mask flag is set
pam_masker = None
if args.pam_mask:
if args.pam_hmm_db is None:
raise ValueError("--pam_mask requires --pam_hmm_db to be specified")
hmmscan_runner = HMMSCAN(
hmm_db_path=args.pam_hmm_db,
hmmscan_bin=args.hmmscan_bin,
cpu=args.hmmscan_cpu,
)
pam_masker = Cas9PIMasker(
hmmscan=hmmscan_runner,
use_env_coords=False, # ali coords (smaller)
max_mask_len=args.pam_mask_max_len,
min_mask_len=args.pam_mask_min_len,
evalue_cutoff=args.pam_evalue,
min_ali_len=30,
fallback_last_n=None, # avoid masking arbitrary tail when no hit
cache_size=2048,
)
print(f"[PI mask] enabled: db={args.pam_hmm_db}, max_len={args.pam_mask_max_len}, min_len={args.pam_mask_min_len}, evalue<={args.pam_evalue} (blocks all edits: ins/del/sub)")
elif args.pam_scale_edits:
if args.pam_hmm_db is None:
raise ValueError("--pam_scale_edits requires --pam_hmm_db to be specified")
# For scaling mode, we still need the masker to identify the PAM region
hmmscan_runner = HMMSCAN(
hmm_db_path=args.pam_hmm_db,
hmmscan_bin=args.hmmscan_bin,
cpu=args.hmmscan_cpu,
)
pam_masker = Cas9PIMasker(
hmmscan=hmmscan_runner,
use_env_coords=False, # ali coords (smaller)
max_mask_len=args.pam_mask_max_len,
min_mask_len=args.pam_mask_min_len,
evalue_cutoff=args.pam_evalue,
min_ali_len=30,
fallback_last_n=None, # avoid masking arbitrary tail when no hit
cache_size=2048,
)
print(f"[PI scaling] enabled: db={args.pam_hmm_db}, max_len={args.pam_mask_max_len}, min_len={args.pam_mask_min_len}, evalue<={args.pam_evalue} (scales ins/sub rates by {args.pam_edit_scale_factor}x)")
print("Initial Scores:")
input_scores = compute_scores_print([args.input], objective_models, constraint_models, device, return_scores=True).squeeze(0)
# -----------------------------------------------------------------------
# Budget-matched baseline: unguided Edit Flow + terminal rejection.
# No candidate proposal, no rollouts, no Doob-h tilt -- just i.i.d. terminals
# from the base kernel started at the input, each scored once by the same
# oracle, keeping the best feasible one. Because x0 is fixed and the base
# kernel carries no run-level state, the drawn terminals are i.i.d., so this
# pool characterises the sampler at every budget <= --oracle_budget.
# -----------------------------------------------------------------------
if args.rejection_baseline:
import os, csv, json as _json
_instr_reset()
rho_pool = 0.5 # same value pCoMol() is invoked with below
obj_names_pool = []
for obj in objective_models:
_nm, _ = obj(protein_tokens=None, protein_seqs=[args.input])
obj_names_pool.append(_nm)
cons_names = [c.__class__.__name__ for c in constraint_models]
orig_len = len(args.input.replace(' ', ''))
pool_path = args.pool_out or (
(args.output_file.rsplit('.', 1)[0] + '_pool.csv') if args.output_file else 'rejection_pool.csv'
)
os.makedirs(os.path.dirname(os.path.abspath(pool_path)) or '.', exist_ok=True)
header = (["idx", "final_len", "length_diff"]
+ [f"{n}_score" for n in obj_names_pool]
+ ["utility", "feasible"]
+ [f"pass_{n}" for n in cons_names]
+ ["sequence"])
with open(pool_path, 'w', newline='') as _f:
csv.writer(_f).writerow(header)
print(f"[REJECTION] unguided Edit Flow + terminal rejection | budget={args.oracle_budget} "
f"terminals | batch={args.rejection_batch} | steps={args.num_steps}")
print(f"[REJECTION] objectives={obj_names_pool} constraints={cons_names}")
time_grid = torch.linspace(0.0, 1.0, steps=args.num_steps, device=device)
n_done = 0
n_feasible = 0
first_feasible_idx = None
best_u = float('-inf')
best_seq_str = None
t_sample_tot = 0.0
t_oracle_tot = 0.0
pool_t0 = time.time()
while n_done < args.oracle_budget:
k = min(args.rejection_batch, args.oracle_budget - n_done)
_t0 = time.time()
x_T = short_rollout_batch(
model, x0, time_grid, 0, pad_id, bos_id, eos_id, allowed_tokens,
args.max_len_cap, num_rollouts=k, num_steps=args.num_steps,
protected_mask=None, deletion_rate_scale=args.deletion_rate_scale,
zero_lam_ins=args.zero_lam_ins,
)
if torch.cuda.is_available():
torch.cuda.synchronize()
t_sample_tot += time.time() - _t0
_t1 = time.time()
_, _, det = _G_T(
x_T, objective_models, constraint_models, objective_weights, rho_pool, ref_z,
args.beta_end, tokenizer, ws_for_invalid=True,
debug_context="REJECTION_TERMINAL", count_terminal=True, return_details=True,
)
if torch.cuda.is_available():
torch.cuda.synchronize()
t_oracle_tot += time.time() - _t1
f_vals = det["f_vals"]
u_all = _augmented_tchebycheff(f_vals, objective_weights, rho_pool, ref_z)
cons = det["constraint_results"]
feas = (cons == 1).all(dim=0)
seqs = det["protein_seqs"]
rows = []
for i in range(len(seqs)):
gidx = n_done + i
s = seqs[i]
is_feas = bool(feas[i].item())
u_i = float(u_all[i].item())
if is_feas:
n_feasible += 1
if first_feasible_idx is None:
first_feasible_idx = gidx + 1 # 1-based oracle-call index, exact
if u_i > best_u:
best_u = u_i
best_seq_str = s
rows.append([gidx, len(s), orig_len - len(s)]
+ [f"{float(f_vals[i, j].item()):.6f}" for j in range(f_vals.shape[1])]
+ [f"{u_i:.6f}", int(is_feas)]
+ [int(cons[j, i].item()) for j in range(cons.shape[0])]
+ [s])
with open(pool_path, 'a', newline='') as _f:
csv.writer(_f).writerows(rows)
n_done += k
_best_disp = best_u if best_u > float('-inf') else float('nan')
print(f"[REJECTION] {n_done}/{args.oracle_budget} terminals | "
f"feasible={n_feasible} ({100.0 * n_feasible / n_done:.2f}%) | "
f"best_utility={_best_disp:.4f} | "
f"sample={t_sample_tot:.0f}s oracle={t_oracle_tot:.0f}s", flush=True)
summary = {
"mode": "rejection_baseline",
"input_length": orig_len,
"oracle_budget": int(args.oracle_budget),
"rejection_batch": int(args.rejection_batch),
"num_steps": int(args.num_steps),
"deletion_rate_scale": float(args.deletion_rate_scale),
"zero_lam_ins": bool(args.zero_lam_ins),
"seed": args.seed,
"n_terminals": int(n_done),
"n_feasible": int(n_feasible),
"first_feasible_oracle_call": first_feasible_idx,
"best_utility": (best_u if best_u > float('-inf') else None),
"best_sequence": best_seq_str,
"objective_names": obj_names_pool,
"constraint_names": cons_names,
"objective_weights": [float(v) for v in objective_weights.tolist()],
"ref_z": [float(v) for v in ref_z.tolist()],
"rho": rho_pool,
"time_sample_s": t_sample_tot,
"time_oracle_s": t_oracle_tot,
"wall_s": time.time() - pool_t0,
"instr_counts": dict(INSTR["counts"]),
"pool_csv": pool_path,
}
sum_path = pool_path.rsplit('.', 1)[0] + '_summary.json'
with open(sum_path, 'w') as _f:
_json.dump(summary, _f, indent=2)
print(f"[REJECTION] wrote pool -> {pool_path}")
print(f"[REJECTION] wrote summary -> {sum_path}")
return
# CSV bookkeeping for optional output file
csv_header_parts = None
obj_names = []
pam_obj = None
generated_csv_rows = [] # Generated rows only (excludes the initial input row)
run_start_time = time.time()
# Write CSV header and input sequence row if file doesn't exist
if args.output_file is not None:
import os
import csv
for obj in objective_models:
name, _ = obj(protein_tokens=None, protein_seqs=[args.input])
obj_names.append(name)
if hasattr(obj, 'predict_pam'):
pam_obj = obj
csv_header_parts = ["final_len", "length_diff"]
csv_header_parts.extend([f"{name}_score" for name in obj_names])
csv_header_parts.extend([f"{name}_weight" for name in obj_names])
if pam_obj is not None:
csv_header_parts.append("predicted_pam")
if hasattr(pam_obj, 'use_ce_loss') and pam_obj.use_ce_loss:
csv_header_parts.append("pam_matching_ce_loss_score")
csv_header_parts.append("pam_matching_probability")
csv_header_parts.append("pam_target_match_binary")
csv_header_parts.extend([
"length_pass",
"cas9_score_pass",
"pam_prob_pass",
"usable_pass",
"total_sequences_generated",
"usable_sequences_generated",
"avg_usable_generation_time",
"wall_elapsed_s",
])
csv_header_parts.append("deletion_rate_scale")
if pam_obj is not None:
csv_header_parts.append("pam_min_confidence")
csv_header_parts.append("generation_time")
csv_header_parts.append("sequence")
def _constraint_passes(sequence):
passes = []
for constraint in constraint_models:
if hasattr(constraint, 'predict_batch'):
result = constraint.predict_batch([sequence])[0]
else:
result = constraint(None, sequence)
passes.append(bool(result))
return passes
def _build_csv_row(sequence, final_len, length_diff, objective_scores_only, generation_time):
row = {k: "" for k in csv_header_parts}
row["final_len"] = final_len
row["length_diff"] = length_diff
row["deletion_rate_scale"] = float(args.deletion_rate_scale)
row["generation_time"] = float(generation_time)
row["wall_elapsed_s"] = float(time.time() - run_start_time)
row["sequence"] = sequence.replace(' ', '')
for name, score in zip(obj_names, objective_scores_only):
row[f"{name}_score"] = float(score.item() if torch.is_tensor(score) else score)
for name, weight in zip(obj_names, objective_weights):
row[f"{name}_weight"] = float(weight.item() if torch.is_tensor(weight) else weight)
if pam_obj is not None:
predicted_pams = pam_obj.predict_pam([sequence])
predicted_pam = predicted_pams[0] if predicted_pams else ""
row["predicted_pam"] = predicted_pam
if hasattr(pam_obj, 'use_ce_loss') and pam_obj.use_ce_loss:
pam_score_idx = None
for i, obj in enumerate(objective_models):
if obj == pam_obj:
pam_score_idx = i
break
if pam_score_idx is not None:
row["pam_matching_ce_loss_score"] = float(objective_scores_only[pam_score_idx].item())
target_pam = pam_obj.target_pam
prob_score = pam_obj.get_score_for_pam([sequence], target_pam, use_temperature_scaling=False)[0]
row["pam_matching_probability"] = float(prob_score)
row["pam_target_match_binary"] = 1.0 if predicted_pam == target_pam else 0.0
row["pam_min_confidence"] = float(args.pam_min_confidence)
constraint_passes = _constraint_passes(row["sequence"])
row["length_pass"] = int(all(constraint_passes[:1]))
if args.target_length is not None:
row["length_pass"] = int(final_len == int(args.target_length))
if computed_min_target_length is not None and final_len < computed_min_target_length:
row["length_pass"] = 0
if computed_max_target_length is not None and final_len > computed_max_target_length:
row["length_pass"] = 0
if args.cas9_score_threshold is not None and row.get("cas9_score", "") != "":
row["cas9_score_pass"] = int(float(row["cas9_score"]) > float(args.cas9_score_threshold))
if args.pam_probability_threshold is not None and row.get("pam_matching_probability", "") != "":
row["pam_prob_pass"] = int(float(row["pam_matching_probability"]) > float(args.pam_probability_threshold))
row["usable_pass"] = int(all(constraint_passes))
return row
if not os.path.exists(args.output_file):
with open(args.output_file, 'w', newline='') as f:
writer = csv.DictWriter(f, fieldnames=csv_header_parts)
writer.writeheader()
input_len = len(args.input.replace(' ', ''))
input_objective_scores_only = input_scores[:len(objective_models)]
input_row = _build_csv_row(
sequence=args.input,
final_len=input_len,
length_diff=0,
objective_scores_only=input_objective_scores_only,
generation_time=0.0,
)
writer.writerow(input_row)
valid = 0
target_valid = args.num_sequences # how many successful designs you want
attempt = 0
max_attempts = max(500, target_valid) # safety cap so you don't infinite-loop if it keeps OOM'ing
instr_records = []
while valid < target_valid and attempt < max_attempts:
if args.time_limit_s is not None and (time.time() - run_start_time) >= args.time_limit_s:
print(f"[TIME_LIMIT] Reached {args.time_limit_s}s generation budget.")
break
attempt += 1
try:
# Start timing for this generation
generation_start_time = time.time()
if args.instrument:
_instr_reset()
if unguided_generation_mode:
# Objective-free mode: pure model rollout without pCoMol guidance.
time_grid = torch.linspace(0.0, 1.0, steps=args.num_steps, device=device)
x_T = short_rollout_batch(
model,
x0,
time_grid,
start_idx=0,
pad_id=pad_id,
bos_id=bos_id,
eos_id=eos_id,
allowed_tokens=allowed_tokens,
max_len_cap=args.max_len_cap,
num_rollouts=1,
num_steps=args.num_steps,
protected_mask=None,
deletion_rate_scale=args.deletion_rate_scale,
zero_lam_ins=args.zero_lam_ins,
)
else:
x_T = pCoMol(
model=model,
x0=x0,
pad_id=pad_id,
bos_id=bos_id,
eos_id=eos_id,
allowed_tokens=allowed_tokens,
objective_models=objective_models,
constraint_models=constraint_models,
w=objective_weights,
rho=0.5,
ref_z=ref_z,
beta_start=args.beta_start,
beta_end=args.beta_end,
num_steps=args.num_steps,
num_candidates=args.num_candidates,
num_rollouts=args.num_rollouts,
max_len_cap=args.max_len_cap,
num_final_rollouts=args.num_final_rollouts,
cfg=cfg,
tokenizer=tokenizer,
pam_masker=pam_masker,
pam_mask_refresh_every=max(1, args.pam_refresh_every),
pam_debug=args.pam_debug,
pam_scale_edits=args.pam_scale_edits,
pam_edit_scale_factor=args.pam_edit_scale_factor,
deletion_rate_scale=args.deletion_rate_scale,
zero_lam_ins=args.zero_lam_ins,
legacy_beta_incumbent=args.legacy_beta_incumbent,
)
# End timing for this generation
generation_time = time.time() - generation_start_time
out_str = tokenizer.batch_decode(x_T.tolist(), skip_special_tokens=True)[0]
if args.instrument and INSTR is not None:
_rec = {"generation_time": generation_time, "final_len": len(out_str.replace(' ', ''))}
_rec.update({f"time_{k}": v for k, v in INSTR["timers"].items()})
_rec.update({f"n_{k}": v for k, v in INSTR["counts"].items()})
_rec["first_feasible_terminal_at"] = INSTR.get("first_feasible_terminal_at")
_rec["legacy_beta_incumbent"] = bool(args.legacy_beta_incumbent)
instr_records.append(_rec)
print("----------------------------")
print(f"\nDesigned Sequence: {out_str}\n")
print("Final scores:")
scores = compute_scores_print(
[out_str],
objective_models, constraint_models,
device, return_scores=True
).squeeze(0)
# Only count + save on success
valid += 1
orig_len = len(args.input.replace(' ', ''))
final_len = len(out_str.replace(' ', ''))
length_diff = orig_len - final_len
if args.output_file is not None:
import csv
objective_scores_only = scores[:len(objective_models)]
row = _build_csv_row(
sequence=out_str,
final_len=final_len,
length_diff=length_diff,
objective_scores_only=objective_scores_only,
generation_time=generation_time,
)
generated_csv_rows.append(row)
with open(args.output_file, 'a', newline='') as f:
writer = csv.DictWriter(f, fieldnames=csv_header_parts)
writer.writerow(row)
except torch.cuda.OutOfMemoryError:
# OOM error occurred - discard this generation and restart a new one
print(f"[WARN] CUDA OOM during pCoMol generation (attempt {attempt}/{max_attempts}). Discarding this generation and restarting.")
print(f" Current progress: {valid}/{target_valid} valid sequences generated")
# Clear CUDA cache to free memory
torch.cuda.empty_cache()
torch.cuda.ipc_collect() # optional, can help in some cases
# Continue to next iteration - this will start a new generation attempt
# Note: attempt counter already incremented, so this doesn't count as a valid sequence
continue
if args.instrument and instr_records:
import json
_iout = args.instr_out or ((args.output_file.rsplit('.', 1)[0] + '_instr.json') if args.output_file else 'pcomol_instr.json')
with open(_iout, 'w') as _f:
json.dump(instr_records, _f, indent=2)
print(f"[INSTRUMENT] wrote {len(instr_records)} per-sequence records to {_iout}")
usable_csv_rows = [row for row in generated_csv_rows if int(row.get("usable_pass", 0)) == 1]
total_generated = len(generated_csv_rows)
total_usable = len(usable_csv_rows)
final_wall_elapsed_s = float(time.time() - run_start_time)
avg_usable_generation_time = final_wall_elapsed_s / total_usable if total_usable else ""
print(
"[TIME_BUDGET_SUMMARY] "
f"usable_sequences={total_usable}, total_sequences={total_generated}, "
f"avg_usable_generation_time={avg_usable_generation_time if avg_usable_generation_time != '' else 'nan'}, "
f"wall_elapsed_s={final_wall_elapsed_s:.2f}"
)
# Append an average row over generated outputs only (excludes the initial input row).
if args.output_file is not None and generated_csv_rows:
import csv
avg_row = {k: "" for k in csv_header_parts}
for col in csv_header_parts:
vals = []
for row in generated_csv_rows:
v = row.get(col, "")
try:
vals.append(float(v))
except (TypeError, ValueError):
continue
if vals:
avg_row[col] = sum(vals) / len(vals)
avg_row["sequence"] = "AVERAGE_EXCL_INPUT"
if "predicted_pam" in avg_row:
avg_row["predicted_pam"] = ""
if "pam_target_match_binary" in avg_row and avg_row["pam_target_match_binary"] != "":
avg_row["pam_target_match_binary"] = avg_row["pam_target_match_binary"] * 100.0
avg_row["total_sequences_generated"] = total_generated
avg_row["usable_sequences_generated"] = total_usable
avg_row["avg_usable_generation_time"] = avg_usable_generation_time
with open(args.output_file, 'a', newline='') as f:
writer = csv.DictWriter(f, fieldnames=csv_header_parts)
writer.writerow(avg_row)
print(f"[CSV] Appended average row over {len(generated_csv_rows)} generated sequences (excluding input row).")
summary_row = {k: "" for k in csv_header_parts}
summary_row["sequence"] = "SUMMARY_TIME_BUDGET"
summary_row["total_sequences_generated"] = total_generated
summary_row["usable_sequences_generated"] = total_usable
summary_row["avg_usable_generation_time"] = avg_usable_generation_time
summary_row["wall_elapsed_s"] = final_wall_elapsed_s
with open(args.output_file, 'a', newline='') as f:
writer = csv.DictWriter(f, fieldnames=csv_header_parts)
writer.writerow(summary_row)
print("[CSV] Appended time-budget summary row.")
# for _ in range(100):
# x_T = pCoMol(
# model=model,
# x0=x0,
# pad_id=pad_id,
# bos_id=bos_id,
# eos_id=eos_id,
# allowed_tokens=allowed_tokens,
# objective_models=objective_models,
# constraint_models=constraint_models,
# w=objective_weights,
# rho=0.5,
# ref_z=ref_z,
# beta_start=args.beta_start,
# beta_end=args.beta_end,
# num_steps=args.num_steps,
# num_candidates=args.num_candidates,
# num_rollouts=args.num_rollouts,
# max_len_cap=args.max_len_cap,
# num_final_rollouts=args.num_final_rollouts,
# cfg=cfg,
# selfies_tokenizer=selfies_tokenizer,
# smiles_tokenizer=smiles_tokenizer
# )
# out_str = selfies_tokenizer.batch_decode(x_T.tolist())[0]
# smiles_token = smiles_tokenizer(out_str, return_tensors='pt')['input_ids'].to(device)
# print("----------------------------")
# # print(f"Initial Sequence: {args.input}\n")
# # print(f"Initial Scores:")
# # compute_scores_print([args.input], objective_models, constraint_models, device)
# print(f"\nDesigned Sequence: {out_str}\n")
# print("Final scores:")
# scores = compute_scores_print(smiles_token, [out_str], objective_models, constraint_models, device, return_scores=True).squeeze(0)
# with open(args.output_file, 'a') as f:
# f.write(f"{smiles_token.shape[1]}")
# for score in scores:
# f.write(f",{score.item()}")
# f.write(f",{out_str}\n")
if __name__ == "__main__":
main()