| 80 | |
| 81 | |
| 82 | class EnsembleModel(torch.nn.Module): |
| 83 | def __init__(self, models, mode='mean'): |
| 84 | super(EnsembleModel, self).__init__() |
| 85 | self.device = next(models[0].parameters()).device |
| 86 | for model in models: |
| 87 | model.to(self.device) |
| 88 | self.models = models |
| 89 | self.softmax = torch.nn.Softmax(dim=1) |
| 90 | self.type_name = 'ensemble' |
| 91 | self.num_models = len(models) |
| 92 | self.mode = mode |
| 93 | |
| 94 | def forward(self, x): |
| 95 | outputs = [] |
| 96 | for model in self.models: |
| 97 | outputs.append(model(x)) |
| 98 | outputs = torch.stack(outputs, dim=0) |
| 99 | if self.mode == 'mean': |
| 100 | outputs = torch.mean(outputs, dim=0) |
| 101 | return outputs |
| 102 | elif self.mode == 'ind': |
| 103 | return outputs |
| 104 | else: |
| 105 | raise NotImplementedError |
| 106 | |
| 107 | |
| 108 | class AdvDataset(torch.utils.data.Dataset): |
no outgoing calls
no test coverage detected