| """Generate 5' UTR candidates and rank them with FramePool and MTtrans.""" |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| from pathlib import Path |
|
|
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| MODEL_ROOT = PROJECT_ROOT / "model" |
| MODULE_ROOT = MODEL_ROOT / "src" / "mrl_te_optimization" |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser( |
| description="Generate UTRGAN candidates and rank by MRL and TE." |
| ) |
| parser.add_argument("--num-candidates", type=int, default=1024) |
| parser.add_argument("--batch-size", type=int, default=128) |
| parser.add_argument("--seed", type=int, default=33) |
| parser.add_argument("--device", choices=("dcu", "cpu"), default="dcu") |
| parser.add_argument("--device-id", default="0") |
| parser.add_argument( |
| "--output-dir", |
| default=str(PROJECT_ROOT / "outputs" / "pretrained_batch_ranking"), |
| ) |
| return parser.parse_args() |
|
|
|
|
| def configure_runtime(args): |
| os.environ.setdefault("TF_USE_LEGACY_KERAS", "1") |
| os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2") |
| if args.device == "cpu": |
| os.environ["HIP_VISIBLE_DEVICES"] = "-1" |
| os.environ["CUDA_VISIBLE_DEVICES"] = "-1" |
| else: |
| os.environ["HIP_VISIBLE_DEVICES"] = args.device_id |
| os.environ["CUDA_VISIBLE_DEVICES"] = args.device_id |
| for import_root in (MODEL_ROOT, MODULE_ROOT): |
| if str(import_root) not in sys.path: |
| sys.path.insert(0, str(import_root)) |
|
|
|
|
| def main(): |
| args = parse_args() |
| if args.num_candidates < 1 or args.batch_size < 1: |
| raise ValueError("--num-candidates and --batch-size must be positive") |
| configure_runtime(args) |
|
|
| import numpy as np |
| import pandas as pd |
| import tensorflow as tf |
| import torch |
|
|
| import framepool |
| import util |
|
|
| output_dir = Path(args.output_dir).expanduser().resolve() |
| output_dir.mkdir(parents=True, exist_ok=True) |
|
|
| generator_path = PROJECT_ROOT / "weight" / "checkpoint_3000.h5" |
| framepool_path = PROJECT_ROOT / "weight" / "utr_model_combined_residual_new.h5" |
| mttrans_path = ( |
| PROJECT_ROOT |
| / "weight" |
| / "mttrans" |
| / "RL_hard_share_MTL" |
| / "3R" |
| / "schedule_MTL-model_best_cv1.pth" |
| ) |
| for path in (generator_path, framepool_path, mttrans_path): |
| if not path.is_file(): |
| raise FileNotFoundError(path) |
|
|
| tf_device = "/GPU:0" if args.device == "dcu" else "/CPU:0" |
| torch_device = torch.device("cuda:0" if args.device == "dcu" else "cpu") |
| if args.device == "dcu": |
| tf_gpus = tf.config.list_physical_devices("GPU") |
| if not tf_gpus: |
| raise RuntimeError("TensorFlow did not detect a DCU") |
| if not torch.cuda.is_available(): |
| raise RuntimeError("PyTorch did not detect a DCU") |
| for gpu in tf_gpus: |
| try: |
| tf.config.experimental.set_memory_growth(gpu, True) |
| except RuntimeError: |
| pass |
|
|
| |
| |
| with tf.device("/CPU:0"): |
| generator = tf.keras.models.load_model(generator_path, compile=False) |
| mrl_model = framepool.load_framepool(str(framepool_path)) |
| generator.trainable = False |
| mrl_model.trainable = False |
|
|
| checkpoint = torch.load( |
| mttrans_path, map_location="cpu", weights_only=False |
| ) |
| te_model = checkpoint["state_dict"].to(torch_device) |
| te_model.eval() |
|
|
| np.random.seed(args.seed) |
| tf.random.set_seed(args.seed) |
| torch.manual_seed(args.seed) |
| if args.device == "dcu": |
| torch.cuda.manual_seed_all(args.seed) |
|
|
| noise = np.random.RandomState(args.seed).normal( |
| size=(args.num_candidates, 40) |
| ).astype(np.float32) |
|
|
| generated_batches = [] |
| with tf.device(tf_device): |
| for start in range(0, args.num_candidates, args.batch_size): |
| stop = min(start + args.batch_size, args.num_candidates) |
| generated_batches.append( |
| generator(tf.convert_to_tensor(noise[start:stop]), training=False).numpy() |
| ) |
| generated = np.concatenate(generated_batches, axis=0) |
| if generated.shape != (args.num_candidates, 128, 5): |
| raise RuntimeError(f"Unexpected generator shape: {generated.shape}") |
| if not np.isfinite(generated).all(): |
| raise RuntimeError("Generator output contains NaN/Inf") |
|
|
| sequences = list(util.recover_seq(generated, util.rev_rna_vocab)) |
| mrl_scores = [] |
| with tf.device(tf_device): |
| for start in range(0, len(sequences), args.batch_size): |
| chunk = sequences[start : start + args.batch_size] |
| encoded = np.asarray( |
| [util.encode_seq_framepool(seq) for seq in chunk], |
| dtype=np.float32, |
| ) |
| prediction = mrl_model(tf.convert_to_tensor(encoded), training=False) |
| mrl_scores.extend(tf.reshape(prediction, (-1,)).numpy().tolist()) |
|
|
| te_scores = [] |
| with torch.inference_mode(): |
| for start in range(0, len(sequences), args.batch_size): |
| chunk = sequences[start : start + args.batch_size] |
| encoded = np.asarray(util.one_hot_all_motif(chunk), dtype=np.float32) |
| encoded = torch.from_numpy(encoded).transpose(1, 2).to(torch_device) |
| prediction = te_model(encoded) |
| te_scores.extend(prediction.reshape(-1).cpu().numpy().tolist()) |
|
|
| mrl_scores = np.asarray(mrl_scores, dtype=np.float32) |
| te_scores = np.asarray(te_scores, dtype=np.float32) |
| if not np.isfinite(mrl_scores).all() or not np.isfinite(te_scores).all(): |
| raise RuntimeError("MRL/TE scores contain NaN/Inf") |
|
|
| table = pd.DataFrame( |
| { |
| "candidate_id": [ |
| f"UTRGAN_{index + 1:05d}" for index in range(len(sequences)) |
| ], |
| "sequence": sequences, |
| "length": [len(sequence) for sequence in sequences], |
| "mrl_score": mrl_scores, |
| "te_score": te_scores, |
| } |
| ) |
| table["is_duplicate"] = table.duplicated("sequence", keep="first") |
| table["mrl_rank"] = table["mrl_score"].rank( |
| method="first", ascending=False |
| ).astype(int) |
| table["te_rank"] = table["te_score"].rank( |
| method="first", ascending=False |
| ).astype(int) |
| unique = table.drop_duplicates("sequence", keep="first").copy() |
|
|
| table.to_csv(output_dir / "all_candidates_scores.csv", index=False) |
| unique.sort_values("mrl_score", ascending=False).to_csv( |
| output_dir / "ranked_by_mrl.csv", index=False |
| ) |
| unique.sort_values("te_score", ascending=False).to_csv( |
| output_dir / "ranked_by_te.csv", index=False |
| ) |
| np.save(output_dir / "generator_probabilities.npy", generated) |
|
|
| summary = { |
| "requested_candidates": args.num_candidates, |
| "generated_candidates": len(table), |
| "unique_sequences": len(unique), |
| "duplicate_sequences": int(table["is_duplicate"].sum()), |
| "generator_shape": list(generated.shape), |
| "generator_probability_max_error": float( |
| np.max(np.abs(generated.sum(axis=-1) - 1.0)) |
| ), |
| "length_min": int(table["length"].min()), |
| "length_max": int(table["length"].max()), |
| "mrl_min": float(mrl_scores.min()), |
| "mrl_max": float(mrl_scores.max()), |
| "mrl_mean": float(mrl_scores.mean()), |
| "te_min": float(te_scores.min()), |
| "te_max": float(te_scores.max()), |
| "te_mean": float(te_scores.mean()), |
| "tensorflow_version": tf.__version__, |
| "torch_version": torch.__version__, |
| "torch_hip": torch.version.hip, |
| "device": args.device, |
| "seed": args.seed, |
| } |
| (output_dir / "summary.json").write_text( |
| json.dumps(summary, indent=2), encoding="utf-8" |
| ) |
| print(json.dumps(summary, indent=2)) |
| print("UTRGAN_PRETRAINED_BATCH_MRL_TE_RANKING_PASS") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|