Team Ai
Modelpublic

OneScience-Group/UTRGAN

sourceHugging Facecc-by-nc-sa-2.0updated 1mo agoView on Hugging Face
0likes9downloads
predict.py214 linesDownload Raw Back to scripts
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