OneScience-Group/UTRGAN
09
1"""Generate 5' UTR candidates and rank them with FramePool and MTtrans."""2 3import argparse4import json5import os6import sys7from pathlib import Path8 9 10PROJECT_ROOT = Path(__file__).resolve().parents[1]11MODEL_ROOT = PROJECT_ROOT / "model"12MODULE_ROOT = MODEL_ROOT / "src" / "mrl_te_optimization"13 14 15def parse_args():16 parser = argparse.ArgumentParser(17 description="Generate UTRGAN candidates and rank by MRL and TE."18 )19 parser.add_argument("--num-candidates", type=int, default=1024)20 parser.add_argument("--batch-size", type=int, default=128)21 parser.add_argument("--seed", type=int, default=33)22 parser.add_argument("--device", choices=("dcu", "cpu"), default="dcu")23 parser.add_argument("--device-id", default="0")24 parser.add_argument(25 "--output-dir",26 default=str(PROJECT_ROOT / "outputs" / "pretrained_batch_ranking"),27 )28 return parser.parse_args()29 30 31def configure_runtime(args):32 os.environ.setdefault("TF_USE_LEGACY_KERAS", "1")33 os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2")34 if args.device == "cpu":35 os.environ["HIP_VISIBLE_DEVICES"] = "-1"36 os.environ["CUDA_VISIBLE_DEVICES"] = "-1"37 else:38 os.environ["HIP_VISIBLE_DEVICES"] = args.device_id39 os.environ["CUDA_VISIBLE_DEVICES"] = args.device_id40 for import_root in (MODEL_ROOT, MODULE_ROOT):41 if str(import_root) not in sys.path:42 sys.path.insert(0, str(import_root))43 44 45def main():46 args = parse_args()47 if args.num_candidates < 1 or args.batch_size < 1:48 raise ValueError("--num-candidates and --batch-size must be positive")49 configure_runtime(args)50 51 import numpy as np52 import pandas as pd53 import tensorflow as tf54 import torch55 56 import framepool57 import util58 59 output_dir = Path(args.output_dir).expanduser().resolve()60 output_dir.mkdir(parents=True, exist_ok=True)61 62 generator_path = PROJECT_ROOT / "weight" / "checkpoint_3000.h5"63 framepool_path = PROJECT_ROOT / "weight" / "utr_model_combined_residual_new.h5"64 mttrans_path = (65 PROJECT_ROOT66 / "weight"67 / "mttrans"68 / "RL_hard_share_MTL"69 / "3R"70 / "schedule_MTL-model_best_cv1.pth"71 )72 for path in (generator_path, framepool_path, mttrans_path):73 if not path.is_file():74 raise FileNotFoundError(path)75 76 tf_device = "/GPU:0" if args.device == "dcu" else "/CPU:0"77 torch_device = torch.device("cuda:0" if args.device == "dcu" else "cpu")78 if args.device == "dcu":79 tf_gpus = tf.config.list_physical_devices("GPU")80 if not tf_gpus:81 raise RuntimeError("TensorFlow did not detect a DCU")82 if not torch.cuda.is_available():83 raise RuntimeError("PyTorch did not detect a DCU")84 for gpu in tf_gpus:85 try:86 tf.config.experimental.set_memory_growth(gpu, True)87 except RuntimeError:88 pass89 90 # Loading on CPU avoids device-side random-initializer kernels; inference91 # is explicitly placed on the requested device below.92 with tf.device("/CPU:0"):93 generator = tf.keras.models.load_model(generator_path, compile=False)94 mrl_model = framepool.load_framepool(str(framepool_path))95 generator.trainable = False96 mrl_model.trainable = False97 98 checkpoint = torch.load(99 mttrans_path, map_location="cpu", weights_only=False100 )101 te_model = checkpoint["state_dict"].to(torch_device)102 te_model.eval()103 104 np.random.seed(args.seed)105 tf.random.set_seed(args.seed)106 torch.manual_seed(args.seed)107 if args.device == "dcu":108 torch.cuda.manual_seed_all(args.seed)109 110 noise = np.random.RandomState(args.seed).normal(111 size=(args.num_candidates, 40)112 ).astype(np.float32)113 114 generated_batches = []115 with tf.device(tf_device):116 for start in range(0, args.num_candidates, args.batch_size):117 stop = min(start + args.batch_size, args.num_candidates)118 generated_batches.append(119 generator(tf.convert_to_tensor(noise[start:stop]), training=False).numpy()120 )121 generated = np.concatenate(generated_batches, axis=0)122 if generated.shape != (args.num_candidates, 128, 5):123 raise RuntimeError(f"Unexpected generator shape: {generated.shape}")124 if not np.isfinite(generated).all():125 raise RuntimeError("Generator output contains NaN/Inf")126 127 sequences = list(util.recover_seq(generated, util.rev_rna_vocab))128 mrl_scores = []129 with tf.device(tf_device):130 for start in range(0, len(sequences), args.batch_size):131 chunk = sequences[start : start + args.batch_size]132 encoded = np.asarray(133 [util.encode_seq_framepool(seq) for seq in chunk],134 dtype=np.float32,135 )136 prediction = mrl_model(tf.convert_to_tensor(encoded), training=False)137 mrl_scores.extend(tf.reshape(prediction, (-1,)).numpy().tolist())138 139 te_scores = []140 with torch.inference_mode():141 for start in range(0, len(sequences), args.batch_size):142 chunk = sequences[start : start + args.batch_size]143 encoded = np.asarray(util.one_hot_all_motif(chunk), dtype=np.float32)144 encoded = torch.from_numpy(encoded).transpose(1, 2).to(torch_device)145 prediction = te_model(encoded)146 te_scores.extend(prediction.reshape(-1).cpu().numpy().tolist())147 148 mrl_scores = np.asarray(mrl_scores, dtype=np.float32)149 te_scores = np.asarray(te_scores, dtype=np.float32)150 if not np.isfinite(mrl_scores).all() or not np.isfinite(te_scores).all():151 raise RuntimeError("MRL/TE scores contain NaN/Inf")152 153 table = pd.DataFrame(154 {155 "candidate_id": [156 f"UTRGAN_{index + 1:05d}" for index in range(len(sequences))157 ],158 "sequence": sequences,159 "length": [len(sequence) for sequence in sequences],160 "mrl_score": mrl_scores,161 "te_score": te_scores,162 }163 )164 table["is_duplicate"] = table.duplicated("sequence", keep="first")165 table["mrl_rank"] = table["mrl_score"].rank(166 method="first", ascending=False167 ).astype(int)168 table["te_rank"] = table["te_score"].rank(169 method="first", ascending=False170 ).astype(int)171 unique = table.drop_duplicates("sequence", keep="first").copy()172 173 table.to_csv(output_dir / "all_candidates_scores.csv", index=False)174 unique.sort_values("mrl_score", ascending=False).to_csv(175 output_dir / "ranked_by_mrl.csv", index=False176 )177 unique.sort_values("te_score", ascending=False).to_csv(178 output_dir / "ranked_by_te.csv", index=False179 )180 np.save(output_dir / "generator_probabilities.npy", generated)181 182 summary = {183 "requested_candidates": args.num_candidates,184 "generated_candidates": len(table),185 "unique_sequences": len(unique),186 "duplicate_sequences": int(table["is_duplicate"].sum()),187 "generator_shape": list(generated.shape),188 "generator_probability_max_error": float(189 np.max(np.abs(generated.sum(axis=-1) - 1.0))190 ),191 "length_min": int(table["length"].min()),192 "length_max": int(table["length"].max()),193 "mrl_min": float(mrl_scores.min()),194 "mrl_max": float(mrl_scores.max()),195 "mrl_mean": float(mrl_scores.mean()),196 "te_min": float(te_scores.min()),197 "te_max": float(te_scores.max()),198 "te_mean": float(te_scores.mean()),199 "tensorflow_version": tf.__version__,200 "torch_version": torch.__version__,201 "torch_hip": torch.version.hip,202 "device": args.device,203 "seed": args.seed,204 }205 (output_dir / "summary.json").write_text(206 json.dumps(summary, indent=2), encoding="utf-8"207 )208 print(json.dumps(summary, indent=2))209 print("UTRGAN_PRETRAINED_BATCH_MRL_TE_RANKING_PASS")210 211 212if __name__ == "__main__":213 main()214 