MCPcopy Create free account
hub / github.com/VisionRush/DeepFakeDefenders / final_model

Class final_model

toolkit/chelper.py:21–47  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

19
20
21class final_model(nn.Module): # Total parameters: 158.64741325378418 MB
22 def __init__(self):
23 super(final_model, self).__init__()
24
25 self.convnext = convnext_base(num_classes=2)
26 self.convnext = augment_inputs_network(self.convnext)
27
28 self.replknet = create_RepLKNet31B(num_classes=2)
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
50def load_model(model_name, ctg_num, use_sync_bn):

Callers 2

merge.pyFile · 0.90
load_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected