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

Class INFER_API

main_infer.py:69–132  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

67
68
69class INFER_API:
70
71 _instance = None
72
73 def __new__(cls):
74 if cls._instance is None:
75 cls._instance = super(INFER_API, cls).__new__(cls)
76 cls._instance.initialize()
77 return cls._instance
78
79 def initialize(self):
80 self.transformer_ = [create_transforms_inference(h=512, w=512),
81 create_transforms_inference1(h=512, w=512),
82 create_transforms_inference2(h=512, w=512),
83 create_transforms_inference3(h=512, w=512),
84 create_transforms_inference4(h=512, w=512),
85 create_transforms_inference5(h=512, w=512)]
86 self.srm = SRMConv2d_simple()
87
88 # model init
89 self.model = load_model('all', 2)
90 model_path = './final_model_csv/final_model.pth'
91 self.model = extract_model_from_pth(model_path, self.model)
92
93 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
94 self.model = self.model.to(device)
95
96 self.model.eval()
97
98 def _add_new_channels_worker(self, image):
99 new_channels = []
100
101 image = einops.rearrange(image, "h w c -> c h w")
102 image = (image - torch.as_tensor(timm.data.constants.IMAGENET_DEFAULT_MEAN).view(-1, 1, 1)) / torch.as_tensor(
103 timm.data.constants.IMAGENET_DEFAULT_STD).view(-1, 1, 1)
104 srm = self.srm(image.unsqueeze(0)).squeeze(0)
105 new_channels.append(einops.rearrange(srm, "c h w -> h w c").numpy())
106
107 new_channels = np.concatenate(new_channels, axis=2)
108 return torch.from_numpy(new_channels).float()
109
110 def add_new_channels(self, images):
111 images_copied = einops.rearrange(images, "c h w -> h w c")
112 new_channels = self._add_new_channels_worker(images_copied)
113 images_copied = torch.concatenate([images_copied, new_channels], dim=-1)
114 images_copied = einops.rearrange(images_copied, "h w c -> c h w")
115
116 return images_copied
117
118 def test(self, img_path):
119 # img load
120 img_data = Image.open(img_path).convert('RGB')
121
122 # transform
123 all_data = []
124 for transform in self.transformer_:
125 current_data = transform(img_data)
126 current_data = self.add_new_channels(current_data)

Callers 3

infer_api.pyFile · 0.90
inter_apiFunction · 0.90
mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected