| 8 | # from src.data.batch import Batch |
| 9 | |
| 10 | class Collater(object): |
| 11 | def __init__(self, follow_batch, ndevices): |
| 12 | self.follow_batch = follow_batch |
| 13 | self.ndevices = ndevices |
| 14 | |
| 15 | def collate(self, batch): |
| 16 | elem = batch[0] |
| 17 | if isinstance(elem, Data): |
| 18 | data = batch |
| 19 | count = torch.tensor([data.num_nodes for data in data]) |
| 20 | cumsum = count.cumsum(0) |
| 21 | cumsum = torch.cat([cumsum.new_zeros(1), cumsum], dim=0) |
| 22 | device_id = self.ndevices * cumsum.to(torch.float) / cumsum[-1].item() |
| 23 | device_id = (device_id[:-1] + device_id[1:]) / 2.0 |
| 24 | device_id = device_id.to(torch.long) # round. |
| 25 | split = device_id.bincount().cumsum( |
| 26 | 0) # Count the frequency of each value in an array of non-negative ints. |
| 27 | split = torch.cat([split.new_zeros(1), split], dim=0) |
| 28 | split = torch.unique(split, sorted=True) |
| 29 | split = split.tolist() |
| 30 | |
| 31 | graphs = [] |
| 32 | for i in range(len(split) - 1): |
| 33 | data1 = data[split[i]:split[i + 1]] |
| 34 | graph = Batch.from_data_list(data1, self.follow_batch) |
| 35 | graphs += [graph] |
| 36 | |
| 37 | return graphs #Batch.from_data_list(batch, self.follow_batch) |
| 38 | elif isinstance(elem, torch.Tensor): |
| 39 | return default_collate(batch) |
| 40 | elif isinstance(elem, float): |
| 41 | return torch.tensor(batch, dtype=torch.float) |
| 42 | elif isinstance(elem, int_classes): |
| 43 | return torch.tensor(batch) |
| 44 | elif isinstance(elem, string_classes): |
| 45 | return batch |
| 46 | elif isinstance(elem, container_abcs.Mapping): |
| 47 | return {key: self.collate([d[key] for d in batch]) for key in elem} |
| 48 | elif isinstance(elem, tuple) and hasattr(elem, '_fields'): |
| 49 | return type(elem)(*(self.collate(s) for s in zip(*batch))) |
| 50 | elif isinstance(elem, container_abcs.Sequence): |
| 51 | return [self.collate(s) for s in zip(*batch)] |
| 52 | |
| 53 | raise TypeError('DataLoader found invalid type: {}'.format(type(elem))) |
| 54 | |
| 55 | def __call__(self, batch): |
| 56 | return self.collate(batch) |
| 57 | |
| 58 | |
| 59 | class DataLoader(torch.utils.data.DataLoader): |