MCPcopy Create free account
hub / github.com/DeepGraphLearning/S3F / load_dataset

Function load_dataset

script/evaluate.py:126–181  ·  view source on GitHub ↗
(csv_file, protein)

Source from the content-addressed store, hash-verified

124
125
126def 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

Callers 1

evaluate.pyFile · 0.85

Calls 2

mutation_siteFunction · 0.85
get_optimal_windowFunction · 0.85

Tested by

no test coverage detected