Team Ai
Modelpublic

OneScience-Group/UTRGAN

sourceHugging Facecc-by-nc-sa-2.0updated 1mo agoView on Hugging Face
0likes9downloads
optimize_te_mrl.py425 linesDownload Raw Back to scripts
1 2import os3import sys4from pathlib import Path5 6PROJECT_ROOT = Path(__file__).resolve().parents[1]7MODEL_ROOT = PROJECT_ROOT / "model"8MODULE_ROOT = MODEL_ROOT / "src" / "mrl_te_optimization"9for import_root in (MODEL_ROOT, MODULE_ROOT):10    if str(import_root) not in sys.path:11        sys.path.insert(0, str(import_root))12 13os.environ.setdefault("TF_USE_LEGACY_KERAS", "1")14 15from tqdm import tqdm16import random17random.seed(1337)18import matplotlib.pyplot as plt19import argparse20import numpy as np21np.random.seed(1337)22import pandas as pd23import torch24from framepool import *25from util import *26 27import random28random.seed(1337)29import scipy.stats as stats30 31import tensorflow as tf32from tensorflow.keras import backend as K33from tensorflow.keras.models import load_model34 35tf.compat.v1.enable_eager_execution()36 37import pandas as pd38import numpy as np39import requests40 41parser = argparse.ArgumentParser()42parser.add_argument('-d', type=str, required=False,43                    default=str(PROJECT_ROOT / 'conf' / 'data' / 'utrdb2.csv'))44parser.add_argument('-bs', type=int, required=False ,default=64)45parser.add_argument('-lr', type=int, required=False ,default=1)46parser.add_argument('-task', type=str, required=False ,default="mrl")47parser.add_argument('-gpu', type=str, required=False ,default='-1')48parser.add_argument('-s', type=int, required=False ,default=10000)49parser.add_argument('--output-dir', type=str,50                    default=str(PROJECT_ROOT / 'outputs' / 'optimization'))51args = parser.parse_args()52 53if args.gpu == '-1':54    device = 'cpu'55else:56    os.environ['CUDA_VISIBLE_DEVICES'] = args.gpu57    device = 'cuda'58 59def prepare_mttrans(seqs):60    seqs_init = torch.tensor(np.array(one_hot_all_motif(seqs),dtype=np.float32))61 62    seqs_init = torch.transpose(seqs_init, 1, 2)63    seqs_init = torch.tensor(seqs_init,dtype=torch.float32).to(device)64    return seqs_init65 66def prepare_framepool(seqs):67    return tf.convert_to_tensor(np.array([encode_seq_framepool(seq) for seq in seqs]),dtype=tf.float32)68 69 70 71BATCH_SIZE = args.bs72motifs_path = str(PROJECT_ROOT / 'conf' / 'data' / 'motifs.csv')73STEPS = args.s74LR = args.lr75DIM = 4076SEQ_LEN = 12877UTR_LEN = 12878 79TASK = args.task80 81gpath = str(PROJECT_ROOT / 'weight' / 'checkpoint_3000.h5')82 83 84if TASK == 'te':85    path = str(PROJECT_ROOT / 'weight' / 'mttrans' / 'RL_hard_share_MTL' /86               '3R' / 'schedule_MTL-model_best_cv1.pth')87    OPT = 'TE'88else:89    path = str(PROJECT_ROOT / 'weight' / 'utr_model_combined_residual_new.h5')90    OPT = 'FMRL'91 92 93# Check for GPU availability94gpus = tf.config.list_physical_devices('GPU')95 96if gpus:97  print(f"GPU is available. Using GPU:{args.gpu} for computation.")98  print("List of GPUs:", gpus)99else:100  print("GPU is not available. Using CPU instead.")101 102out_folder = str(Path(args.output_dir).expanduser().resolve())103os.makedirs(out_folder, exist_ok=True)104 105 106 107def select_best(scores, seqs):108    selected_scores = []109    selected_seqs = []110    for i in range(len(scores[0])):111        best = scores[0][i]112        best_seq = seqs[0][i]113        for j in range(len(scores)-1):114            if scores[j+1][i] > best:115                best = scores[j+1][i]116                best_seq = seqs[j+1][i]117        selected_scores.append(best)118        selected_seqs.append(best_seq)119 120    return selected_seqs, selected_scores121 122if __name__ == '__main__':123    124    if OPT == 'FMRL':125        Optimize_FrameSlice = True126    else:127        Optimize_FrameSlice = False128 129 130 131    if Optimize_FrameSlice:132        model = load_framepool(path)133 134    else:135 136        model = torch.load(path,map_location=torch.device(device))['state_dict']  137        model.train()   138   139 140    wgan = tf.keras.models.load_model(gpath)141 142    """143    Data:144    """145 146    tf.random.set_seed(33)147    np.random.seed(33)148 149    diffs = []150    init_exps = []151    opt_exps = []152    orig_vals = []153 154    DIM = 40155    MAX_LEN = 128156    LR = np.exp(-LR)157 158    tempnoise = tf.random.normal(shape=[BATCH_SIZE,DIM])159    selectednoise = tempnoise160 161    best = 10162 163    LOW_START = False164 165 166    if LOW_START:167    168        for i in range(10000):169            tempnoise = tf.random.normal(shape=[BATCH_SIZE,DIM])170            sequences = wgan(tempnoise)171 172            seqs_gen = recover_seq(sequences, rev_rna_vocab)173            seqs_str = seqs_gen174 175            shape_ = tf.shape(np.array([encode_seq_framepool(seq) for seq in recover_seq(sequences, rev_rna_vocab)]))176 177            seqs = tf.convert_to_tensor(np.array([encode_seq_framepool(seq) for seq in recover_seq(sequences, rev_rna_vocab)]),dtype=tf.float32)178 179            180            pred =  model(seqs)181 182            t = tf.reshape(pred,(-1))183            t = t.numpy().astype('float')184            score = np.mean(t)185 186            if score < best:187                best = score188                selectednoise = tempnoise189        noise = tf.Variable(selectednoise)190    else:191        noise = tf.Variable(tf.random.normal(shape=[BATCH_SIZE,DIM]))192    193 194    noise_small = tf.random.normal(shape=[BATCH_SIZE,DIM],stddev=1e-4)195 196    optimizer = tf.keras.optimizers.Adam(learning_rate=np.power(np.e,LR))197 198    '''199    Optimization takes place here.200    '''201 202    bind_scores_list = []203    bind_scores_means = []204    sequences_list = []205 206    means = []207    maxes = []208    iters_ = []209 210    OPTIMIZE = True211 212    DNA_SEL = False213 214 215    sequences_init = wgan(noise)216 217    gen_seqs_init = sequences_init.numpy().astype('float')218 219    seqs_gen_init = recover_seq(gen_seqs_init, rev_rna_vocab)220 221    init_pos, init_neg = motif_count(seqs_gen_init,motifs_path)222    223    if Optimize_FrameSlice:224        seqs = prepare_framepool(seqs_gen_init)225 226        seqs_init = prepare_mttrans(seqs_gen_init)227 228        pred_init = model(seqs)229        230    else:231 232 233        one_hots = one_hot_all_motif(np.array(seqs_gen_init))234        seqs = torch.tensor(one_hots,dtype=torch.double)235        seqs = torch.transpose(seqs, 1, 2)236        seqs = seqs.float().to(device)237 238 239        pred_init = model.forward(seqs)240    241    if Optimize_FrameSlice:242 243        t = tf.reshape(pred_init,(-1))244 245        init_t = t.numpy().astype('float')246        247    else:248        249        t = torch.flatten(pred_init)250        t.float()251        252        init_t = t.cpu().detach().numpy()253 254    init_exp = np.mean(init_t)255 256    max_init = np.max(init_t)257 258    min_init = np.min(init_t)259    260    predicted_mrls = []261 262    STEPS = STEPS263 264    seqs_collection = []265    scores_collection = []266    if OPTIMIZE:267        iter_ = 0268        for opt_iter in tqdm(range(int(STEPS))):269            270            with tf.GradientTape() as gtape:271                gtape.watch(noise)272                sequences = wgan(noise)273 274                seqs_gen = recover_seq(sequences, rev_rna_vocab)275                seqs_collection.append(seqs_gen)276                seqs_str = seqs_gen277                278                if Optimize_FrameSlice:279 280                    seqs = tf.convert_to_tensor(np.array([encode_seq_framepool(seq) for seq in recover_seq(sequences, rev_rna_vocab)]),dtype=tf.float32)281                282                else:283                    seqs = torch.tensor(np.array(one_hot_all_motif(seqs_gen),dtype=np.float32))    284 285                if Optimize_FrameSlice:286 287                    with tf.GradientTape() as ptape:288                        ptape.watch(seqs)289 290                        pred =  model(seqs)291                        score = tf.reduce_mean(pred)292                        t = tf.reshape(pred,(-1))293                        mx = t.numpy().astype('float')294                        scores_collection.append(mx)295                        mx = np.max(mx)296                        297                        sum_ = tf.reduce_sum(t).numpy().astype('float')298                        299                        maxes.append(mx)300                        predicted_mrls.append(sum_/BATCH_SIZE)301                        means.append(sum_/BATCH_SIZE)302 303                    g1 = ptape.gradient(score,seqs)304 305                    OPTIMIZE_FULL = False306                    if OPTIMIZE_FULL:307                        tmp_g = g1.numpy().astype('float')308                        tmp_seqs = seqs_gen309                        tmp_lst = np.zeros(shape=(BATCH_SIZE,MAX_LEN,5))310                        for i in range(len(tmp_seqs)):311                            312                            len_ = len(tmp_seqs[i])313                            edited_g = tmp_g[i][:len_,:]314                            edited_g = np.pad(edited_g,((0,MAX_LEN-len_),(0,1)),'constant')   315                            tmp_lst[i] = edited_g   316                        317                        g1 = tf.convert_to_tensor(tmp_lst,dtype=tf.float32)318 319                    else:320                        321                        g1 = tf.pad(g1,tf.constant([[0, 0], [0, 0], [0, 1]]),"CONSTANT")322 323                    g1 = tf.math.scalar_mul(-1.0,g1)324 325                326                else:327                    328                    seqs = torch.transpose(seqs, 1, 2)329                    seqs = seqs.float()330                    seqs = torch.tensor(seqs.to(device), requires_grad=True)331                    pred = model(seqs)332                    pred = torch.flatten(pred)333                    predicted_mrls.append(np.average(pred.cpu().detach().numpy()))334                    scores_collection.append(pred.cpu().detach().numpy())335                    score = torch.mean(pred)336                    t = torch.flatten(pred)337                    mx = t.cpu().detach().numpy()338                    mx = np.max(mx)339                    340                    sum_ = torch.mean(t).cpu().detach().numpy()341                    342                    maxes.append(mx)343                    means.append(sum_/BATCH_SIZE)344                    pred.backward(torch.ones_like(pred))345                    346                    g1 = seqs.grad347                    348                    g1 = g1.cpu().detach().numpy()349                    g1 = tf.convert_to_tensor(g1)350                    g1 = tf.transpose(g1, perm=[0,2,1])351                    g1 = tf.pad(g1,tf.constant([[0, 0], [0, 0], [0, 1]]),"CONSTANT")352                    g1 = tf.math.scalar_mul(-1.0,g1)353                354                355                g2 = gtape.gradient(sequences,noise,output_gradients=g1)356 357            a1 = g2 + noise_small358            change = [(a1,noise)]359            optimizer.apply_gradients(change)360 361            iters_.append(iter_)362            iter_ += 1363 364        best_seqs, best_scores = select_best(scores_collection, seqs_collection)365 366        sequences_opt = wgan(noise)367        368        gen_seqs_opt = sequences_opt.numpy().astype('float')369 370        seqs_gen_opt = recover_seq(gen_seqs_opt, rev_rna_vocab)371 372        opt_pos, opt_neg = motif_count(seqs_gen_opt,motifs_path)373        374        if Optimize_FrameSlice:375            376            seqs_opt = prepare_framepool(seqs_gen_opt)377 378 379        380        else: 381 382            one_hots = np.array(one_hot_all_motif(seqs_gen_opt))383            # print(np.shape(one_hots))384            seqs = torch.tensor(one_hots,dtype=torch.double)385            seqs = torch.transpose(seqs, 1, 2)386            seqs = seqs.float().to(device)387 388        pred_opt = model(seqs)389        390        if Optimize_FrameSlice:391 392            t = tf.reshape(pred_opt,(-1))393            394            opt_t = t.numpy().astype('float')395            396        else:397            398            t = torch.flatten(pred_opt)399        400        401            opt_t = t.cpu().detach().numpy()402 403        opt_exp = np.mean(opt_t)404 405        min_opt = np.min(opt_t)406        max_opt = np.max(opt_t)407 408        with open(os.path.join(out_folder, f'init_mrl_{OPT}.txt'), 'w') as f:409            f.writelines([str(x)+'\n' for x in init_t])410 411        with open(os.path.join(out_folder, f'opt_mrl_{OPT}.txt'), 'w') as f:412            f.writelines([str(x)+'\n' for x in best_scores])413 414        with open(os.path.join(out_folder, f'opt_seqs_{OPT}.txt'), 'w') as f:415            f.writelines([str(x)+'\n' for x in best_seqs])416 417        with open(os.path.join(out_folder, f'init_seqs_{OPT}.txt'), 'w') as f:418            f.writelines([str(x)+'\n' for x in seqs_gen_init])419    420 421        print(f"Average Initial Pred: {np.average(init_t)}")422        print(f"Max Initial Pred: {np.max(init_t)}")423        print(f"Average Opt. Pred: {np.average(best_scores)}")424        print(f"Max Opt. Pred: {np.max(best_scores)}")425