andro1241 commited on
Commit
7875448
·
verified ·
1 Parent(s): a70da48

Upload 4 files

Browse files
processors/modules/face_enhancer/choices.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import List, Sequence, get_args
2
+
3
+ from facefusion.common_helper import create_float_range, create_int_range
4
+ from facefusion.processors.modules.face_enhancer.types import FaceEnhancerModel
5
+
6
+ face_enhancer_models : List[FaceEnhancerModel] = list(get_args(FaceEnhancerModel))
7
+
8
+ face_enhancer_blend_range : Sequence[int] = create_int_range(0, 100, 1)
9
+
10
+ face_enhancer_weight_range : Sequence[float] = create_float_range(0.0, 1.0, 0.05)
processors/modules/face_enhancer/core.py ADDED
@@ -0,0 +1,437 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from argparse import ArgumentParser
2
+ from functools import lru_cache
3
+ from types import ModuleType
4
+ from typing import List
5
+
6
+ import numpy
7
+
8
+ import facefusion.jobs.job_manager
9
+ import facefusion.jobs.job_store
10
+ from facefusion import config, content_analyser, face_classifier, face_detector, face_landmarker, face_masker, face_recognizer, inference_manager, logger, state_manager, translator, video_manager
11
+ from facefusion.common_helper import create_float_metavar, create_int_metavar, get_middle
12
+ from facefusion.download import conditional_download_hashes, conditional_download_sources, resolve_download_url
13
+ from facefusion.face_creator import scale_face
14
+ from facefusion.face_helper import paste_back, warp_face_by_face_landmark_5
15
+ from facefusion.face_masker import create_box_mask, create_occlusion_mask
16
+ from facefusion.face_selector import select_faces
17
+ from facefusion.filesystem import in_directory, is_image, is_video, resolve_relative_path, same_file_extension
18
+ from facefusion.processors.modules.face_enhancer import choices as face_enhancer_choices
19
+ from facefusion.processors.modules.face_enhancer.types import FaceEnhancerInputs, FaceEnhancerWeight
20
+ from facefusion.processors.types import ProcessorOutputs
21
+ from facefusion.program_helper import find_argument_group
22
+ from facefusion.thread_helper import thread_semaphore
23
+ from facefusion.types import ApplyStateItem, Args, DownloadScope, Face, InferencePool, ModelOptions, ModelSet, ProcessMode, VisionFrame
24
+ from facefusion.vision import blend_frame, read_static_image, read_static_video_frame
25
+
26
+
27
+ @lru_cache()
28
+ def create_static_model_set(download_scope : DownloadScope) -> ModelSet:
29
+ return\
30
+ {
31
+ 'codeformer':
32
+ {
33
+ '__metadata__':
34
+ {
35
+ 'vendor': 'sczhou',
36
+ 'license': 'S-Lab-1.0',
37
+ 'year': 2022
38
+ },
39
+ 'hashes':
40
+ {
41
+ 'face_enhancer':
42
+ {
43
+ 'url': resolve_download_url('models-3.0.0', 'codeformer.hash'),
44
+ 'path': resolve_relative_path('../.assets/models/codeformer.hash')
45
+ }
46
+ },
47
+ 'sources':
48
+ {
49
+ 'face_enhancer':
50
+ {
51
+ 'url': resolve_download_url('models-3.0.0', 'codeformer.onnx'),
52
+ 'path': resolve_relative_path('../.assets/models/codeformer.onnx')
53
+ }
54
+ },
55
+ 'template': 'ffhq_512',
56
+ 'size': (512, 512)
57
+ },
58
+ 'gfpgan_1.2':
59
+ {
60
+ '__metadata__':
61
+ {
62
+ 'vendor': 'TencentARC',
63
+ 'license': 'Apache-2.0',
64
+ 'year': 2022
65
+ },
66
+ 'hashes':
67
+ {
68
+ 'face_enhancer':
69
+ {
70
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.2.hash'),
71
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.2.hash')
72
+ }
73
+ },
74
+ 'sources':
75
+ {
76
+ 'face_enhancer':
77
+ {
78
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.2.onnx'),
79
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.2.onnx')
80
+ }
81
+ },
82
+ 'template': 'ffhq_512',
83
+ 'size': (512, 512)
84
+ },
85
+ 'gfpgan_1.3':
86
+ {
87
+ '__metadata__':
88
+ {
89
+ 'vendor': 'TencentARC',
90
+ 'license': 'Apache-2.0',
91
+ 'year': 2022
92
+ },
93
+ 'hashes':
94
+ {
95
+ 'face_enhancer':
96
+ {
97
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.3.hash'),
98
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.3.hash')
99
+ }
100
+ },
101
+ 'sources':
102
+ {
103
+ 'face_enhancer':
104
+ {
105
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.3.onnx'),
106
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.3.onnx')
107
+ }
108
+ },
109
+ 'template': 'ffhq_512',
110
+ 'size': (512, 512)
111
+ },
112
+ 'gfpgan_1.4':
113
+ {
114
+ '__metadata__':
115
+ {
116
+ 'vendor': 'TencentARC',
117
+ 'license': 'Apache-2.0',
118
+ 'year': 2022
119
+ },
120
+ 'hashes':
121
+ {
122
+ 'face_enhancer':
123
+ {
124
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.4.hash'),
125
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.4.hash')
126
+ }
127
+ },
128
+ 'sources':
129
+ {
130
+ 'face_enhancer':
131
+ {
132
+ 'url': resolve_download_url('models-3.0.0', 'gfpgan_1.4.onnx'),
133
+ 'path': resolve_relative_path('../.assets/models/gfpgan_1.4.onnx')
134
+ }
135
+ },
136
+ 'template': 'ffhq_512',
137
+ 'size': (512, 512)
138
+ },
139
+ 'gpen_bfr_256':
140
+ {
141
+ '__metadata__':
142
+ {
143
+ 'vendor': 'yangxy',
144
+ 'license': 'Non-Commercial',
145
+ 'year': 2021
146
+ },
147
+ 'hashes':
148
+ {
149
+ 'face_enhancer':
150
+ {
151
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_256.hash'),
152
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_256.hash')
153
+ }
154
+ },
155
+ 'sources':
156
+ {
157
+ 'face_enhancer':
158
+ {
159
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_256.onnx'),
160
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_256.onnx')
161
+ }
162
+ },
163
+ 'template': 'arcface_128',
164
+ 'size': (256, 256)
165
+ },
166
+ 'gpen_bfr_512':
167
+ {
168
+ '__metadata__':
169
+ {
170
+ 'vendor': 'yangxy',
171
+ 'license': 'Non-Commercial',
172
+ 'year': 2021
173
+ },
174
+ 'hashes':
175
+ {
176
+ 'face_enhancer':
177
+ {
178
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_512.hash'),
179
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_512.hash')
180
+ }
181
+ },
182
+ 'sources':
183
+ {
184
+ 'face_enhancer':
185
+ {
186
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_512.onnx'),
187
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_512.onnx')
188
+ }
189
+ },
190
+ 'template': 'ffhq_512',
191
+ 'size': (512, 512)
192
+ },
193
+ 'gpen_bfr_1024':
194
+ {
195
+ '__metadata__':
196
+ {
197
+ 'vendor': 'yangxy',
198
+ 'license': 'Non-Commercial',
199
+ 'year': 2021
200
+ },
201
+ 'hashes':
202
+ {
203
+ 'face_enhancer':
204
+ {
205
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_1024.hash'),
206
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_1024.hash')
207
+ }
208
+ },
209
+ 'sources':
210
+ {
211
+ 'face_enhancer':
212
+ {
213
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_1024.onnx'),
214
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_1024.onnx')
215
+ }
216
+ },
217
+ 'template': 'ffhq_512',
218
+ 'size': (1024, 1024)
219
+ },
220
+ 'gpen_bfr_2048':
221
+ {
222
+ '__metadata__':
223
+ {
224
+ 'vendor': 'yangxy',
225
+ 'license': 'Non-Commercial',
226
+ 'year': 2021
227
+ },
228
+ 'hashes':
229
+ {
230
+ 'face_enhancer':
231
+ {
232
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_2048.hash'),
233
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_2048.hash')
234
+ }
235
+ },
236
+ 'sources':
237
+ {
238
+ 'face_enhancer':
239
+ {
240
+ 'url': resolve_download_url('models-3.0.0', 'gpen_bfr_2048.onnx'),
241
+ 'path': resolve_relative_path('../.assets/models/gpen_bfr_2048.onnx')
242
+ }
243
+ },
244
+ 'template': 'ffhq_512',
245
+ 'size': (2048, 2048)
246
+ },
247
+ 'restoreformer_plus_plus':
248
+ {
249
+ '__metadata__':
250
+ {
251
+ 'vendor': 'wzhouxiff',
252
+ 'license': 'Apache-2.0',
253
+ 'year': 2022
254
+ },
255
+ 'hashes':
256
+ {
257
+ 'face_enhancer':
258
+ {
259
+ 'url': resolve_download_url('models-3.0.0', 'restoreformer_plus_plus.hash'),
260
+ 'path': resolve_relative_path('../.assets/models/restoreformer_plus_plus.hash')
261
+ }
262
+ },
263
+ 'sources':
264
+ {
265
+ 'face_enhancer':
266
+ {
267
+ 'url': resolve_download_url('models-3.0.0', 'restoreformer_plus_plus.onnx'),
268
+ 'path': resolve_relative_path('../.assets/models/restoreformer_plus_plus.onnx')
269
+ }
270
+ },
271
+ 'template': 'ffhq_512',
272
+ 'size': (512, 512)
273
+ }
274
+ }
275
+
276
+
277
+ def get_inference_pool() -> InferencePool:
278
+ model_names = [ state_manager.get_item('face_enhancer_model') ]
279
+ model_source_set = get_model_options().get('sources')
280
+
281
+ return inference_manager.get_inference_pool(__name__, model_names, model_source_set)
282
+
283
+
284
+ def clear_inference_pool() -> None:
285
+ model_names = [ state_manager.get_item('face_enhancer_model') ]
286
+ inference_manager.clear_inference_pool(__name__, model_names)
287
+
288
+
289
+ def get_model_options() -> ModelOptions:
290
+ model_name = state_manager.get_item('face_enhancer_model')
291
+ return create_static_model_set('full').get(model_name)
292
+
293
+
294
+ def register_args(program : ArgumentParser) -> None:
295
+ group_processors = find_argument_group(program, 'processors')
296
+ if group_processors:
297
+ group_processors.add_argument('--face-enhancer-model', help = translator.get('help.model', __package__), default = config.get_str_value('processors', 'face_enhancer_model', 'gfpgan_1.4'), choices = face_enhancer_choices.face_enhancer_models)
298
+ group_processors.add_argument('--face-enhancer-blend', help = translator.get('help.blend', __package__), type = int, default = config.get_int_value('processors', 'face_enhancer_blend', '80'), choices = face_enhancer_choices.face_enhancer_blend_range, metavar = create_int_metavar(face_enhancer_choices.face_enhancer_blend_range))
299
+ group_processors.add_argument('--face-enhancer-weight', help = translator.get('help.weight', __package__), type = float, default = config.get_float_value('processors', 'face_enhancer_weight', '0.5'), choices = face_enhancer_choices.face_enhancer_weight_range, metavar = create_float_metavar(face_enhancer_choices.face_enhancer_weight_range))
300
+ facefusion.jobs.job_store.register_step_keys([ 'face_enhancer_model', 'face_enhancer_blend', 'face_enhancer_weight' ])
301
+
302
+
303
+ def apply_args(args : Args, apply_state_item : ApplyStateItem) -> None:
304
+ apply_state_item('face_enhancer_model', args.get('face_enhancer_model'))
305
+ apply_state_item('face_enhancer_blend', args.get('face_enhancer_blend'))
306
+ apply_state_item('face_enhancer_weight', args.get('face_enhancer_weight'))
307
+
308
+
309
+ def get_common_modules() -> List[ModuleType]:
310
+ return [ content_analyser, face_classifier, face_detector, face_landmarker, face_masker, face_recognizer ]
311
+
312
+
313
+ def pre_check() -> bool:
314
+ model_hash_set = get_model_options().get('hashes')
315
+ model_source_set = get_model_options().get('sources')
316
+
317
+ for common_module in get_common_modules():
318
+ if not common_module.pre_check():
319
+ return False
320
+
321
+ return conditional_download_hashes(model_hash_set) and conditional_download_sources(model_source_set)
322
+
323
+
324
+ def pre_process(mode : ProcessMode) -> bool:
325
+ if mode in [ 'output', 'preview' ] and not is_image(state_manager.get_item('target_path')) and not is_video(state_manager.get_item('target_path')):
326
+ logger.error(translator.get('choose_image_or_video_target') + translator.get('exclamation_mark'), __name__)
327
+ return False
328
+ if mode == 'output' and not in_directory(state_manager.get_item('output_path')):
329
+ logger.error(translator.get('specify_image_or_video_output') + translator.get('exclamation_mark'), __name__)
330
+ return False
331
+ if mode == 'output' and not same_file_extension(state_manager.get_item('target_path'), state_manager.get_item('output_path')):
332
+ logger.error(translator.get('match_target_and_output_extension') + translator.get('exclamation_mark'), __name__)
333
+ return False
334
+ return True
335
+
336
+
337
+ def post_process() -> None:
338
+ read_static_image.cache_clear()
339
+ read_static_video_frame.cache_clear()
340
+ video_manager.clear_video_pool()
341
+
342
+ if state_manager.get_item('video_memory_strategy') in [ 'strict', 'moderate' ]:
343
+ clear_inference_pool()
344
+
345
+ if state_manager.get_item('video_memory_strategy') == 'strict':
346
+ for common_module in get_common_modules():
347
+ common_module.clear_inference_pool()
348
+
349
+
350
+ def enhance_face(target_face : Face, temp_vision_frame : VisionFrame) -> VisionFrame:
351
+ model_template = get_model_options().get('template')
352
+ model_size = get_model_options().get('size')
353
+ crop_vision_frame, affine_matrix = warp_face_by_face_landmark_5(temp_vision_frame, target_face.landmark_set.get('5/68'), model_template, model_size)
354
+ box_mask = create_box_mask(crop_vision_frame, state_manager.get_item('face_mask_blur'), (0, 0, 0, 0))
355
+ crop_masks =\
356
+ [
357
+ box_mask
358
+ ]
359
+
360
+ if 'occlusion' in state_manager.get_item('face_mask_types'):
361
+ occlusion_mask = create_occlusion_mask(crop_vision_frame)
362
+ crop_masks.append(occlusion_mask)
363
+
364
+ crop_vision_frame = prepare_crop_frame(crop_vision_frame)
365
+ face_enhancer_weight = numpy.array([ state_manager.get_item('face_enhancer_weight') ]).astype(numpy.double)
366
+ crop_vision_frame = forward(crop_vision_frame, face_enhancer_weight)
367
+ crop_vision_frame = normalize_crop_frame(crop_vision_frame)
368
+ crop_mask = numpy.minimum.reduce(crop_masks).clip(0, 1)
369
+ paste_vision_frame = paste_back(temp_vision_frame, crop_vision_frame, crop_mask, affine_matrix)
370
+ temp_vision_frame = blend_paste_frame(temp_vision_frame, paste_vision_frame)
371
+ return temp_vision_frame
372
+
373
+
374
+ def forward(crop_vision_frame : VisionFrame, face_enhancer_weight : FaceEnhancerWeight) -> VisionFrame:
375
+ face_enhancer = get_inference_pool().get('face_enhancer')
376
+ face_enhancer_inputs = {}
377
+
378
+ for face_enhancer_input in face_enhancer.get_inputs():
379
+ if face_enhancer_input.name == 'input':
380
+ face_enhancer_inputs[face_enhancer_input.name] = crop_vision_frame
381
+ if face_enhancer_input.name == 'weight':
382
+ face_enhancer_inputs[face_enhancer_input.name] = face_enhancer_weight
383
+
384
+ with thread_semaphore():
385
+ crop_vision_frame = face_enhancer.run(None, face_enhancer_inputs)[0][0]
386
+
387
+ return crop_vision_frame
388
+
389
+
390
+ def has_weight_input() -> bool:
391
+ face_enhancer = get_inference_pool().get('face_enhancer')
392
+
393
+ for deep_swapper_input in face_enhancer.get_inputs():
394
+ if deep_swapper_input.name == 'weight':
395
+ return True
396
+
397
+ return False
398
+
399
+
400
+ def prepare_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
401
+ crop_vision_frame = crop_vision_frame[:, :, ::-1] / 255.0
402
+ crop_vision_frame = (crop_vision_frame - 0.5) / 0.5
403
+ crop_vision_frame = numpy.expand_dims(crop_vision_frame.transpose(2, 0, 1), axis = 0).astype(numpy.float32)
404
+ return crop_vision_frame
405
+
406
+
407
+ def normalize_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
408
+ crop_vision_frame = numpy.clip(crop_vision_frame, -1, 1)
409
+ crop_vision_frame = (crop_vision_frame + 1) / 2
410
+ crop_vision_frame = crop_vision_frame.transpose(1, 2, 0)
411
+ crop_vision_frame = (crop_vision_frame * 255.0).round()
412
+ crop_vision_frame = crop_vision_frame.astype(numpy.uint8)[:, :, ::-1]
413
+ return crop_vision_frame
414
+
415
+
416
+ def blend_paste_frame(temp_vision_frame : VisionFrame, paste_vision_frame : VisionFrame) -> VisionFrame:
417
+ face_enhancer_blend = 1 - (state_manager.get_item('face_enhancer_blend') / 100)
418
+ temp_vision_frame = blend_frame(temp_vision_frame, paste_vision_frame, 1 - face_enhancer_blend)
419
+ return temp_vision_frame
420
+
421
+
422
+ def process_frame(inputs : FaceEnhancerInputs) -> ProcessorOutputs:
423
+ reference_vision_frame = inputs.get('reference_vision_frame')
424
+ source_vision_frames = inputs.get('source_vision_frames')
425
+ target_vision_frames = inputs.get('target_vision_frames')
426
+ temp_vision_frame = inputs.get('temp_vision_frame')
427
+ temp_vision_mask = inputs.get('temp_vision_mask')
428
+
429
+ target_vision_frame = get_middle(target_vision_frames)
430
+ target_faces = select_faces(reference_vision_frame, source_vision_frames, target_vision_frames)
431
+
432
+ if target_faces:
433
+ for target_face in target_faces:
434
+ target_face = scale_face(target_face, target_vision_frame, temp_vision_frame)
435
+ temp_vision_frame = enhance_face(target_face, temp_vision_frame)
436
+
437
+ return temp_vision_frame, temp_vision_mask
processors/modules/face_enhancer/locales.py ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from facefusion.types import Locales
2
+
3
+ LOCALES : Locales =\
4
+ {
5
+ 'en':
6
+ {
7
+ 'help':
8
+ {
9
+ 'model': 'choose the model responsible for enhancing the face',
10
+ 'blend': 'blend the enhanced into the previous face',
11
+ 'weight': 'specify the degree of weight applied to the face'
12
+ },
13
+ 'uis':
14
+ {
15
+ 'blend_slider': 'FACE ENHANCER BLEND',
16
+ 'model_dropdown': 'FACE ENHANCER MODEL',
17
+ 'weight_slider': 'FACE ENHANCER WEIGHT'
18
+ }
19
+ }
20
+ }
processors/modules/face_enhancer/types.py ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any, List, Literal, TypeAlias, TypedDict
2
+
3
+ from numpy.typing import NDArray
4
+
5
+ from facefusion.types import Mask, VisionFrame
6
+
7
+ FaceEnhancerInputs = TypedDict('FaceEnhancerInputs',
8
+ {
9
+ 'reference_vision_frame' : VisionFrame,
10
+ 'source_vision_frames' : List[VisionFrame],
11
+ 'target_vision_frames' : List[VisionFrame],
12
+ 'temp_vision_frame' : VisionFrame,
13
+ 'temp_vision_mask' : Mask
14
+ })
15
+
16
+ FaceEnhancerModel = Literal['codeformer', 'gfpgan_1.2', 'gfpgan_1.3', 'gfpgan_1.4', 'gpen_bfr_256', 'gpen_bfr_512', 'gpen_bfr_1024', 'gpen_bfr_2048', 'restoreformer_plus_plus']
17
+
18
+ FaceEnhancerWeight : TypeAlias = NDArray[Any]