(self, x)
| 29 | self.replknet = augment_inputs_network(self.replknet) |
| 30 | |
| 31 | def forward(self, x): |
| 32 | B, N, C, H, W = x.shape |
| 33 | x = x.view(-1, C, H, W) |
| 34 | |
| 35 | pred1 = self.convnext(x) |
| 36 | pred2 = self.replknet(x) |
| 37 | |
| 38 | outputs_score1 = nn.functional.softmax(pred1, dim=1) |
| 39 | outputs_score2 = nn.functional.softmax(pred2, dim=1) |
| 40 | |
| 41 | predict_score1 = outputs_score1[:, 1] |
| 42 | predict_score2 = outputs_score2[:, 1] |
| 43 | |
| 44 | predict_score1 = predict_score1.view(B, N).mean(dim=-1) |
| 45 | predict_score2 = predict_score2.view(B, N).mean(dim=-1) |
| 46 | |
| 47 | return torch.stack((predict_score1, predict_score2), dim=-1).mean(dim=-1) |
| 48 | |
| 49 | |
| 50 | def load_model(model_name, ctg_num, use_sync_bn): |
nothing calls this directly
no outgoing calls
no test coverage detected