| 94 | |
| 95 | |
| 96 | def get_prob(seq_prob, mutations, offsets): |
| 97 | i = 0 |
| 98 | preds = [] |
| 99 | targets = [] |
| 100 | last_sites = None |
| 101 | for j, item in tqdm(enumerate(mutations)): |
| 102 | sites, muts, target = item |
| 103 | if j > 0 and sites != last_sites: |
| 104 | i += 1 |
| 105 | |
| 106 | node_index = torch.tensor(sites, dtype=torch.long) |
| 107 | offset = offsets[i] |
| 108 | node_index = node_index - offset |
| 109 | mt_target = [data.Protein.residue_symbol2id.get(mut[-1], -1) for mut in muts] |
| 110 | wt_target = [data.Protein.residue_symbol2id.get(mut[0], -1) for mut in muts] |
| 111 | log_prob = torch.log_softmax(seq_prob[i], dim=-1) |
| 112 | mt_log_prob = log_prob[node_index, mt_target] |
| 113 | wt_log_prob = log_prob[node_index, wt_target] |
| 114 | log_prob = mt_log_prob - wt_log_prob |
| 115 | score = log_prob.sum(dim=0) |
| 116 | |
| 117 | preds.append(score) |
| 118 | targets.append(target) |
| 119 | last_sites = sites |
| 120 | |
| 121 | pred = torch.stack(preds) |
| 122 | target = torch.tensor(targets).cuda() |
| 123 | return pred, target |
| 124 | |
| 125 | |
| 126 | def load_dataset(csv_file, protein): |