(csv_file, protein)
| 124 | |
| 125 | |
| 126 | def load_dataset(csv_file, protein): |
| 127 | with open(csv_file, "r") as fin: |
| 128 | reader = csv.reader(fin) |
| 129 | fields = next(reader) |
| 130 | mutations = [] |
| 131 | targets = [] |
| 132 | for i, values in enumerate(reader): |
| 133 | for field, value in zip(fields, values): |
| 134 | if field == "mutant": |
| 135 | mutations.append(value.split(":")) |
| 136 | elif field == "DMS_score": |
| 137 | value = utils.literal_eval(value) |
| 138 | targets.append(value) |
| 139 | |
| 140 | def mutation_site(x): |
| 141 | return [int(y[1:-1])-1 for y in x] |
| 142 | |
| 143 | mutations = [(tuple(mutation_site(mut)), mut, tar) for mut, tar in zip(mutations, targets)] |
| 144 | mutations = sorted(mutations) |
| 145 | sequences = [] |
| 146 | offsets = [] |
| 147 | for i, mut in enumerate(mutations): |
| 148 | if i > 0 and mut[0] == mutations[i-1][0]: |
| 149 | continue |
| 150 | masked_seq = protein.clone() |
| 151 | _mutation_site = mut[0] |
| 152 | node_index = torch.tensor(_mutation_site, dtype=torch.long) |
| 153 | |
| 154 | # truncate long sequences and those only with substructures |
| 155 | if os.path.basename(csv_file) == "POLG_HCVJF_Qi_2014.csv": |
| 156 | start, end = 1981, 2225 |
| 157 | elif os.path.basename(csv_file) == "A0A140D2T1_ZIKV_Sourisseau_2019.csv": |
| 158 | start, end = 290, 794 |
| 159 | elif os.path.basename(csv_file) == "B2L11_HUMAN_Dutta_2010_binding-Mcl-1.csv": |
| 160 | start, end = 119, 197 # keep high plddt part |
| 161 | elif masked_seq.num_residue > 1022: |
| 162 | seq_len = masked_seq.num_residue |
| 163 | start, end = get_optimal_window(mutation_position_relative=mut[0][0], seq_len_wo_special=seq_len, model_window=1022) |
| 164 | else: |
| 165 | start, end = 0, masked_seq.num_residue |
| 166 | node_index = node_index - start |
| 167 | residue_mask = torch.zeros((masked_seq.num_residue, ), dtype=torch.bool) |
| 168 | residue_mask[start:end] = 1 |
| 169 | masked_seq = masked_seq.subresidue(residue_mask) |
| 170 | with masked_seq.graph(): |
| 171 | masked_seq.start = torch.as_tensor(start) |
| 172 | masked_seq.end = torch.as_tensor(end) |
| 173 | offsets.append(start) |
| 174 | |
| 175 | mask_id = task.model.sequence_model.alphabet.get_idx("<mask>") |
| 176 | with masked_seq.residue(): |
| 177 | masked_seq.residue_feature[node_index] = 0 |
| 178 | masked_seq.residue_type[node_index] = mask_id |
| 179 | sequences.append(masked_seq) |
| 180 | |
| 181 | return sequences, mutations, offsets |
| 182 | |
| 183 |
no test coverage detected