MCPcopy Create free account
hub / github.com/Trustworthy-AI-Group/TransferAttack / EnsembleModel

Class EnsembleModel

transferattack/utils.py:82–105  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

80
81
82class 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
108class AdvDataset(torch.utils.data.Dataset):

Callers 2

load_modelMethod · 0.85
load_modelMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected