File size: 3,058 Bytes
f211cf7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
#!/usr/bin/env python3

import sys
import os

import random
import torch
import pandas as pd
import numpy as np

from tqdm import tqdm
from datetime import datetime
from transformers import AutoTokenizer, AutoModelForMaskedLM

from src.lm.memdlm.diffusion_module import MembraneDiffusion
from src.sampling.olig_sampler import NOSSampler

from src.utils.generate_utils import calc_blosum_score, calc_ppl
from src.utils.model_utils import _print
from src.utils.config_utils import load_config, repo_path


device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
config = load_config("oligo.yaml")

date = datetime.now().strftime("%Y-%m-%d")




def main():
    csv_save_path = repo_path('results', 'oligo', config.wandb.name, date)
    
    try: os.makedirs(csv_save_path, exist_ok=False)
    except FileExistsError: pass

    tokenizer = AutoTokenizer.from_pretrained(config.lm.pretrained_evoflow)
    
    memdlm = MembraneDiffusion(config).to(device)
    state_dict = memdlm.get_state_dict(str(repo_path("checkpoints", config.lm.ft_evoflow, "best_model.ckpt")))
    memdlm.load_state_dict(state_dict)
    memdlm.eval()

    esm_pth = config.lm.pretrained_esm
    esm_model = AutoModelForMaskedLM.from_pretrained(esm_pth).to(device)
    esm_model.eval()

    generator = NOSSampler(config, device, memdlm, esm_model, tokenizer)

    # Determine length from positive controls
    df = pd.read_csv(str(repo_path('data', 'olig_clf', 'test.csv')))
    seqs = df['Sequence'].tolist()


    generation_results = []
    for seq in tqdm(seqs, desc=f"Generating sequences: "):
        seq_res = []

        seq_len = len(seq)
        tokens = tokenizer(seq, return_tensors='pt')

        gen_seq = ""
        attempts = 0

        while len(gen_seq) != seq_len and attempts < 3:
            gen_seq, og_pred, final_pred = generator.sample_guidance(
                tokens,
                config.olig_guidance.guide_steps,
                config.olig_guidance.diffusion_steps
            )
            attempts += 1

        if len(gen_seq) != seq_len:
            esm_ppl, memdlm_ppl = None, None
        else:
            esm_ppl = calc_ppl(esm_model, tokenizer, gen_seq, [i for i in range(len(gen_seq))], model_type='esm')
            memdlm_ppl = calc_ppl(memdlm, tokenizer, gen_seq, [i for i in range(len(gen_seq))], model_type='diffusion')
            blosum = calc_blosum_score(seq, gen_seq, indices=[i for i in range(len(gen_seq))])

        seq_res.append(seq)
        seq_res.append(gen_seq)
        seq_res.append(og_pred)
        seq_res.append(final_pred)
        seq_res.append(final_pred - og_pred)
        seq_res.append(esm_ppl)
        seq_res.append(memdlm_ppl)
        seq_res.append(blosum)
        generation_results.append(seq_res)

    df = pd.DataFrame(generation_results, columns=['Original Sequence', 'Generated Sequence', 'OG Olig Value', 'New Olig Value', 'Olig Increase', 'ESM PPL', 'MeMDLM PPL', 'MemDLM Blosum'])
    df.to_csv(str(csv_save_path / "seqs_with_ppl.csv"), index=False)


if __name__ == "__main__":
    main()