MCPcopy Create free account
hub / github.com/THUNLP-MT/MEAN / PaddingCollate

Class PaddingCollate

evaluation/ddg/utils/data.py:8–68  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class PaddingCollate(object):
9
10 def __init__(self, length_ref_key='mutation_mask', pad_values={'aa': 20, 'pos14': float('999'), 'icode': ' ', 'chain_id': '-'}, donot_pad={'foldx'}, eight=False):
11 super().__init__()
12 self.length_ref_key = length_ref_key
13 self.pad_values = pad_values
14 self.donot_pad = donot_pad
15 self.eight = eight
16
17 def _pad_last(self, x, n, value=0):
18 if isinstance(x, torch.Tensor):
19 assert x.size(0) <= n
20 if x.size(0) == n:
21 return x
22 pad_size = [n - x.size(0)] + list(x.shape[1:])
23 pad = torch.full(pad_size, fill_value=value).to(x)
24 return torch.cat([x, pad], dim=0)
25 elif isinstance(x, list):
26 pad = [value] * (n - len(x))
27 return x + pad
28 elif isinstance(x, str):
29 if value == 0: # Won't pad strings if not specified
30 return x
31 pad = value * (n - len(x))
32 return x + pad
33 elif isinstance(x, dict):
34 padded = {}
35 for k, v in x.items():
36 if k in self.donot_pad:
37 padded[k] = v
38 else:
39 padded[k] = self._pad_last(v, n, value=self._get_pad_value(k))
40 return padded
41 else:
42 return x
43
44 @staticmethod
45 def _get_pad_mask(l, n):
46 return torch.cat([
47 torch.ones([l], dtype=torch.bool),
48 torch.zeros([n-l], dtype=torch.bool)
49 ], dim=0)
50
51 def _get_pad_value(self, key):
52 if key not in self.pad_values:
53 return 0
54 return self.pad_values[key]
55
56 def __call__(self, data_list):
57 max_length = max([data[self.length_ref_key].size(0) for data in data_list])
58 if self.eight:
59 max_length = math.ceil(max_length / 8) * 8
60 data_list_padded = []
61 for data in data_list:
62 data_padded = {
63 k: self._pad_last(v, max_length, value=self._get_pad_value(k))
64 for k, v in data.items() if k in ('wt', 'mut', 'ddG', 'mutation_mask', 'index', 'mutation')
65 }

Callers 1

load_wt_mut_pdb_pairFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected