"""Evaluate FASTA with a trained one-hot Keras model (optional labeled CSV for metrics).""" from __future__ import annotations import argparse import json import os import sys from pathlib import Path import numpy as np import pandas as pd from tensorflow.keras.layers import Dense, Flatten, Input from tensorflow.keras.models import Sequential, load_model, model_from_json # # ===================================================================== # Default checkpoint: repository `weights/nn_one_hot/` (see evaluation/README.md). # After re-training, replace that folder or pass --model-dir. # ===================================================================== _REPO_ROOT = Path(__file__).resolve().parents[1] MODEL_CHECKPOINT_DIR = _REPO_ROOT / 'weights' / 'nn_one_hot' # _SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) if _SCRIPT_DIR not in sys.path: sys.path.insert(0, _SCRIPT_DIR) import eval_metrics as em # noqa: E402 def _checkpoint_search_roots(model_dir: str) -> list[str]: """Train outputs often put weights in .../model_weights/ while model_parameters.json sits one level up.""" roots: list[str] = [] seen: set[str] = set() def add(p: str) -> None: ap = os.path.abspath(p) if ap not in seen: seen.add(ap) roots.append(ap) add(model_dir) parent = os.path.dirname(os.path.abspath(model_dir)) if parent: add(parent) return roots def load_inference_operating_point(model_dir: str) -> tuple[float, float]: temperature = 1.0 threshold = 0.5 name_candidates = ['one_hot_nn', 'one_hot_model', 'one_hot_embeddings'] roots = _checkpoint_search_roots(model_dir) temp_loaded = False for root in roots: for name in name_candidates: temp_path = os.path.join(root, f'{name}_validation_temperature.csv') if os.path.exists(temp_path): try: temp_df = pd.read_csv(temp_path) if 'temperature' in temp_df.columns and len(temp_df) > 0: temperature = float(temp_df['temperature'].iloc[0]) temp_loaded = True break except Exception as e: print(f'Warning: Could not read temperature from {temp_path}: {e}') if temp_loaded: break if not temp_loaded: print('Note: No validation temperature file found; using temperature=1.0') thr_loaded = False for root in roots: for name in name_candidates: summary_path = os.path.join( root, f'{name}_validation_validation_diagnostics_summary.csv' ) if os.path.exists(summary_path): try: summary_df = pd.read_csv(summary_path) if 'optimal_threshold_f1' in summary_df.columns and len(summary_df) > 0: threshold = float(summary_df['optimal_threshold_f1'].iloc[0]) thr_loaded = True break except Exception as e: print(f'Warning: Could not read threshold from {summary_path}: {e}') if thr_loaded: break if not thr_loaded: print('Note: No optimized threshold file found; using threshold=0.5') return temperature, threshold def resolve_model_dir(cli_value: str | None) -> str: if cli_value: return os.path.abspath(cli_value) return str(MODEL_CHECKPOINT_DIR.resolve()) def parse_arguments(): p = argparse.ArgumentParser( description='Evaluate FASTA sequences with a trained one-hot Keras classifier.', formatter_class=argparse.ArgumentDefaultsHelpFormatter, ) p.add_argument('--fasta', required=True, help='Input FASTA file') p.add_argument( '--model-dir', default=None, help=f'Overrides MODEL_CHECKPOINT_DIR in this script (default: use that constant)', ) p.add_argument('--output', required=True, help='Output directory') p.add_argument( '--target-length', type=int, default=None, help='Sequence padding length (else from full model / model_parameters.json / inferred from weights)', ) p.add_argument( '--csv', default=None, help='Optional CSV with id column (variant/sequence_id/id) and label (highFRET/lowFRET)', ) p.add_argument( '--no-detailed-metrics', action='store_true', help='With --csv: print a one-line summary only; skip diagnostic CSV/report output', ) p.add_argument( '--long-sequences', choices=['truncate', 'skip'], default='truncate', help='Sequences longer than model target_length: keep first target_length residues or skip', ) return p.parse_args() def create_neural_net_model(input_shape, n_units_l1=64, n_units_l2=32): model = Sequential([ Input(shape=input_shape), Flatten(), Dense(n_units_l1, activation='relu'), Dense(n_units_l2, activation='relu'), Dense(2, activation='softmax'), ]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) return model def load_model_parameters(model_dir: str) -> dict | None: for root in _checkpoint_search_roots(model_dir): path = os.path.join(root, 'model_parameters.json') if os.path.isfile(path): with open(path, 'r') as f: return json.load(f) print(f'Warning: model_parameters.json not found in {model_dir} (or its parent directory)') return None def _load_best_hyperparameters_csv(model_dir: str) -> dict | None: for root in _checkpoint_search_roots(model_dir): path = os.path.join(root, 'best_hyperparameters.csv') if os.path.isfile(path): try: df = pd.read_csv(path) if len(df) < 1: return None row = df.iloc[0].to_dict() if 'n_units_l1' in row and 'n_units_l2' in row: return row except Exception as e: print(f'Warning: could not read {path}: {e}') return None def _collect_h5_kernel_shapes(weights_path: str) -> list[tuple[int, int]]: import h5py shapes: list[tuple[int, int]] = [] def visitor(name: str, item: object) -> None: if isinstance(item, h5py.Dataset) and item.ndim == 2 and 'kernel' in name.lower(): shapes.append(tuple(int(x) for x in item.shape)) with h5py.File(weights_path, 'r') as f: f.visititems(visitor) return shapes def infer_one_hot_mlp_dims_from_weights_h5(weights_path: str) -> tuple[int, int, int]: """Recover (target_length, n_units_l1, n_units_l2) from a Keras weights-only H5 (Flatten + 3 Dense).""" shapes = _collect_h5_kernel_shapes(weights_path) shapes = list(dict.fromkeys(shapes)) if len(shapes) < 3: raise ValueError( f'Expected at least 3 Dense kernel matrices in {weights_path}, found {len(shapes)}' ) ends = [s for s in shapes if s[1] == 2] if len(ends) != 1: raise ValueError( f'Expected exactly one output kernel with shape (*, 2); got candidates {ends} in {weights_path}' ) n2, _ = ends[0] mids = [s for s in shapes if s[1] == n2 and s != ends[0]] if len(mids) != 1: raise ValueError(f'Could not resolve middle layer kernel before (*, 2); got {mids}') n1, mid_out = mids[0] if mid_out != n2: raise ValueError('Inconsistent MLP chain (middle layer)') firsts = [s for s in shapes if s[1] == n1 and s[0] > 0 and s[0] % 20 == 0] if len(firsts) != 1: raise ValueError( f'Could not resolve first Dense kernel (flatten_dim %% 20 == 0, out={n1}); got {firsts}' ) flat, _ = firsts[0] target_length = flat // 20 return target_length, n1, n2 def _candidate_weight_paths(model_dir: str) -> list[str]: out: list[str] = [] seen: set[str] = set() for root in _checkpoint_search_roots(model_dir): for rel in ( os.path.join('model_weights', 'one_hot_model_weights.h5'), 'one_hot_model_weights.h5', ): p = os.path.join(root, rel) if os.path.isfile(p) and p not in seen: seen.add(p) out.append(p) return out def _candidate_full_model_paths(model_dir: str) -> list[str]: paths: list[str] = [] seen: set[str] = set() for root in _checkpoint_search_roots(model_dir): for rel in ( os.path.join('model_weights', 'one_hot_model_full_model.h5'), 'one_hot_model_full_model.h5', ): p = os.path.join(root, rel) if os.path.isfile(p) and p not in seen: seen.add(p) paths.append(p) return paths def _candidate_architecture_json_paths(model_dir: str) -> list[str]: paths: list[str] = [] seen: set[str] = set() for root in _checkpoint_search_roots(model_dir): for rel in ( os.path.join('model_weights', 'one_hot_model_architecture.json'), 'one_hot_model_architecture.json', ): p = os.path.join(root, rel) if os.path.isfile(p) and p not in seen: seen.add(p) paths.append(p) return paths def load_model_and_weights(model_dir: str, target_length: int | None): full_paths = _candidate_full_model_paths(model_dir) for path in full_paths: try: model = load_model(path, compile=False) sh = model.input_shape if sh is None or len(sh) < 3: raise ValueError('Unexpected model input shape') inferred_len = int(sh[1]) if target_length is not None and target_length != inferred_len: print( f'Warning: --target-length={target_length} does not match full model ' f'input length {inferred_len}; using {inferred_len}' ) target_length = inferred_len print(f'Loaded full model (architecture + weights): {path}') print(f'Architecture input length: {target_length}') return model, target_length except Exception as e: print(f'Failed full-model load {path}: {e}') n_units_l1: int | None = None n_units_l2: int | None = None model_params = load_model_parameters(model_dir) if model_params and 'one_hot_model' in model_params: model_info = model_params['one_hot_model'] if target_length is None: input_shape = model_info.get('input_shape', (120, 20)) target_length = int(input_shape[0]) print(f'Loaded target_length from model parameters: {target_length}') if 'best_hyperparameters' in model_info: bh = model_info['best_hyperparameters'] n_units_l1 = int(bh.get('n_units_l1', 64)) n_units_l2 = int(bh.get('n_units_l2', 32)) print(f'Architecture (Optuna): [{n_units_l1}, {n_units_l2}]') elif 'hidden_layers' in model_info: n_units_l1, n_units_l2 = (int(x) for x in model_info['hidden_layers']) print(f'Architecture: [{n_units_l1}, {n_units_l2}]') else: n_units_l1, n_units_l2 = 64, 32 else: bh_row = _load_best_hyperparameters_csv(model_dir) if bh_row: n_units_l1 = int(bh_row['n_units_l1']) n_units_l2 = int(bh_row['n_units_l2']) print(f'Architecture from best_hyperparameters.csv: [{n_units_l1}, {n_units_l2}]') if target_length is None and n_units_l1 is None: target_length = 120 print(f'Warning: default target_length={target_length} (no model_parameters.json)') weights_paths = _candidate_weight_paths(model_dir) arch_paths = _candidate_architecture_json_paths(model_dir) for arch in arch_paths: for wpath in weights_paths: if not os.path.isfile(wpath): continue try: with open(arch, 'r') as f: model = model_from_json(f.read()) model.load_weights(wpath) sh = model.input_shape tl = int(sh[1]) if sh and len(sh) >= 3 else None if target_length is None: target_length = tl elif tl is not None and target_length != tl: print( f'Warning: --target-length={target_length} vs architecture.json {tl}; ' f'using architecture value {tl}' ) target_length = tl if target_length is None: raise ValueError('Could not determine target_length') print(f'Loaded model from architecture JSON + weights:\n {arch}\n {wpath}') return model, target_length except Exception as e: print(f'Failed architecture+json load ({arch}, {wpath}): {e}') inferred_from_h5: tuple[int, int, int] | None = None for wpath in weights_paths: if os.path.isfile(wpath): try: inferred_from_h5 = infer_one_hot_mlp_dims_from_weights_h5(wpath) break except Exception as e: print(f'Could not infer architecture from {wpath}: {e}') if inferred_from_h5 is not None: inf_L, inf_u1, inf_u2 = inferred_from_h5 if n_units_l1 is None: n_units_l1, n_units_l2 = inf_u1, inf_u2 print( f'Inferred architecture from weight file: target_length={inf_L}, ' f'hidden=[{n_units_l1}, {n_units_l2}]' ) if target_length is None: target_length = inf_L elif target_length != inf_L: print( f'Warning: --target-length={target_length} vs weights (flatten implies L={inf_L}); ' f'using L={inf_L}' ) target_length = inf_L elif n_units_l1 is None: n_units_l1, n_units_l2 = 64, 32 print('Warning: using default hidden sizes [64, 32] (could not read metadata or infer weights)') if target_length is None: target_length = 120 print(f'Warning: default target_length={target_length}') model = create_neural_net_model((target_length, 20), n_units_l1, n_units_l2) for path in weights_paths: if os.path.exists(path): try: model.load_weights(path) print(f'Loaded weights: {path}') return model, target_length except Exception as e: print(f'Failed weights load {path}: {e}') raise FileNotFoundError(f'Could not load weights in {model_dir}. Tried: {weights_paths}') def read_fasta_robust(fasta_path): current_header = None current_seq = [] with open(fasta_path, 'r') as f: for line_num, line in enumerate(f, 1): line = line.strip() if not line: continue if line.startswith('>'): if current_header is not None: seq = ''.join(current_seq) if seq: yield (current_header, seq) current_header = line[1:].strip() current_seq = [] else: if current_header is None: print(f'Warning: Sequence before header line {line_num}, skipping') continue current_seq.append(line) if current_header is not None: seq = ''.join(current_seq) if seq: yield (current_header, seq) def main(): args = parse_arguments() model_dir = resolve_model_dir(args.model_dir) detailed = args.csv and not args.no_detailed_metrics os.makedirs(args.output, exist_ok=True) print('Loading model...') print(f'Checkpoint directory: {model_dir}') model, target_length = load_model_and_weights(model_dir, args.target_length) temperature, threshold = load_inference_operating_point(model_dir) print(f'Operating point: temperature={temperature:.4f}, threshold={threshold:.4f}') fasta_records = list(read_fasta_robust(args.fasta)) sequence_ids, sequences, original_lengths, truncated_flags, skipped = em.prepare_sequences_from_fasta( fasta_records, target_length, long_sequence_policy=args.long_sequences ) n_trunc = sum(truncated_flags) if n_trunc: print( f'Note: {n_trunc} sequence(s) longer than target_length={target_length}; ' f'using N-terminal {target_length} residues (--long-sequences truncate).' ) for header, length, reason in skipped: print(f'Warning: skip {header} ({reason}, len={length})') if not sequences: print('ERROR: No sequences to process.') return X = em.encode_aa_sequences(sequences, target_length) probs_all = model.predict(X, verbose=0) probs = probs_all[:, 1] probs = em.apply_temperature_scaling(probs, temperature) results = pd.DataFrame( { 'sequence_id': sequence_ids, 'original_length': original_lengths, 'encoded_length': target_length, 'truncated': truncated_flags, 'prediction_probability': probs, 'predicted_class': np.where(probs >= threshold, 'highFRET', 'lowFRET'), } ) out_predictions = os.path.join(args.output, 'one_hot_predictions.csv') results.to_csv(out_predictions, index=False) print(f'Predictions saved: {out_predictions}') if skipped: pd.DataFrame(skipped, columns=['sequence_id', 'length', 'reason']).to_csv( os.path.join(args.output, 'skipped_sequences.csv'), index=False ) if not args.csv or not os.path.isfile(args.csv): print('No labeled CSV; predictions only.') return try: _, _, label_map = em.build_label_map_int(pd.read_csv(args.csv)) except ValueError as e: print(f'Skipping metrics: {e}') return y_true_list: list[int] = [] y_prob_list: list[float] = [] for sid in sequence_ids: key = sid.strip() lab = label_map.get(key) if lab is None: continue y_true_list.append(lab) y_prob_list.append(float(results.loc[results['sequence_id'] == sid, 'prediction_probability'].iloc[0])) if not y_true_list: print('No overlapping labels between FASTA IDs and CSV. Skipping metrics.') return y_true = np.asarray(y_true_list) y_prob = np.asarray(y_prob_list) y_pred = (y_prob >= threshold).astype(int) model_name = 'one_hot_embeddings' em.run_supervised_evaluation( y_true, y_prob, y_pred, model_name, args.output, threshold, detailed_metrics=detailed ) if detailed: print(f'Detailed diagnostics written under {args.output}') if __name__ == '__main__': main()