Download cas9/evaluate_val.py from ChatterjeeLab/pCoMole: direct link, hf CLI and curl.
- Browser
- Download file 47.6 kB
-
https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/evaluate_val.py
- Command line
-
hf download hf://ChatterjeeLab/pCoMole/cas9/evaluate_val.py
-
curl -L -o evaluate_val.py https://huggingface.co/ChatterjeeLab/pCoMole/resolve/main/cas9/evaluate_val.py
47.6 kB
| #!/usr/bin/env python3 | |
| """ | |
| Script to evaluate Editflows models by generating sequences and evaluating them. | |
| 1. Decodes test set into actual string sequences | |
| 2. Calculates diversity, sequence loss (val_unweighted_total_loss), and optionally Cas9 scores for decoded sequences | |
| 3. Uses multi-step generation (generate_from_x0_multi_edit) to sample new sequences | |
| 4. Evaluates generated sequences on diversity, sequence loss, and optionally Cas9 scores | |
| 5. Calculates and plots pLDDT scores for both decoded and generated sequences (as separate subplots) | |
| If Cas9 classifier is not provided, only diversity and sequence loss metrics will be calculated. | |
| """ | |
| import os | |
| import argparse | |
| import torch | |
| from datasets import load_from_disk | |
| import yaml | |
| from easydict import EasyDict as edict | |
| from tqdm import tqdm | |
| import numpy as np | |
| import random | |
| from collections import Counter | |
| import matplotlib.pyplot as plt | |
| from transformers import AutoTokenizer, EsmForProteinFolding | |
| from cas9.generate import build_model_and_stuff, tokenize_input_str, detokenize_output | |
| from cas9.model.utils import generate_from_x0_multi_edit | |
| from cas9.objectives import Cas9Classification | |
| # Try to import rapidfuzz for Levenshtein diversity (optional) | |
| HAS_RAPIDFUZZ = False | |
| HAS_LEVENSHTEIN = False | |
| try: | |
| from rapidfuzz.distance import Levenshtein as RapidLevenshtein | |
| HAS_RAPIDFUZZ = True | |
| except ImportError: | |
| try: | |
| from Levenshtein import distance as levenshtein_distance_fallback | |
| HAS_LEVENSHTEIN = True | |
| except ImportError: | |
| print("Note: rapidfuzz and python-Levenshtein not available. Levenshtein diversity will be skipped.") | |
| def decode_validation_set(val_dataset, tokenizer, pad_id, bos_id, eos_id, num_samples=1000): | |
| """ | |
| Decode validation dataset into actual string sequences. | |
| Returns: | |
| Tuple of (all_sequences, sampled_sequences) | |
| """ | |
| all_sequences = [] | |
| print(f"Decoding sequences from validation dataset...") | |
| for batch_item in tqdm(val_dataset, desc="Decoding"): | |
| input_ids_batch = batch_item["input_ids"] | |
| # Convert to tensor if needed | |
| if isinstance(input_ids_batch[0], list): | |
| batch_tensor = torch.tensor(input_ids_batch, dtype=torch.long) | |
| elif isinstance(input_ids_batch, torch.Tensor): | |
| batch_tensor = input_ids_batch | |
| else: | |
| batch_tensor = torch.tensor(input_ids_batch, dtype=torch.long) | |
| # Decode each sequence to string | |
| for seq_tensor in batch_tensor: | |
| # Find valid length (non-padding tokens) | |
| seq_list = seq_tensor.tolist() | |
| valid_length = 0 | |
| for i, tok in enumerate(seq_list): | |
| if tok == pad_id: | |
| break | |
| valid_length = i + 1 | |
| seq_list = seq_list[:valid_length] | |
| # Decode to string - tokenizer will handle BOS/EOS automatically | |
| if len(seq_list) > 0: | |
| try: | |
| # For ESM tokenizer (protein), use batch_decode | |
| if hasattr(tokenizer, 'batch_decode'): | |
| seq_str = tokenizer.batch_decode([seq_list], skip_special_tokens=True)[0] | |
| else: | |
| seq_str = tokenizer.decode(seq_list, skip_special_tokens=True) | |
| # Remove spaces (ESM tokenizer adds spaces between tokens) | |
| seq_str = seq_str.replace(" ", "") | |
| # Verify we got a valid sequence string | |
| if seq_str and len(seq_str) > 0: | |
| # Check if it contains valid amino acids | |
| valid_chars = sum(1 for c in seq_str if c in "ACDEFGHIKLMNPQRSTVWY") | |
| if valid_chars > len(seq_str) * 0.9: # At least 90% valid amino acids | |
| all_sequences.append(seq_str) | |
| except Exception as e: | |
| pass | |
| # Sample random subset if needed | |
| if len(all_sequences) > num_samples: | |
| sampled = random.sample(all_sequences, num_samples) | |
| print(f"Sampled {num_samples} sequences from {len(all_sequences)} total") | |
| else: | |
| sampled = all_sequences | |
| print(f"Using all {len(all_sequences)} sequences from validation set") | |
| # Debug: Print first few sequences to verify they're decoded correctly | |
| if len(sampled) > 0: | |
| print(f"\nSample of decoded sequences (first 3):") | |
| for i, seq in enumerate(sampled[:3]): | |
| print(f" Sequence {i+1}: Length={len(seq)}, First 50 chars: {seq[:50]}") | |
| return all_sequences, sampled | |
| def calculate_plddt_from_sequence_string(sequence_string, esmfold_tokenizer, esm_model, device): | |
| """ | |
| Calculate pLDDT score for a sequence string using ESMFold. | |
| Based on generate_and_analyze_plddt.py | |
| """ | |
| try: | |
| tok = esmfold_tokenizer([sequence_string], return_tensors="pt", add_special_tokens=False).to(device) | |
| with torch.no_grad(): | |
| out = esm_model(**tok) | |
| plddt = out.plddt.mean(-1).mean(-1) # Average across both confidence and sequence length | |
| # Handle scalar or tensor output | |
| if plddt.dim() > 0: | |
| plddt = plddt[0] # Take first element if batch dimension exists | |
| return plddt.cpu().item() | |
| except Exception as e: | |
| print(f"Error calculating pLDDT for sequence (length {len(sequence_string)}): {e}") | |
| return None | |
| def calculate_plddt_scores(sequences, esmfold_tokenizer, esm_model, device, batch_size=1): | |
| """ | |
| Calculate pLDDT scores for a list of sequences. | |
| Args: | |
| sequences: List of sequence strings | |
| esmfold_tokenizer: ESMFold tokenizer | |
| esm_model: ESMFold model | |
| device: torch device | |
| batch_size: Batch size for processing (default 1 for ESMFold) | |
| Returns: | |
| List of pLDDT scores (None for failed calculations) | |
| """ | |
| plddt_scores = [] | |
| for i in tqdm(range(0, len(sequences), batch_size), desc="Computing pLDDT scores"): | |
| batch_seqs = sequences[i:i+batch_size] | |
| for seq in batch_seqs: | |
| score = calculate_plddt_from_sequence_string(seq, esmfold_tokenizer, esm_model, device) | |
| plddt_scores.append(score) | |
| return plddt_scores | |
| def plot_plddt_histogram(decoded_plddt, generated_plddt, output_dir, num_bins=50, dataset_type="test"): | |
| """ | |
| Plot histogram comparing pLDDT scores between decoded and generated sequences. | |
| Uses two subplots (one above the other) instead of overlapping distributions. | |
| Args: | |
| decoded_plddt: List of pLDDT scores for decoded sequences | |
| generated_plddt: List of pLDDT scores for generated sequences | |
| output_dir: Directory to save the plot | |
| num_bins: Number of bins for histogram | |
| dataset_type: Type of dataset ("test", "validation", etc.) for labeling | |
| """ | |
| # Filter out None values | |
| decoded_plddt_valid = [s for s in decoded_plddt if s is not None] | |
| generated_plddt_valid = [g for g in generated_plddt if g is not None] | |
| if len(decoded_plddt_valid) == 0 or len(generated_plddt_valid) == 0: | |
| print("Warning: No valid pLDDT scores to plot. Skipping histogram.") | |
| return | |
| # Create output directory if it doesn't exist | |
| os.makedirs(output_dir, exist_ok=True) | |
| # Professional color scheme - clear, distinct colors | |
| color_decoded = '#2E86AB' # Professional blue | |
| color_generated = '#E63946' # Clear red/coral | |
| # Determine bin edges based on all scores | |
| all_scores = decoded_plddt_valid + generated_plddt_valid | |
| min_score = min(all_scores) | |
| max_score = max(all_scores) | |
| bin_edges = np.linspace(min_score, max_score, num_bins + 1) | |
| # Create figure with two subplots stacked vertically | |
| fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(10, 8), sharex=True) | |
| # Top subplot: Decoded sequences (true data) | |
| ax1.hist(decoded_plddt_valid, bins=bin_edges, alpha=0.7, color=color_decoded, | |
| density=True, edgecolor=color_decoded, linewidth=1.2) | |
| ax1.set_ylabel('Density', fontsize=14, fontweight='medium') | |
| ax1.set_title(f'Decoded sequences ({dataset_type} set)', fontsize=13, fontweight='medium', pad=10) | |
| ax1.grid(True, alpha=0.2, linestyle='--', linewidth=0.5) | |
| ax1.spines['top'].set_visible(False) | |
| ax1.spines['right'].set_visible(False) | |
| ax1.spines['left'].set_linewidth(0.8) | |
| ax1.spines['bottom'].set_linewidth(0.8) | |
| ax1.tick_params(axis='both', which='major', labelsize=12, length=4, width=0.8) | |
| # Bottom subplot: Generated sequences | |
| ax2.hist(generated_plddt_valid, bins=bin_edges, alpha=0.7, color=color_generated, | |
| density=True, edgecolor=color_generated, linewidth=1.2) | |
| ax2.set_xlabel('pLDDT Score', fontsize=14, fontweight='medium') | |
| ax2.set_ylabel('Density', fontsize=14, fontweight='medium') | |
| ax2.set_title('Generated sequences', fontsize=13, fontweight='medium', pad=10) | |
| ax2.grid(True, alpha=0.2, linestyle='--', linewidth=0.5) | |
| ax2.spines['top'].set_visible(False) | |
| ax2.spines['right'].set_visible(False) | |
| ax2.spines['left'].set_linewidth(0.8) | |
| ax2.spines['bottom'].set_linewidth(0.8) | |
| ax2.tick_params(axis='both', which='major', labelsize=12, length=4, width=0.8) | |
| # Adjust spacing between subplots | |
| plt.tight_layout() | |
| # Save plot | |
| filename = os.path.join(output_dir, 'plddt_histogram_comparison.png') | |
| plt.savefig(filename, dpi=300, bbox_inches='tight', facecolor='white') | |
| print(f"\nSaved pLDDT histogram to {filename}") | |
| plt.close() | |
| def calculate_sequence_loss(editflow, sequences, tokenizer, pad_id, bos_id, eos_id, device, cfg, batch_size=32): | |
| """ | |
| Calculate sequence loss (val_unweighted_total_loss) for sequences. | |
| Uses the same loss calculation as validation_step in base_models.py. | |
| Args: | |
| editflow: EditFlow LightningModule | |
| sequences: List of sequence strings | |
| tokenizer: Tokenizer | |
| pad_id, bos_id, eos_id: Special token IDs | |
| device: torch device | |
| cfg: Config object | |
| batch_size: Batch size for processing | |
| Returns: | |
| List of loss values (one per sequence) | |
| """ | |
| editflow.eval() | |
| all_losses = [] | |
| print(f"Calculating sequence loss for {len(sequences)} sequences...") | |
| # Process in batches | |
| for i in tqdm(range(0, len(sequences), batch_size), desc="Computing losses"): | |
| batch_seqs = sequences[i:i+batch_size] | |
| # Tokenize batch | |
| x1_batch = [] | |
| for seq_str in batch_seqs: | |
| x1 = tokenize_input_str(seq_str, cfg, tokenizer, bos_id, eos_id, pad_id, device) | |
| x1_batch.append(x1.squeeze(0)) | |
| # Pad to same length | |
| max_len = max(x.shape[0] for x in x1_batch) | |
| padded_batch = [] | |
| for x in x1_batch: | |
| padding = torch.full((max_len - x.shape[0],), pad_id, dtype=torch.long, device=device) | |
| padded_batch.append(torch.cat([x, padding])) | |
| x1_tensor = torch.stack(padded_batch).to(device) | |
| # Calculate loss for this batch | |
| with torch.no_grad(): | |
| # Call preparation to get necessary tensors | |
| if editflow.reparameterize: | |
| lam_total, logits_type, logits_ins, logits_sub, z_t, z_1, x_t, mask, weight, M_t = editflow.preparation(x1_tensor) | |
| if editflow.loc_prop_path: | |
| loss, loss_components = editflow.loss_fn.reparameterized_forward_localized( | |
| lam_total, logits_type, logits_ins, logits_sub, | |
| z_t, z_1, x_t, mask, weight, M_t, editflow.lam_prop, | |
| editflow.eps_id, editflow.bos_id, editflow.eos_id, | |
| editflow.gamma_rate, editflow.gamma_edit, | |
| editflow.use_aux_ce, editflow.aux_ce_weight | |
| ) | |
| else: | |
| loss, loss_components = editflow.loss_fn.reparameterized_forward( | |
| lam_total, logits_type, logits_ins, logits_sub, | |
| z_t, z_1, x_t, mask, weight, editflow.eps_id, | |
| editflow.bos_id, editflow.eos_id, editflow.gamma_rate, | |
| editflow.gamma_edit, editflow.use_aux_ce, editflow.aux_ce_weight | |
| ) | |
| else: | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub, z_t, z_1, x_t, mask, weight, M_t = editflow.preparation(x1_tensor) | |
| if editflow.loc_prop_path: | |
| loss, loss_components = editflow.loss_fn.forward_localized( | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub, | |
| z_t, z_1, x_t, mask, weight, M_t, editflow.lam_prop, | |
| editflow.eps_id, editflow.bos_id, editflow.eos_id, | |
| editflow.use_aux_ce, editflow.aux_ce_weight | |
| ) | |
| else: | |
| loss, loss_components = editflow.loss_fn.forward( | |
| lam_ins, logits_ins, lam_del, lam_sub, logits_sub, | |
| z_t, z_1, x_t, mask, weight, editflow.eps_id, | |
| editflow.bos_id, editflow.eos_id, | |
| editflow.use_aux_ce, editflow.aux_ce_weight | |
| ) | |
| # Extract unweighted total loss (matching val_unweighted_total_loss) | |
| if "loss_total_unweighted" in loss_components: | |
| # For reparameterized models | |
| unweighted_loss = loss_components["loss_total_unweighted"] | |
| elif "loss_base" in loss_components: | |
| # For non-reparameterized models, loss_base is rate + edit (unweighted) | |
| unweighted_loss = loss_components["loss_base"] | |
| else: | |
| # Fallback to total loss | |
| unweighted_loss = loss | |
| # Get per-sequence losses (loss is averaged over batch, so we need to compute per-sequence) | |
| # Since loss is batch-averaged, we'll use the batch loss for all sequences in the batch | |
| # For more accurate per-sequence loss, we'd need to compute individually | |
| batch_loss_value = unweighted_loss.item() | |
| all_losses.extend([batch_loss_value] * len(batch_seqs)) | |
| return all_losses | |
| def kgrams(s: str, k: int): | |
| """Extract k-grams from a string.""" | |
| s = s.strip() | |
| if len(s) < k: | |
| return {s} if s else set() | |
| return {s[i:i+k] for i in range(len(s) - k + 1)} | |
| def jaccard(a: set, b: set) -> float: | |
| """Calculate Jaccard similarity between two sets.""" | |
| if not a and not b: | |
| return 1.0 | |
| inter = len(a & b) | |
| union = len(a | b) | |
| return inter / union if union else 1.0 | |
| def diversity_kmer_jaccard(seqs, k=3, pairs=50000, seed=0): | |
| """ | |
| Diversity = 1 - average Jaccard similarity over random pairs. | |
| Works for variable-length strings. | |
| Returns: | |
| diversity (1 - avg_sim), avg_similarity | |
| """ | |
| if len(seqs) < 2: | |
| return 0.0, 1.0 | |
| rng = random.Random(seed) | |
| grams = [kgrams(s, k) for s in seqs] | |
| n = len(seqs) | |
| # Limit pairs to avoid excessive computation | |
| max_pairs = min(pairs, n * (n - 1) // 2) | |
| if max_pairs == 0: | |
| return 0.0, 1.0 | |
| total_sim = 0.0 | |
| for _ in range(max_pairs): | |
| i = rng.randrange(n) | |
| j = rng.randrange(n - 1) | |
| if j >= i: | |
| j += 1 | |
| total_sim += jaccard(grams[i], grams[j]) | |
| avg_sim = total_sim / max_pairs | |
| return 1.0 - avg_sim, avg_sim | |
| def diversity_levenshtein(seqs, pairs=20000, seed=0): | |
| """ | |
| Diversity = 1 - average normalized Levenshtein similarity over random pairs. | |
| Returns: | |
| diversity (1 - avg_sim), avg_similarity | |
| """ | |
| if len(seqs) < 2: | |
| return 0.0, 1.0 | |
| if not (HAS_RAPIDFUZZ or HAS_LEVENSHTEIN): | |
| # No Levenshtein implementation available | |
| return 0.0, 1.0 | |
| rng = random.Random(seed) | |
| n = len(seqs) | |
| # Limit pairs to avoid excessive computation | |
| max_pairs = min(pairs, n * (n - 1) // 2) | |
| if max_pairs == 0: | |
| return 0.0, 1.0 | |
| total_sim = 0.0 | |
| valid_pairs = 0 | |
| for _ in range(max_pairs): | |
| i = rng.randrange(n) | |
| j = rng.randrange(n - 1) | |
| if j >= i: | |
| j += 1 | |
| if HAS_RAPIDFUZZ: | |
| # Use rapidfuzz for fast normalized similarity | |
| sim = RapidLevenshtein.normalized_similarity(seqs[i], seqs[j]) | |
| elif HAS_LEVENSHTEIN: | |
| # Fallback: use python-Levenshtein distance and normalize | |
| # Simple normalization: 1 - (distance / max_length) | |
| dist = levenshtein_distance_fallback(seqs[i], seqs[j]) | |
| max_len = max(len(seqs[i]), len(seqs[j])) | |
| sim = 1.0 - (dist / max_len) if max_len > 0 else 1.0 | |
| else: | |
| # Should not reach here, but skip if somehow we do | |
| continue | |
| total_sim += sim | |
| valid_pairs += 1 | |
| if valid_pairs == 0: | |
| return 0.0, 1.0 | |
| avg_sim = total_sim / valid_pairs | |
| return 1.0 - avg_sim, avg_sim | |
| def calculate_diversity(sequences): | |
| """ | |
| Calculate diversity metrics for a set of sequences. | |
| Uses k-mer Jaccard diversity and optionally Levenshtein diversity. | |
| Args: | |
| sequences: List of sequence strings | |
| Returns: | |
| Dictionary with diversity metrics: | |
| - unique_count: Number of unique sequences | |
| - uniqueness_ratio: Fraction of unique sequences | |
| - kmer_diversity: k-mer Jaccard diversity (1 - avg_similarity) | |
| - kmer_avg_similarity: Average k-mer Jaccard similarity | |
| - levenshtein_diversity: Levenshtein diversity (if available) | |
| - levenshtein_avg_similarity: Average Levenshtein similarity (if available) | |
| """ | |
| if len(sequences) == 0: | |
| return { | |
| 'unique_count': 0, | |
| 'uniqueness_ratio': 0.0, | |
| 'kmer_diversity': 0.0, | |
| 'kmer_avg_similarity': 1.0, | |
| 'levenshtein_diversity': 0.0, | |
| 'levenshtein_avg_similarity': 1.0 | |
| } | |
| # Unique fraction | |
| unique_count = len(set(sequences)) | |
| uniqueness_ratio = unique_count / len(sequences) if len(sequences) > 0 else 0.0 | |
| # k-mer Jaccard diversity | |
| kmer_div, kmer_sim = diversity_kmer_jaccard(sequences, k=3, pairs=50000, seed=0) | |
| # Levenshtein diversity (if available) | |
| if HAS_RAPIDFUZZ or HAS_LEVENSHTEIN: | |
| lev_div, lev_sim = diversity_levenshtein(sequences, pairs=20000, seed=0) | |
| else: | |
| lev_div, lev_sim = 0.0, 1.0 | |
| return { | |
| 'unique_count': unique_count, | |
| 'uniqueness_ratio': uniqueness_ratio, | |
| 'kmer_diversity': kmer_div, | |
| 'kmer_avg_similarity': kmer_sim, | |
| 'levenshtein_diversity': lev_div, | |
| 'levenshtein_avg_similarity': lev_sim | |
| } | |
| def evaluate_cas9_scores(sequences, cas9_classifier, threshold=0.5): | |
| """ | |
| Evaluate Cas9 scores for sequences. | |
| Args: | |
| sequences: List of sequence strings | |
| cas9_classifier: Cas9Classification object | |
| threshold: Score threshold for validity (default 0.5) | |
| Returns: | |
| validity_rate, average_score, list of scores | |
| """ | |
| if len(sequences) == 0: | |
| return 0.0, 0.0, [] | |
| # Get scores in batches | |
| scores = [] | |
| batch_size = 32 | |
| for i in tqdm(range(0, len(sequences), batch_size), desc="Computing Cas9 scores"): | |
| batch_seqs = sequences[i:i+batch_size] | |
| batch_scores = cas9_classifier.get_scores(batch_seqs) | |
| scores.extend(batch_scores) | |
| scores = np.array(scores) | |
| validity_rate = np.mean(scores > threshold) | |
| avg_score = np.mean(scores) | |
| return validity_rate, avg_score, scores.tolist() | |
| def generate_sequences_multi_edit(model, source_dist, tokenizer, pad_id, bos_id, eos_id, eps_id, | |
| input_sequences, device, cfg, num_steps=20, batch_size=32, num_generations_per_sequence=1): | |
| """ | |
| Generate sequences from input sequences using multi-edit generation. | |
| Args: | |
| model: Editflows model | |
| input_sequences: List of input sequence strings | |
| device: torch device | |
| cfg: Config object | |
| num_steps: Number of generation steps | |
| batch_size: Batch size for generation | |
| num_generations_per_sequence: Number of sequences to generate per input sequence | |
| Returns: | |
| List of generated sequence strings | |
| """ | |
| model.eval() | |
| generated_sequences = [] | |
| # Get allowed tokens | |
| allowed_tokens = torch.tensor( | |
| [tok for tok in source_dist._allowed_tokens if tok not in (eps_id,)], | |
| device=device, | |
| dtype=torch.long, | |
| ) | |
| print(f"Generating sequences using multi-edit generation (num_steps={num_steps}, batch_size={batch_size}, num_generations_per_sequence={num_generations_per_sequence})...") | |
| print(f"Total sequences to generate: {len(input_sequences)} input sequences × {num_generations_per_sequence} = {len(input_sequences) * num_generations_per_sequence}") | |
| # Process in batches | |
| for i in tqdm(range(0, len(input_sequences), batch_size), desc="Generating"): | |
| batch_seqs = input_sequences[i:i+batch_size] | |
| # Generate num_generations_per_sequence for each sequence in the batch | |
| for gen_idx in range(num_generations_per_sequence): | |
| # Tokenize batch | |
| x0_batch = [] | |
| for seq_str in batch_seqs: | |
| x0 = tokenize_input_str(seq_str, cfg, tokenizer, bos_id, eos_id, pad_id, device) | |
| x0_batch.append(x0.squeeze(0)) | |
| # Pad to same length | |
| max_len = max(x.shape[0] for x in x0_batch) | |
| padded_batch = [] | |
| for x in x0_batch: | |
| padding = torch.full((max_len - x.shape[0],), pad_id, dtype=torch.long, device=device) | |
| padded_batch.append(torch.cat([x, padding])) | |
| x0_tensor = torch.stack(padded_batch).to(device) | |
| # Generate using multi-edit generation | |
| with torch.no_grad(): | |
| x_gen = generate_from_x0_multi_edit( | |
| model, | |
| x0_tensor, | |
| pad_id=pad_id, | |
| bos_id=bos_id, | |
| eos_id=eos_id, | |
| allowed_tokens=allowed_tokens, | |
| num_steps=num_steps, | |
| device=device, | |
| ) | |
| # Decode generated sequences | |
| for j in range(x_gen.shape[0]): | |
| gen_seq = detokenize_output(x_gen[j:j+1], cfg, tokenizer, bos_id, eos_id, pad_id) | |
| generated_sequences.append(gen_seq) | |
| return generated_sequences | |
| def get_evaluate_val_argument_parser(): | |
| parser = argparse.ArgumentParser(description='Evaluate Editflows model by generating sequences') | |
| parser.add_argument('--config', type=str, required=True, | |
| help='Path to config YAML file') | |
| parser.add_argument('--ckpt', type=str, required=True, | |
| help='Path to model checkpoint (.ckpt file)') | |
| parser.add_argument('--val_data', type=str, default=None, | |
| help='Path to validation dataset (output from uniref_pre_batching.py). Deprecated: use --test_data instead.') | |
| parser.add_argument('--test_data', type=str, default=None, | |
| help='Path to test dataset (output from uniref_pre_batching.py)') | |
| parser.add_argument('--cas9_classifier_ckpt', type=str, default=None, | |
| help='Path to Cas9 classifier checkpoint (optional, if not provided, validity rate will not be calculated)') | |
| parser.add_argument('--cas9_classifier_config', type=str, default=None, | |
| help='Path to Cas9 classifier config (optional, required if --cas9_classifier_ckpt is provided)') | |
| parser.add_argument('--num_samples', type=int, default=1000, | |
| help='Number of sequences to sample from test set (default: 1000)') | |
| parser.add_argument('--num_generations_per_sequence', type=int, default=1, | |
| help='Number of sequences to generate per sampled sequence (default: 1). Total generated = num_samples × num_generations_per_sequence') | |
| parser.add_argument('--num_steps', type=int, default=20, | |
| help='Number of generation steps (default: 20)') | |
| parser.add_argument('--batch_size', type=int, default=32, | |
| help='Batch size for generation (default: 32)') | |
| parser.add_argument('--validity_threshold', type=float, default=0.5, | |
| help='Cas9 score threshold for validity (default: 0.5)') | |
| parser.add_argument('--device', type=str, default='cuda' if torch.cuda.is_available() else 'cpu', | |
| help='Device to use (cuda/cpu)') | |
| parser.add_argument('--seed', type=int, default=42, | |
| help='Random seed (default: 42)') | |
| parser.add_argument('--output_dir', type=str, default='./evaluation_output', | |
| help='Directory to save pLDDT plots (default: ./evaluation_output)') | |
| parser.add_argument('--output_fasta', type=str, default=None, | |
| help='Path to save generated sequences as FASTA file (default: {output_dir}/generated_sequences.fasta)') | |
| parser.add_argument('--num_bins', type=int, default=50, | |
| help='Number of bins for pLDDT histogram (default: 50)') | |
| parser.add_argument('--calculate_plddt', action='store_true', | |
| help='Enable pLDDT calculation and plotting (default: False, pLDDT calculation is disabled by default)') | |
| parser.add_argument('--include_train', action='store_true', | |
| help='Also compute losses for the train set (default: False)') | |
| return parser | |
| def _set_eval_run_seeds(seed: int) -> None: | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(seed) | |
| def run_evaluate_val( | |
| args, | |
| *, | |
| run_seed=None, | |
| save_fasta=True, | |
| fasta_tag=None, | |
| ): | |
| """ | |
| Run the full evaluate_val pipeline and return metrics as a nested dict of JSON-serializable values. | |
| Args: | |
| args: Namespace from get_evaluate_val_argument_parser().parse_args() (or compatible). | |
| run_seed: If int, seeds random/numpy/torch before stochastic steps. If None, behavior matches | |
| legacy evaluate_val (no seeding). | |
| save_fasta: If False, skip writing generated sequences to FASTA. | |
| fasta_tag: If set (and save_fasta), write to output_dir/generated_sequences_{fasta_tag}.fasta | |
| unless output_fasta is explicitly set on args. | |
| """ | |
| if run_seed is not None: | |
| _set_eval_run_seeds(int(run_seed)) | |
| device = torch.device(args.device) | |
| # Load config | |
| with open(args.config, 'r') as f: | |
| cfg = edict(yaml.safe_load(f)) | |
| # Build model | |
| print("Building model...") | |
| editflow, source_dist, tokenizer, pad_id, bos_id, eos_id, eps_id = build_model_and_stuff(cfg, device) | |
| # Load checkpoint | |
| print(f"Loading checkpoint from {args.ckpt}...") | |
| ckpt = torch.load(args.ckpt, map_location=device, weights_only=False) | |
| editflow.load_state_dict(ckpt["state_dict"], strict=False) | |
| model = editflow.model.to(device) | |
| model.eval() | |
| # Initialize Cas9 classifier (optional) | |
| cas9_classifier = None | |
| if args.cas9_classifier_ckpt is not None: | |
| if args.cas9_classifier_config is None: | |
| raise ValueError("--cas9_classifier_config is required when --cas9_classifier_ckpt is provided") | |
| print("Initializing Cas9 classifier...") | |
| 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 to save memory | |
| ) | |
| # Test classifier on a known Cas9 sequence to verify it works | |
| test_cas9_seq = "MDKKYSIGLDIGTNSVGWAVITDEYKVPSKKFKVLGNTDRHSIKKNLIGALLFDSGETAEATRLKRTARRRYTRRKNRICYLQEIFSNEMAKVDDSFFHRLEESFLVEEDKKHERHPIFGNIVDEVAYHEKYPTIYHLRKKLVDSTDKADLRLIYLALAHMIKFRGHFLIEGDLNPDNSDVDKLFIQLVQTYNQLFEENPINASGVDAKAILSARLSKSRRLENLIAQLPGEKKNGLFGNLIALSLGLTPNFKSNFDLAEDAKLQLSKDTYDDDLDNLLAQIGDQYADLFLAAKNLSDAILLSDILRVNTEITKAPLSASMIKRYDEHHQDLTLLKALVRQQLPEKYKEIFFDQSKNGYAGYIDGGASQEEFYKFIKPILEKMDGTEELLVKLNREDLLRKQRTFDNGSIPHQIHLGELHAILRRQEDFYPFLKDNREKIEKILTFRIPYYVGPLARGNSRFAWMTRKSEETITPWNFEEVVDKGASAQSFIERMTNFDKNLPNEKVLPKHSLLYEYFTVYNELTKVKYVTEGMRKPAFLSGEQKKAIVDLLFKTNRKVTVKQLKEDYFKKIECFDSVEISGVEDRFNASLGTYHDLLKIIKDKDFLDNEENEDILEDIVLTLTLFEDREMIEERLKTYAHLFDDKVMKQLKRRRYTGWGRLSRKLINGIRDKQSGKTILDFLKSDGFANRNFMQLIHDDSLTFKEDIQKAQVSGQGDSLHEHIANLAGSPAIKKGILQTVKVVDELVKVMGRHKPENIVIEMARENQTTQKGQKNSRERMKRIEEGIKELGSQILKEHPVENTQLQNEKLYLYYLQNGRDMYVDQELDINRLSDYDVDHIVPQSFLKDDSIDNKVLTRSDKNRGKSDNVPSEEVVKKMKNYWRQLLNAKLITQRKFDNLTKAERGGLSELDKAGFIKRQLVETRQITKHVAQILDSRMNTKYDENDKLIREVKVITLKSKLVSDFRKDFQFYKVREINNYHHAHDAYLNAVVGTALIKKYPKLESEFVYGDYKVYDVRKMIAKSEQEIGKATAKYFFYSNIMNFFKTEITLANGEIRKRPLIETNGETGEIVWDKGRDFATVRKVLSMPQVNIVKKTEVQTGGFSKESILPKRNSDKLIARKKDWDPKKYGGFDSPTVAYSVLVVAKVEKGKSKKLKSVKELLGITIMERSSFEKNPIDFLEAKGYKEVKKDLIIKLPKYSLFELENGRKRMLASAGELQKGNELALPSKYVNFLYLASHYEKLKGSPEDNEQKQLFVEQHKHYLDEIIEQISEFSKRVILADANLDKVLSAYNKHRDKPIREQAENIIHLFTLTNLGAPAAFKYFDTTIDRKRYTSTKEVLDATLIHQSITGLYETRIDLSQLGGD" | |
| test_score = cas9_classifier.get_scores([test_cas9_seq])[0] | |
| print(f"Test Cas9 classifier on known Cas9 sequence: score = {test_score:.4f}") | |
| if test_score < 0.5: | |
| print(f"WARNING: Known Cas9 sequence got low score! This suggests a problem with the classifier or sequence format.") | |
| else: | |
| print(f"Classifier working correctly (score > 0.5 for known Cas9 sequence)") | |
| else: | |
| print("No Cas9 classifier provided. Validity rate will not be calculated.") | |
| # Determine which dataset path to use (test_data takes precedence) | |
| if args.test_data is not None: | |
| dataset_path = args.test_data | |
| split_name = "test" | |
| dataset_type = "test" | |
| elif args.val_data is not None: | |
| dataset_path = args.val_data | |
| # Try "validation" or "val" first when --val_data is used | |
| split_name = None | |
| dataset_type = "validation" | |
| else: | |
| raise ValueError("Either --test_data or --val_data must be provided") | |
| # Load dataset | |
| print(f"\nLoading {dataset_type} dataset from {dataset_path}...") | |
| dataset_dict = load_from_disk(dataset_path) | |
| # Determine which split to use | |
| if split_name is None: | |
| # When --val_data is used, prioritize validation splits | |
| if "validation" in dataset_dict: | |
| test_dataset = dataset_dict["validation"] | |
| split_name = "validation" | |
| elif "val" in dataset_dict: | |
| test_dataset = dataset_dict["val"] | |
| split_name = "val" | |
| elif "test" in dataset_dict: | |
| # Fallback to test if validation splits not found | |
| test_dataset = dataset_dict["test"] | |
| split_name = "test" | |
| dataset_type = "test" | |
| else: | |
| raise ValueError(f"Dataset must have 'validation', 'val', or 'test' split. Found splits: {list(dataset_dict.keys())}") | |
| else: | |
| if split_name in dataset_dict: | |
| test_dataset = dataset_dict[split_name] | |
| else: | |
| raise ValueError(f"Dataset must have '{split_name}' split. Found splits: {list(dataset_dict.keys())}") | |
| # Load train dataset if --include_train flag is set | |
| train_dataset = None | |
| all_train_sequences = None | |
| train_avg_loss = None | |
| if args.include_train: | |
| if "train" in dataset_dict: | |
| train_dataset = dataset_dict["train"] | |
| print(f"\nLoading train dataset from {dataset_path}...") | |
| print(f"Train dataset loaded: {len(train_dataset)} items") | |
| else: | |
| print(f"Warning: --include_train flag set but 'train' split not found in dataset. Found splits: {list(dataset_dict.keys())}") | |
| print("Skipping train set evaluation.") | |
| # Step 1: Decode test set into string sequences | |
| print(f"\n{'='*70}") | |
| print("STEP 1: Decoding test set") | |
| print(f"{'='*70}") | |
| all_decoded_sequences, sampled_decoded_sequences = decode_validation_set( | |
| test_dataset, | |
| tokenizer, | |
| pad_id, | |
| bos_id, | |
| eos_id, | |
| num_samples=args.num_samples | |
| ) | |
| # Step 2: Calculate metrics on decoded sequences | |
| print(f"\n{'='*70}") | |
| print("STEP 2: Evaluating decoded sequences") | |
| print(f"{'='*70}") | |
| # Calculate sequence loss on FULL test set | |
| print(f"\nCalculating sequence loss for FULL test set ({len(all_decoded_sequences)} sequences)...") | |
| orig_losses_full = calculate_sequence_loss( | |
| editflow, all_decoded_sequences, tokenizer, pad_id, bos_id, eos_id, device, cfg, batch_size=args.batch_size | |
| ) | |
| orig_avg_loss_full = np.mean(orig_losses_full) | |
| print(f" Average sequence loss (full test set): {orig_avg_loss_full:.4f}") | |
| # Calculate other metrics on sampled sequences only | |
| print(f"\nCalculating diversity for sampled decoded sequences ({len(sampled_decoded_sequences)} sequences)...") | |
| orig_diversity = calculate_diversity(sampled_decoded_sequences) | |
| # Calculate sequence loss for sampled sequences (for comparison) | |
| print(f"\nCalculating sequence loss for sampled decoded sequences ({len(sampled_decoded_sequences)} sequences)...") | |
| orig_losses = calculate_sequence_loss( | |
| editflow, sampled_decoded_sequences, tokenizer, pad_id, bos_id, eos_id, device, cfg, batch_size=args.batch_size | |
| ) | |
| orig_avg_loss = np.mean(orig_losses) | |
| print(f" Average sequence loss (sampled): {orig_avg_loss:.4f}") | |
| # Calculate sequence loss for train set if --include_train flag is set | |
| if args.include_train and train_dataset is not None: | |
| print(f"\n{'='*70}") | |
| print("TRAIN SET EVALUATION") | |
| print(f"{'='*70}") | |
| print(f"\nDecoding train set into string sequences...") | |
| all_train_sequences, _ = decode_validation_set( | |
| train_dataset, | |
| tokenizer, | |
| pad_id, | |
| bos_id, | |
| eos_id, | |
| num_samples=10**9 # Decode all sequences, don't sample (use very large number) | |
| ) | |
| print(f"\nCalculating sequence loss for FULL train set ({len(all_train_sequences)} sequences)...") | |
| train_losses = calculate_sequence_loss( | |
| editflow, all_train_sequences, tokenizer, pad_id, bos_id, eos_id, device, cfg, batch_size=args.batch_size | |
| ) | |
| train_avg_loss = np.mean(train_losses) | |
| print(f" Average sequence loss (full train set): {train_avg_loss:.4f}") | |
| # Calculate Cas9 scores only if classifier is provided (on sampled sequences) | |
| orig_validity_rate, orig_avg_score, orig_scores = None, None, None | |
| if cas9_classifier is not None: | |
| print(f"\nCalculating Cas9 scores for sampled decoded sequences...") | |
| orig_validity_rate, orig_avg_score, orig_scores = evaluate_cas9_scores( | |
| sampled_decoded_sequences, cas9_classifier, threshold=args.validity_threshold | |
| ) | |
| # Step 3: Generate new sequences using multi-edit generation | |
| print(f"\n{'='*70}") | |
| print("STEP 3: Generating new sequences using multi-edit generation") | |
| print(f"{'='*70}") | |
| generated_sequences = generate_sequences_multi_edit( | |
| model, source_dist, tokenizer, pad_id, bos_id, eos_id, eps_id, | |
| sampled_decoded_sequences, device, cfg, num_steps=args.num_steps, batch_size=args.batch_size, | |
| num_generations_per_sequence=args.num_generations_per_sequence | |
| ) | |
| print(f"Generated {len(generated_sequences)} sequences") | |
| # Save generated sequences as FASTA file | |
| fasta_path_written = None | |
| if save_fasta: | |
| if args.output_fasta is None: | |
| os.makedirs(args.output_dir, exist_ok=True) | |
| if fasta_tag: | |
| fasta_path = os.path.join(args.output_dir, f"generated_sequences_{fasta_tag}.fasta") | |
| else: | |
| fasta_path = os.path.join(args.output_dir, "generated_sequences.fasta") | |
| else: | |
| fasta_path = args.output_fasta | |
| fasta_dir = os.path.dirname(fasta_path) | |
| if fasta_dir: | |
| os.makedirs(fasta_dir, exist_ok=True) | |
| print(f"\nSaving generated sequences to FASTA file: {fasta_path}") | |
| with open(fasta_path, 'w') as f: | |
| for i, seq in enumerate(generated_sequences): | |
| f.write(f">generated_sequence_{i+1}\n") | |
| for j in range(0, len(seq), 80): | |
| f.write(seq[j:j+80] + "\n") | |
| print(f"Saved {len(generated_sequences)} sequences to {fasta_path}") | |
| fasta_path_written = fasta_path | |
| # Step 4: Calculate metrics on generated sequences | |
| print(f"\n{'='*70}") | |
| print("STEP 4: Evaluating generated sequences") | |
| print(f"{'='*70}") | |
| print(f"\nCalculating diversity for generated sequences...") | |
| gen_diversity = calculate_diversity(generated_sequences) | |
| # Calculate sequence loss for generated sequences | |
| print(f"\nCalculating sequence loss for generated sequences...") | |
| gen_losses = calculate_sequence_loss( | |
| editflow, generated_sequences, tokenizer, pad_id, bos_id, eos_id, device, cfg, batch_size=args.batch_size | |
| ) | |
| gen_avg_loss = np.mean(gen_losses) | |
| print(f" Average sequence loss: {gen_avg_loss:.4f}") | |
| # Calculate Cas9 scores only if classifier is provided | |
| gen_validity_rate, gen_avg_score, gen_scores = None, None, None | |
| if cas9_classifier is not None: | |
| print(f"\nCalculating Cas9 scores for generated sequences...") | |
| gen_validity_rate, gen_avg_score, gen_scores = evaluate_cas9_scores( | |
| generated_sequences, cas9_classifier, threshold=args.validity_threshold | |
| ) | |
| def _to_float(x): | |
| if x is None: | |
| return None | |
| return float(x) | |
| def _diversity_plain(d): | |
| return {k: int(v) if k == "unique_count" else float(v) for k, v in d.items()} | |
| metrics = { | |
| "run_seed": int(run_seed) if run_seed is not None else None, | |
| "dataset_type": dataset_type, | |
| "fasta_path": fasta_path_written, | |
| "counts": { | |
| "full_train": len(all_train_sequences) if all_train_sequences is not None else None, | |
| "full_test": len(all_decoded_sequences), | |
| "sampled_decoded": len(sampled_decoded_sequences), | |
| "generated": len(generated_sequences), | |
| }, | |
| "loss": { | |
| "train_avg": _to_float(train_avg_loss), | |
| "decoded_full_avg": _to_float(orig_avg_loss_full), | |
| "decoded_sampled_avg": _to_float(orig_avg_loss), | |
| "generated_avg": _to_float(gen_avg_loss), | |
| }, | |
| "cas9": None, | |
| "diversity_decoded": _diversity_plain(orig_diversity), | |
| "diversity_generated": _diversity_plain(gen_diversity), | |
| "plddt": None, | |
| } | |
| if cas9_classifier is not None: | |
| metrics["cas9"] = { | |
| "decoded_validity_rate": _to_float(orig_validity_rate), | |
| "decoded_avg_score": _to_float(orig_avg_score), | |
| "generated_validity_rate": _to_float(gen_validity_rate), | |
| "generated_avg_score": _to_float(gen_avg_score), | |
| } | |
| # Step 5: Calculate and plot pLDDT scores (optional) | |
| if args.calculate_plddt: | |
| print(f"\n{'='*70}") | |
| print("STEP 5: Calculating pLDDT scores") | |
| print(f"{'='*70}") | |
| print("\nClearing GPU memory...") | |
| del model | |
| del editflow | |
| if cas9_classifier is not None: | |
| del cas9_classifier | |
| if device.type == 'cuda': | |
| torch.cuda.empty_cache() | |
| print("GPU memory cleared.") | |
| print("Loading ESMFold model for pLDDT calculation...") | |
| esmfold_tokenizer_path = "facebook/esmfold_v1" | |
| esmfold_tokenizer = AutoTokenizer.from_pretrained(esmfold_tokenizer_path) | |
| esmfold_device = device | |
| esm_model = EsmForProteinFolding.from_pretrained( | |
| esmfold_tokenizer_path, | |
| torch_dtype=torch.bfloat16 | |
| ).to(esmfold_device).eval() | |
| print("ESMFold model loaded successfully!") | |
| print(f"\nCalculating pLDDT scores for sampled decoded sequences...") | |
| decoded_plddt = calculate_plddt_scores( | |
| sampled_decoded_sequences, esmfold_tokenizer, esm_model, esmfold_device, batch_size=1 | |
| ) | |
| decoded_plddt_valid = [s for s in decoded_plddt if s is not None] | |
| if len(decoded_plddt_valid) > 0: | |
| print(f" Valid scores: {len(decoded_plddt_valid)}/{len(decoded_plddt)}") | |
| print(f" Mean pLDDT: {np.mean(decoded_plddt_valid):.2f}, Std: {np.std(decoded_plddt_valid):.2f}") | |
| print(f"\nCalculating pLDDT scores for generated sequences...") | |
| generated_plddt = calculate_plddt_scores( | |
| generated_sequences, esmfold_tokenizer, esm_model, esmfold_device, batch_size=1 | |
| ) | |
| generated_plddt_valid = [g for g in generated_plddt if g is not None] | |
| if len(generated_plddt_valid) > 0: | |
| print(f" Valid scores: {len(generated_plddt_valid)}/{len(generated_plddt)}") | |
| print(f" Mean pLDDT: {np.mean(generated_plddt_valid):.2f}, Std: {np.std(generated_plddt_valid):.2f}") | |
| if len(decoded_plddt_valid) > 0 and len(generated_plddt_valid) > 0: | |
| print(f"\nPlotting pLDDT histograms...") | |
| plot_plddt_histogram(decoded_plddt_valid, generated_plddt_valid, args.output_dir, args.num_bins, dataset_type) | |
| else: | |
| print("Warning: Not enough valid pLDDT scores to plot.") | |
| metrics["plddt"] = { | |
| "decoded_mean": float(np.mean(decoded_plddt_valid)) if decoded_plddt_valid else None, | |
| "decoded_std": float(np.std(decoded_plddt_valid)) if decoded_plddt_valid else None, | |
| "decoded_valid_n": len(decoded_plddt_valid), | |
| "decoded_total_n": len(decoded_plddt), | |
| "generated_mean": float(np.mean(generated_plddt_valid)) if generated_plddt_valid else None, | |
| "generated_std": float(np.std(generated_plddt_valid)) if generated_plddt_valid else None, | |
| "generated_valid_n": len(generated_plddt_valid), | |
| "generated_total_n": len(generated_plddt), | |
| } | |
| return metrics | |
| def print_evaluation_results(metrics, args): | |
| """Print the summary block (same layout as historical evaluate_val.py).""" | |
| dataset_type = metrics["dataset_type"] | |
| counts = metrics["counts"] | |
| loss_m = metrics["loss"] | |
| orig_diversity = metrics["diversity_decoded"] | |
| gen_diversity = metrics["diversity_generated"] | |
| print("\n" + "="*70) | |
| print("EVALUATION RESULTS") | |
| print("="*70) | |
| print(f"\nDataset Statistics:") | |
| if counts["full_train"] is not None: | |
| print(f" Full train set sequences: {counts['full_train']}") | |
| print(f" Full test set sequences: {counts['full_test']}") | |
| print(f" Sampled decoded sequences: {counts['sampled_decoded']}") | |
| print(f" Generated sequences: {counts['generated']}") | |
| print(f"\n" + "-"*70) | |
| print("SEQUENCE LOSS (val_unweighted_total_loss)") | |
| print("-"*70) | |
| if loss_m["train_avg"] is not None: | |
| print(f"Train set (FULL train set, {counts['full_train']} sequences):") | |
| print(f" Average loss: {loss_m['train_avg']:.4f}") | |
| print(f"\nDecoded sequences (FULL {dataset_type} set, {counts['full_test']} sequences):") | |
| print(f" Average loss: {loss_m['decoded_full_avg']:.4f}") | |
| print(f"\nDecoded sequences (sampled, {counts['sampled_decoded']} sequences):") | |
| print(f" Average loss: {loss_m['decoded_sampled_avg']:.4f}") | |
| print(f"\nGenerated sequences ({counts['generated']} sequences):") | |
| print(f" Average loss: {loss_m['generated_avg']:.4f}") | |
| cas9 = metrics["cas9"] | |
| if cas9 is not None: | |
| print(f"\n" + "-"*70) | |
| print("CAS9 SCORES (validity threshold = {})".format(args.validity_threshold)) | |
| print("-"*70) | |
| print(f"NOTE: If {dataset_type} set contains general proteins (not Cas9-specific),") | |
| print(" low Cas9 scores are expected. Cas9 classifier is trained to identify Cas9 proteins.") | |
| print(f"\nDecoded sequences (sampled from {dataset_type} set, {counts['sampled_decoded']} sequences):") | |
| print(f" Validity rate: {cas9['decoded_validity_rate']:.4f} ({cas9['decoded_validity_rate']*100:.2f}%)") | |
| print(f" Average Cas9 score: {cas9['decoded_avg_score']:.4f}") | |
| print(f"\nGenerated sequences ({counts['generated']} sequences):") | |
| print(f" Validity rate: {cas9['generated_validity_rate']:.4f} ({cas9['generated_validity_rate']*100:.2f}%)") | |
| print(f" Average Cas9 score: {cas9['generated_avg_score']:.4f}") | |
| print(f"\n" + "-"*70) | |
| print("DIVERSITY METRICS") | |
| print("-"*70) | |
| print(f"Decoded sequences (sampled from {dataset_type} set, {counts['sampled_decoded']} sequences):") | |
| print(f" Unique sequences: {orig_diversity['unique_count']} / {counts['sampled_decoded']} ({orig_diversity['uniqueness_ratio']*100:.2f}%)") | |
| print(f" k-mer Jaccard diversity (k=3): {orig_diversity['kmer_diversity']:.4f}") | |
| print(f" k-mer Jaccard avg similarity: {orig_diversity['kmer_avg_similarity']:.4f}") | |
| if HAS_RAPIDFUZZ or HAS_LEVENSHTEIN: | |
| print(f" Levenshtein diversity: {orig_diversity['levenshtein_diversity']:.4f}") | |
| print(f" Levenshtein avg similarity: {orig_diversity['levenshtein_avg_similarity']:.4f}") | |
| print(f"\nGenerated sequences ({counts['generated']} sequences):") | |
| print(f" Unique sequences: {gen_diversity['unique_count']} / {counts['generated']} ({gen_diversity['uniqueness_ratio']*100:.2f}%)") | |
| print(f" k-mer Jaccard diversity (k=3): {gen_diversity['kmer_diversity']:.4f}") | |
| print(f" k-mer Jaccard avg similarity: {gen_diversity['kmer_avg_similarity']:.4f}") | |
| if HAS_RAPIDFUZZ or HAS_LEVENSHTEIN: | |
| print(f" Levenshtein diversity: {gen_diversity['levenshtein_diversity']:.4f}") | |
| print(f" Levenshtein avg similarity: {gen_diversity['levenshtein_avg_similarity']:.4f}") | |
| plddt = metrics.get("plddt") | |
| if plddt is not None: | |
| print(f"\n" + "-"*70) | |
| print("PLDDT METRICS") | |
| print("-"*70) | |
| print(f"Decoded sequences (sampled from {dataset_type} set, {counts['sampled_decoded']} sequences):") | |
| if plddt["decoded_mean"] is not None: | |
| print(f" Mean pLDDT: {plddt['decoded_mean']:.2f}") | |
| print(f" Std pLDDT: {plddt['decoded_std']:.2f}") | |
| print(f" Valid scores: {plddt['decoded_valid_n']}/{plddt['decoded_total_n']}") | |
| print(f"\nGenerated sequences ({counts['generated']} sequences):") | |
| if plddt["generated_mean"] is not None: | |
| print(f" Mean pLDDT: {plddt['generated_mean']:.2f}") | |
| print(f" Std pLDDT: {plddt['generated_std']:.2f}") | |
| print(f" Valid scores: {plddt['generated_valid_n']}/{plddt['generated_total_n']}") | |
| print("="*70) | |
| def main(): | |
| parser = get_evaluate_val_argument_parser() | |
| args = parser.parse_args() | |
| # Legacy behavior: do not fix RNG (matches previous commented-out seed lines). | |
| metrics = run_evaluate_val(args) | |
| print_evaluation_results(metrics, args) | |
| if __name__ == "__main__": | |
| main() | |