Compare commits
59
Commits
main
...
tensorrt_testing
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a5e1130088 | ||
|
|
46675b2016 | ||
|
|
806263dd25 | ||
|
|
ac89dc1e2f | ||
|
|
27d745b53e | ||
|
|
052762578c | ||
|
|
e825c51c87 | ||
|
|
177b324fcd | ||
|
|
5c03bd8439 | ||
|
|
ef5ff7075f | ||
|
|
92fad03ee5 | ||
|
|
4cefac79b8 | ||
|
|
5e3c92d55c | ||
|
|
cc0501a2db | ||
|
|
3dc822fd2f | ||
|
|
697b9a78e6 | ||
|
|
2a7bd6116f | ||
|
|
a7d09f5d49 | ||
|
|
8e85d5b96d | ||
|
|
30989a9d37 | ||
|
|
eecf645603 | ||
|
|
1b080706df | ||
|
|
336f3f7c23 | ||
|
|
f27e1cca13 | ||
|
|
86e91a6e9d | ||
|
|
92529f7ca8 | ||
|
|
c0959056ae | ||
|
|
2e40fe3820 | ||
|
|
0a5e187637 | ||
|
|
4e19dbd6d1 | ||
|
|
9c190804a7 | ||
|
|
d9ca40e1d6 | ||
|
|
b68cf8788c | ||
|
|
e702b26895 | ||
|
|
c21705edb5 | ||
|
|
a284bb52b2 | ||
|
|
9884aac18a | ||
|
|
f8aada81db | ||
|
|
6735771664 | ||
|
|
857ddbc6d7 | ||
|
|
a6edcda97d | ||
|
|
ca01d706d0 | ||
|
|
0dc9a8a695 | ||
|
|
ba6b3f5f68 | ||
|
|
ee7d5b4241 | ||
|
|
03df9f35cd | ||
|
|
ec6b5c8c85 | ||
|
|
8509d9a551 | ||
|
|
e724da1161 | ||
|
|
68d0ddf72a | ||
|
|
6f9dba7777 | ||
|
|
811ca557fb | ||
|
|
6d790bdcc3 | ||
|
|
ef8b4263b4 | ||
|
|
eb5fddf4de | ||
|
|
9c7db3c59a | ||
|
|
bf3410cd0d | ||
|
|
24c65627db | ||
|
|
72bb6910e9 |
File diff suppressed because it is too large
Load Diff
+323
-288
@@ -1,52 +1,79 @@
|
||||
{
|
||||
"last_node_id": 31,
|
||||
"last_link_id": 68,
|
||||
"last_node_id": 203,
|
||||
"last_link_id": 477,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadImage",
|
||||
"id": 129,
|
||||
"type": "LivePortraitLoadCropper",
|
||||
"pos": [
|
||||
138,
|
||||
323
|
||||
-1050,
|
||||
-740
|
||||
],
|
||||
"size": {
|
||||
"0": 272.85791015625,
|
||||
"1": 331.60894775390625
|
||||
"0": 315,
|
||||
"1": 82
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"name": "cropper",
|
||||
"type": "LPCROPPER",
|
||||
"links": [
|
||||
59
|
||||
444
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
"Node name for S&R": "LivePortraitLoadCropper"
|
||||
},
|
||||
"widgets_values": [
|
||||
"oldman.jpg",
|
||||
"image"
|
||||
"CPU",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 19,
|
||||
"id": 1,
|
||||
"type": "DownloadAndLoadLivePortraitModels",
|
||||
"pos": [
|
||||
-1040,
|
||||
-850
|
||||
],
|
||||
"size": {
|
||||
"0": 302.43463134765625,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "live_portrait_pipe",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"links": [
|
||||
446,
|
||||
448
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
||||
},
|
||||
"widgets_values": [
|
||||
"fp16"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 165,
|
||||
"type": "ImageResizeKJ",
|
||||
"pos": [
|
||||
507,
|
||||
675
|
||||
-670,
|
||||
-560
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
@@ -59,12 +86,12 @@
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 30
|
||||
"link": 466
|
||||
},
|
||||
{
|
||||
"name": "get_image_size",
|
||||
"type": "IMAGE",
|
||||
"link": 68
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "width_input",
|
||||
@@ -88,7 +115,7 @@
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
32
|
||||
434
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
@@ -112,290 +139,162 @@
|
||||
"widgets_values": [
|
||||
512,
|
||||
512,
|
||||
"nearest-exact",
|
||||
false,
|
||||
"lanczos",
|
||||
true,
|
||||
2,
|
||||
0,
|
||||
0
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "ImageConcatMulti",
|
||||
"id": 196,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
860,
|
||||
679
|
||||
-1050,
|
||||
-550
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
466
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"oldman.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 78,
|
||||
"type": "GetImageSizeAndCount",
|
||||
"pos": [
|
||||
-310,
|
||||
-550
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 150
|
||||
"1": 86
|
||||
},
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image_1",
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 32
|
||||
},
|
||||
{
|
||||
"name": "image_2",
|
||||
"type": "IMAGE",
|
||||
"link": 67
|
||||
"link": 434
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
64
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"right",
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 23,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
1098,
|
||||
240
|
||||
],
|
||||
"size": [
|
||||
1253.234130859375,
|
||||
940.6170654296875
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 64
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "VHS_AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 24,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "LivePortrait",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "LivePortrait_00001.mp4",
|
||||
"subfolder": "",
|
||||
"type": "temp",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 24
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "DownloadAndLoadLivePortraitModels",
|
||||
"pos": [
|
||||
142,
|
||||
205
|
||||
],
|
||||
"size": {
|
||||
"0": 252,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "live_portrait_pipe",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"links": [
|
||||
58
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
||||
},
|
||||
"widgets_values": [
|
||||
"fp16"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VHS_LoadVideo",
|
||||
"pos": [
|
||||
161,
|
||||
714
|
||||
],
|
||||
"size": [
|
||||
235.1999969482422,
|
||||
491.1999969482422
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
30,
|
||||
60
|
||||
445,
|
||||
475
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"name": "512 width",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "VHS_AUDIO",
|
||||
"name": "512 height",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"name": "1 count",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_LoadVideo"
|
||||
},
|
||||
"widgets_values": {
|
||||
"video": "d3.mp4",
|
||||
"force_rate": 0,
|
||||
"force_size": "Disabled",
|
||||
"custom_width": 512,
|
||||
"custom_height": 512,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1,
|
||||
"choose video to upload": "image",
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"force_rate": 0,
|
||||
"filename": "d3.mp4",
|
||||
"type": "input",
|
||||
"format": "video/mp4",
|
||||
"select_every_nth": 1
|
||||
}
|
||||
}
|
||||
"Node name for S&R": "GetImageSizeAndCount"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 30,
|
||||
"id": 190,
|
||||
"type": "LivePortraitProcess",
|
||||
"pos": [
|
||||
500,
|
||||
249
|
||||
563,
|
||||
-418
|
||||
],
|
||||
"size": {
|
||||
"0": 367.79998779296875,
|
||||
"1": 362
|
||||
"0": 430.8000183105469,
|
||||
"1": 282
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "pipeline",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"link": 58
|
||||
"link": 448
|
||||
},
|
||||
{
|
||||
"name": "crop_info",
|
||||
"type": "CROPINFO",
|
||||
"link": 449
|
||||
},
|
||||
{
|
||||
"name": "source_image",
|
||||
"type": "IMAGE",
|
||||
"link": 59
|
||||
"link": 475
|
||||
},
|
||||
{
|
||||
"name": "driving_images",
|
||||
"type": "IMAGE",
|
||||
"link": 60
|
||||
"link": 477
|
||||
},
|
||||
{
|
||||
"name": "opt_retargeting_info",
|
||||
"type": "RETARGETINGINFO",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cropped_images",
|
||||
"name": "cropped_image",
|
||||
"type": "IMAGE",
|
||||
"links": [],
|
||||
"links": [
|
||||
470
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "full_images",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
67,
|
||||
68
|
||||
],
|
||||
"name": "output",
|
||||
"type": "LP_OUT",
|
||||
"links": [],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
@@ -403,85 +302,221 @@
|
||||
"properties": {
|
||||
"Node name for S&R": "LivePortraitProcess"
|
||||
},
|
||||
"widgets_values": [
|
||||
false,
|
||||
0.03,
|
||||
false,
|
||||
1,
|
||||
"constant",
|
||||
"single_frame",
|
||||
0.000003
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 198,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
1027,
|
||||
-409
|
||||
],
|
||||
"size": {
|
||||
"0": 521.2196044921875,
|
||||
"1": 566.1187133789062
|
||||
},
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 470
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 203,
|
||||
"type": "Screencap_mss",
|
||||
"pos": [
|
||||
5,
|
||||
-277
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 178
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
477
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "Screencap_mss"
|
||||
},
|
||||
"widgets_values": [
|
||||
0,
|
||||
0,
|
||||
512,
|
||||
512,
|
||||
1,
|
||||
0.1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 189,
|
||||
"type": "LivePortraitCropper",
|
||||
"pos": [
|
||||
-48,
|
||||
-851
|
||||
],
|
||||
"size": {
|
||||
"0": 330,
|
||||
"1": 242
|
||||
},
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "pipeline",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"link": 446,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "cropper",
|
||||
"type": "LPCROPPER",
|
||||
"link": 444
|
||||
},
|
||||
{
|
||||
"name": "source_image",
|
||||
"type": "IMAGE",
|
||||
"link": 445
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cropped_image",
|
||||
"type": "IMAGE",
|
||||
"links": null,
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "crop_info",
|
||||
"type": "CROPINFO",
|
||||
"links": [
|
||||
449
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LivePortraitCropper"
|
||||
},
|
||||
"widgets_values": [
|
||||
512,
|
||||
2.3,
|
||||
0,
|
||||
-0.11,
|
||||
true,
|
||||
false,
|
||||
1,
|
||||
false,
|
||||
1,
|
||||
true,
|
||||
true,
|
||||
"CPU"
|
||||
-0.125,
|
||||
0,
|
||||
"large-small",
|
||||
false
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
30,
|
||||
8,
|
||||
434,
|
||||
165,
|
||||
0,
|
||||
19,
|
||||
78,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
32,
|
||||
19,
|
||||
444,
|
||||
129,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
58,
|
||||
189,
|
||||
1,
|
||||
0,
|
||||
30,
|
||||
0,
|
||||
"LIVEPORTRAITPIPE"
|
||||
"LPCROPPER"
|
||||
],
|
||||
[
|
||||
59,
|
||||
4,
|
||||
445,
|
||||
78,
|
||||
0,
|
||||
30,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
60,
|
||||
8,
|
||||
0,
|
||||
30,
|
||||
189,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
64,
|
||||
18,
|
||||
446,
|
||||
1,
|
||||
0,
|
||||
23,
|
||||
189,
|
||||
0,
|
||||
"LIVEPORTRAITPIPE"
|
||||
],
|
||||
[
|
||||
448,
|
||||
1,
|
||||
0,
|
||||
190,
|
||||
0,
|
||||
"LIVEPORTRAITPIPE"
|
||||
],
|
||||
[
|
||||
449,
|
||||
189,
|
||||
1,
|
||||
190,
|
||||
1,
|
||||
"CROPINFO"
|
||||
],
|
||||
[
|
||||
466,
|
||||
196,
|
||||
0,
|
||||
165,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
67,
|
||||
30,
|
||||
1,
|
||||
18,
|
||||
1,
|
||||
470,
|
||||
190,
|
||||
0,
|
||||
198,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
68,
|
||||
30,
|
||||
1,
|
||||
19,
|
||||
1,
|
||||
475,
|
||||
78,
|
||||
0,
|
||||
190,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
477,
|
||||
203,
|
||||
0,
|
||||
190,
|
||||
3,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
@@ -489,10 +524,10 @@
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8264462809917354,
|
||||
"scale": 0.7513148009015781,
|
||||
"offset": {
|
||||
"0": 173.40487670898438,
|
||||
"1": -0.9636010527610779
|
||||
"0": 1170.2642381365986,
|
||||
"1": 992.3601372540302
|
||||
}
|
||||
}
|
||||
},
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,44 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
config for user
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
from dataclasses import dataclass
|
||||
#import tyro
|
||||
from typing_extensions import Annotated
|
||||
from .base_config import PrintableConfig, make_abs_path
|
||||
|
||||
|
||||
@dataclass(repr=False) # use repr from PrintableConfig
|
||||
class ArgumentConfig(PrintableConfig):
|
||||
########## input arguments ##########
|
||||
#source_image: Annotated[str, tyro.conf.arg(aliases=["-s"])] = make_abs_path('../../assets/examples/source/s6.jpg') # path to the reference portrait
|
||||
#driving_info: Annotated[str, tyro.conf.arg(aliases=["-d"])] = make_abs_path('../../assets/examples/driving/d0.mp4') # path to driving video or template (.pkl format)
|
||||
#output_dir: Annotated[str, tyro.conf.arg(aliases=["-o"])] = 'animations/' # directory to save output video
|
||||
#####################################
|
||||
|
||||
########## inference arguments ##########
|
||||
device_id: int = 0
|
||||
flag_lip_zero : bool = True # whether let the lip to close state before animation, only take effect when flag_eye_retargeting and flag_lip_retargeting is False
|
||||
flag_eye_retargeting: bool = False
|
||||
flag_lip_retargeting: bool = False
|
||||
flag_stitching: bool = True # we recommend setting it to True!
|
||||
flag_relative: bool = True # whether to use relative pose
|
||||
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
|
||||
flag_do_crop: bool = True # whether to crop the reference portrait to the face-cropping space
|
||||
flag_do_rot: bool = True # whether to conduct the rotation when flag_do_crop is True
|
||||
#########################################
|
||||
|
||||
########## crop arguments ##########
|
||||
dsize: int = 512
|
||||
scale: float = 2.3
|
||||
vx_ratio: float = 0 # vx ratio
|
||||
vy_ratio: float = -0.125 # vy ratio +up, -down
|
||||
####################################
|
||||
|
||||
########## gradio arguments ##########
|
||||
#server_port: Annotated[int, tyro.conf.arg(aliases=["-p"])] = 8890
|
||||
#share: bool = False
|
||||
#server_name: str = "0.0.0.0"
|
||||
@@ -1,18 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
parameters used for crop faces
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
from dataclasses import dataclass
|
||||
from typing import Union, List
|
||||
from .base_config import PrintableConfig
|
||||
|
||||
|
||||
@dataclass(repr=False) # use repr from PrintableConfig
|
||||
class CropConfig(PrintableConfig):
|
||||
dsize: int = 512 # crop size
|
||||
scale: float = 2.3 # scale factor
|
||||
vx_ratio: float = 0 # vx ratio
|
||||
vy_ratio: float = -0.125 # vy ratio +up, -down
|
||||
@@ -29,20 +29,13 @@ class InferenceConfig(PrintableConfig):
|
||||
flag_stitching: bool = True # we recommend setting it to True!
|
||||
|
||||
flag_relative: bool = True # whether to use relative pose
|
||||
anchor_frame: int = 0 # set this value if find_best_frame is True
|
||||
|
||||
input_shape: Tuple[int, int] = (256, 256) # input shape
|
||||
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
|
||||
output_fps: int = 30 # fps for output video
|
||||
crf: int = 15 # crf for output video
|
||||
|
||||
flag_write_result: bool = True # whether to write output video
|
||||
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
|
||||
mask_crop = None
|
||||
flag_write_gif: bool = False
|
||||
size_gif: int = 256
|
||||
ref_max_shape: int = 1280
|
||||
ref_shape_n: int = 2
|
||||
|
||||
device_id: int = 0
|
||||
flag_do_crop: bool = False # whether to crop the reference portrait to the face-cropping space
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .predictor import EfficientLivePortraitPredictor
|
||||
from .config.config import save_config_to_yaml
|
||||
from .utils import *
|
||||
@@ -0,0 +1 @@
|
||||
from .config import Config
|
||||
@@ -0,0 +1,29 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
pretty printing class
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
import os.path as osp
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
|
||||
class PrintableConfig: # pylint: disable=too-few-public-methods
|
||||
"""Printable Config defining str function"""
|
||||
|
||||
def __repr__(self):
|
||||
lines = [self.__class__.__name__ + ":"]
|
||||
for key, val in vars(self).items():
|
||||
if isinstance(val, Tuple):
|
||||
flattened_val = "["
|
||||
for item in val:
|
||||
flattened_val += str(item) + "\n"
|
||||
flattened_val = flattened_val.rstrip("\n")
|
||||
val = flattened_val + "]"
|
||||
lines += f"{key}: {str(val)}".split("\n")
|
||||
return "\n ".join(lines)
|
||||
@@ -0,0 +1,155 @@
|
||||
import os
|
||||
import requests
|
||||
from dataclasses import dataclass, asdict
|
||||
from typing import Literal, Tuple
|
||||
from tqdm import tqdm
|
||||
import torch.cuda
|
||||
import yaml
|
||||
|
||||
# Define the URLs for the model files
|
||||
MODEL_URLS = {
|
||||
'live_portrait': {
|
||||
'grid_sample_3d': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/libgrid_sample_3d_plugin.so?download=true',
|
||||
'F_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/appearance_feature_extractor.onnx?download=true',
|
||||
'M_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/motion_extractor.onnx?download=true',
|
||||
'GW_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/generator_fix_grid.onnx?download=true',
|
||||
'S_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching.onnx?download=true',
|
||||
'SE_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching_eye.onnx?download=true',
|
||||
'SL_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching_lip.onnx?download=true',
|
||||
# TensorRT FP32
|
||||
'F_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/appearance_feature_extractor_fp32.engine?download=true',
|
||||
'M_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/motion_extractor_fp32.engine?download=true',
|
||||
'GW_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/generator_fp32.engine?download=true',
|
||||
'S_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_fp32.engine?download=true',
|
||||
'SE_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_eye_fp32.engine?download=true',
|
||||
'SL_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_lip_fp32.engine?download=true',
|
||||
# TensorRT FP16
|
||||
'F_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/appearance_feature_extractor_fp16.engine?download=true',
|
||||
'M_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/motion_extractor_fp16.engine?download=true',
|
||||
'GW_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/generator_fp16.engine?download=true',
|
||||
'S_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_fp16.engine?download=true',
|
||||
'SE_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_eye_fp16.engine?download=true',
|
||||
'SL_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_lip_fp16.engine?download=true'
|
||||
},
|
||||
'insightface': {
|
||||
'arc_face': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/w600k_r50.onnx?download=true',
|
||||
'2d106det': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/2d106det.onnx?download=true',
|
||||
'det_10g': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/det_10g.onnx?download=true',
|
||||
'landmark': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/landmark.onnx?download=true'
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
# Function to download a file from a URL and save it locally
|
||||
def downloading(url, outf):
|
||||
if not os.path.exists(outf):
|
||||
print(f"Downloading checkpoint to {outf}")
|
||||
response = requests.get(url, stream=True)
|
||||
total_size_in_bytes = int(response.headers.get('content-length', 0))
|
||||
block_size = 1024 # 1 Kibibyte
|
||||
progress_bar = tqdm(total=total_size_in_bytes, unit='iB', unit_scale=True)
|
||||
with open(outf, 'wb') as file:
|
||||
for data in response.iter_content(block_size):
|
||||
progress_bar.update(len(data))
|
||||
file.write(data)
|
||||
progress_bar.close()
|
||||
if total_size_in_bytes != 0 and progress_bar.n != total_size_in_bytes:
|
||||
print("ERROR, something went wrong")
|
||||
print(f"Downloaded successfully to {outf}")
|
||||
else:
|
||||
return outf
|
||||
|
||||
|
||||
def get_efficient_live_portrait():
|
||||
# Download the models and save them in the current working directory
|
||||
current_dir = os.getcwd()
|
||||
face_dir = os.path.join(current_dir, 'live_portrait_weights')
|
||||
model_paths = {}
|
||||
for main_key, sub_dict in MODEL_URLS.items():
|
||||
dir_path = os.path.join(current_dir, 'live_portrait_weights', main_key)
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
model_paths[main_key] = {}
|
||||
for sub_key, url in sub_dict.items():
|
||||
filename = url.split('/')[-1].split('?')[0]
|
||||
save_path = os.path.join(dir_path, filename)
|
||||
downloading(url, save_path)
|
||||
model_paths[main_key][sub_key] = save_path
|
||||
print('Downloaded successfully and already saved')
|
||||
return model_paths, face_dir
|
||||
|
||||
|
||||
@dataclass(repr=False) # use repr from PrintableConfig
|
||||
class Config:
|
||||
model_paths, face_dir = get_efficient_live_portrait()
|
||||
grid_sample_3d: str = model_paths['live_portrait']['grid_sample_3d']
|
||||
# ONNX
|
||||
checkpoint_F: str = model_paths['live_portrait']['F_onnx'] # path to checkpoint
|
||||
checkpoint_M: str = model_paths['live_portrait']['M_onnx'] # path to checkpoint
|
||||
checkpoint_GW: str = model_paths['live_portrait']['GW_onnx']
|
||||
checkpoint_S: str = model_paths['live_portrait']['S_onnx'] # path to checkpoint
|
||||
checkpoint_SE: str = model_paths['live_portrait']['SE_onnx']
|
||||
checkpoint_SL: str = model_paths['live_portrait']['SL_onnx']
|
||||
|
||||
# TensorRT FP32
|
||||
F_rt: str = model_paths['live_portrait']['F_rt'] # path to checkpoint
|
||||
M_rt: str = model_paths['live_portrait']['M_rt'] # path to checkpoint
|
||||
GW_rt: str = model_paths['live_portrait']['GW_rt'] # path to checkpoint
|
||||
S_rt: str = model_paths['live_portrait']['S_rt'] # path to checkpoint
|
||||
SE_rt: str = model_paths['live_portrait']['SE_rt']
|
||||
SL_rt: str = model_paths['live_portrait']['SL_rt']
|
||||
|
||||
# TensorRT FP16
|
||||
F_rt_half: str = model_paths['live_portrait']['F_rt_half'] # path to checkpoint
|
||||
M_rt_half: str = model_paths['live_portrait']['M_rt_half'] # path to checkpoint
|
||||
GW_rt_half: str = model_paths['live_portrait']['GW_rt_half'] # path to checkpoint
|
||||
S_rt_half: str = model_paths['live_portrait']['S_rt_half'] # path to checkpoint
|
||||
SE_rt_half: str = model_paths['live_portrait']['SE_rt_half']
|
||||
SL_rt_half: str = model_paths['live_portrait']['SL_rt_half']
|
||||
|
||||
flag_use_half_precision: bool = True # whether to use half precision
|
||||
flag_lip_zero: bool = True # whether let the lip to close state before animation, only take effect when flag_eye_retargeting and flag_lip_retargeting is False
|
||||
lip_zero_threshold: float = 0.03
|
||||
flag_eye_retargeting: bool = False
|
||||
flag_lip_retargeting: bool = False
|
||||
flag_stitching: bool = True # we recommend setting it to True!
|
||||
flag_relative: bool = True # whether to use relative motion
|
||||
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
|
||||
flag_do_crop: bool = True # whether to crop the source portrait to the face-cropping space
|
||||
flag_do_rot: bool = True # whether to conduct the rotation when flag_do_crop is True
|
||||
flag_write_result: bool = True # whether to write output video
|
||||
flag_write_gif: bool = False
|
||||
|
||||
anchor_frame: int = 0 # set this value if find_best_frame is True
|
||||
|
||||
input_shape: Tuple[int, int] = (256, 256) # input shape
|
||||
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
|
||||
output_fps: int = 30 # fps for output video
|
||||
crf: int = 15 # crf for output video
|
||||
mask_crop: str = 'None'
|
||||
size_gif: int = 256
|
||||
ref_max_shape: int = 1280
|
||||
ref_shape_n: int = 2
|
||||
|
||||
device: str = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
|
||||
# crop config
|
||||
ckpt_landmark: str = model_paths['insightface']['landmark']
|
||||
ckpt_arc_face: str = model_paths['insightface']['arc_face']
|
||||
ckpt_landmark_106: str = model_paths['insightface']['2d106det']
|
||||
ckpt_det: str = model_paths['insightface']['det_10g']
|
||||
ckpt_face: str = face_dir
|
||||
dsize: int = 512 # crop size
|
||||
scale: float = 2.3 # scale factor
|
||||
vx_ratio: float = 0 # vx ratio
|
||||
vy_ratio: float = -0.125 # vy ratio +up, -down
|
||||
|
||||
|
||||
# Function to save the configuration to a YAML file
|
||||
def save_config_to_yaml(filename="efficient-live-portrait.yaml"):
|
||||
# Define the path where the YAML file will be saved
|
||||
file_path = os.path.join(os.getcwd(), filename)
|
||||
if not os.path.exists(file_path):
|
||||
# Save the configuration to the YAML file
|
||||
with open(file_path, 'w') as file:
|
||||
yaml.safe_dump(asdict(Config()), file)
|
||||
return file_path
|
||||
@@ -0,0 +1,47 @@
|
||||
from .utils.onnx_driver import ONNXEngine
|
||||
import numpy as np
|
||||
|
||||
|
||||
class EfficientLivePortraitPredictor:
|
||||
def __init__(self, use_tensorrt=False, half=False, **kwargs):
|
||||
super().__init__()
|
||||
self.use_tensorrt = use_tensorrt
|
||||
self.half = half
|
||||
self.cfg = kwargs
|
||||
if self.use_tensorrt:
|
||||
from .utils.tensorrt_driver import TensorRTEngine
|
||||
self.trt_engine = TensorRTEngine(self.half, **kwargs)
|
||||
else:
|
||||
self.onnx_engine = ONNXEngine().initialize_sessions(self.cfg)
|
||||
|
||||
def run_time(self, engine_name, task, inputs_onnx=None, inputs_tensorrt=None):
|
||||
"""
|
||||
Run inference using either TensorRT or ONNX Runtime based on the configuration.
|
||||
|
||||
Args:
|
||||
- engine_name (str): Name of the engine/model.
|
||||
- task (str): The task or model session name.
|
||||
- inputs_onnx (dict): Input dict for inference.
|
||||
- inputs_tensorrt(np.array or tensor): Input for inference TensorRT
|
||||
Returns:
|
||||
- The outputs from the inference.
|
||||
"""
|
||||
if self.use_tensorrt:
|
||||
return self.trt_engine.inference_tensorrt(engine_name, inputs_tensorrt)
|
||||
else:
|
||||
return self.inference_onnx(task, inputs_onnx)
|
||||
|
||||
def inference_onnx(self, task, inputs):
|
||||
"""
|
||||
Perform inference using ONNX Runtime.
|
||||
|
||||
Args:
|
||||
- task (str): The name of the task/model to use for inference.
|
||||
- inputs (list or array): A list or array of input tensors.
|
||||
|
||||
Returns:
|
||||
- List: The outputs of the inference.
|
||||
"""
|
||||
session = self.onnx_engine[task]
|
||||
outputs = session.run(None, inputs)
|
||||
return outputs
|
||||
@@ -0,0 +1 @@
|
||||
from .utils import *
|
||||
@@ -0,0 +1,51 @@
|
||||
import onnxruntime as ort
|
||||
import torch
|
||||
import numpy as np
|
||||
from typing import Dict
|
||||
|
||||
|
||||
class ONNXEngine:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def get_providers() -> list:
|
||||
"""Returns the list of providers based on the current device."""
|
||||
if ort.get_device() == 'GPU':
|
||||
return ['CUDAExecutionProvider']
|
||||
elif ort.get_device() == 'CPU':
|
||||
return ['CPUExecutionProvider', 'CoreMLExecutionProvider']
|
||||
else:
|
||||
return []
|
||||
|
||||
def initialize_sessions(self, cfg) -> Dict[str, ort.InferenceSession]:
|
||||
"""
|
||||
Initialize ONNX InferenceSession instances for each model checkpoint.
|
||||
|
||||
Args:
|
||||
- cfg (dict): Configuration dictionary containing checkpoint paths.
|
||||
|
||||
Returns:
|
||||
- Dict[str, ort.InferenceSession]: Dictionary mapping session names to InferenceSession objects.
|
||||
"""
|
||||
#providers = self.get_providers()
|
||||
providers = ['CUDAExecutionProvider']
|
||||
|
||||
# Initialize each session manually
|
||||
gw_session = ort.InferenceSession("live_portrait_weights\\live_portrait\\generator_fix_grid.onnx", providers=providers)
|
||||
# m_session = ort.InferenceSession(cfg.get("checkpoint_M"), providers=providers)
|
||||
# f_session = ort.InferenceSession(cfg.get("checkpoint_F"), providers=providers)
|
||||
# s_session = ort.InferenceSession(cfg.get("checkpoint_S"), providers=providers)
|
||||
# se_session = ort.InferenceSession(cfg.get("checkpoint_SE"), providers=providers)
|
||||
# sl_session = ort.InferenceSession(cfg.get("checkpoint_SL"), providers=providers)
|
||||
|
||||
# Return the sessions in a dictionary
|
||||
return {
|
||||
"gw_session": gw_session,
|
||||
# "m_session": m_session,
|
||||
# "f_session": f_session,
|
||||
# "s_session": s_session,
|
||||
# "se_session": se_session,
|
||||
# "sl_session": sl_session
|
||||
}
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
import tensorrt as trt
|
||||
import pycuda.driver as cuda
|
||||
import pycuda.gpuarray
|
||||
import pycuda.autoinit
|
||||
import numpy as np
|
||||
import ctypes
|
||||
from pathlib import Path
|
||||
|
||||
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
|
||||
|
||||
|
||||
class Binding:
|
||||
def __init__(self, engine, idx_or_name):
|
||||
self.name = idx_or_name if isinstance(idx_or_name, str) else engine.get_tensor_name(idx_or_name)
|
||||
if not self.name:
|
||||
raise IndexError(f"Binding index out of range: {idx_or_name}")
|
||||
self.is_input = engine.get_tensor_mode(self.name) == trt.TensorIOMode.INPUT
|
||||
dtype = engine.get_tensor_dtype(self.name)
|
||||
dtype_map = {
|
||||
trt.DataType.FLOAT: np.float32,
|
||||
trt.DataType.HALF: np.float16,
|
||||
trt.DataType.INT8: np.int8,
|
||||
trt.DataType.BOOL: np.bool_,
|
||||
}
|
||||
if hasattr(trt.DataType, 'INT32'):
|
||||
dtype_map[trt.DataType.INT32] = np.int32
|
||||
if hasattr(trt.DataType, 'INT64'):
|
||||
dtype_map[trt.DataType.INT64] = np.int64
|
||||
self.dtype = dtype_map[dtype]
|
||||
self.shape = tuple(engine.get_tensor_shape(self.name))
|
||||
self._host_buf = None
|
||||
self._device_buf = None
|
||||
|
||||
@property
|
||||
def host_buffer(self):
|
||||
if self._host_buf is None:
|
||||
self._host_buf = cuda.pagelocked_empty(self.shape, self.dtype)
|
||||
return self._host_buf
|
||||
|
||||
@property
|
||||
def device_buffer(self):
|
||||
if self._device_buf is None:
|
||||
self._device_buf = pycuda.gpuarray.empty(self.shape, self.dtype)
|
||||
return self._device_buf
|
||||
|
||||
def get_async(self, stream):
|
||||
self.device_buffer.get_async(stream, self.host_buffer)
|
||||
return self.host_buffer
|
||||
|
||||
def cleanup(self):
|
||||
if self._host_buf is not None:
|
||||
del self._host_buf
|
||||
if self._device_buf is not None:
|
||||
del self._device_buf
|
||||
|
||||
|
||||
class TensorRTEngine:
|
||||
def __init__(self, half, **kwargs):
|
||||
self.cfg = kwargs
|
||||
self.cfx = None
|
||||
if kwargs.get("cuda_ctx", None) is None:
|
||||
cuda.init()
|
||||
self.cfx = cuda.Device(0).make_context()
|
||||
else:
|
||||
self.cfx = kwargs.get("cuda_ctx")
|
||||
|
||||
if half:
|
||||
self.model_paths = {
|
||||
#'feature_extractor': self.cfg['F_rt_half'],
|
||||
#'motion_extractor': self.cfg['M_rt_half'],
|
||||
'generator': "live_portrait_weights\\live_portrait\\warping_spade-fix.engine",
|
||||
#'stitching_retargeting': self.cfg['S_rt_half'],
|
||||
#'stitching_retargeting_eye': self.cfg['SE_rt_half'],
|
||||
#'stitching_retargeting_lip': self.cfg['SL_rt_half']
|
||||
}
|
||||
else:
|
||||
self.model_paths = {
|
||||
'feature_extractor': self.cfg['F_rt'],
|
||||
'motion_extractor': self.cfg['M_rt'],
|
||||
'generator': self.cfg['GW_rt'],
|
||||
'stitching_retargeting': self.cfg['S_rt'],
|
||||
'stitching_retargeting_eye': self.cfg['SE_rt'],
|
||||
'stitching_retargeting_lip': self.cfg['SL_rt']
|
||||
}
|
||||
self.plugin_path = Path("N:\\AI\\ComfyUI\\live_portrait_weights\\live_portrait\\grid_sample_3d_plugin.dll")
|
||||
self.load_plugins(TRT_LOGGER)
|
||||
self.engines = {}
|
||||
self.contexts = {}
|
||||
self.bindings = {}
|
||||
self.binding_addresses = {}
|
||||
self.inputs = {}
|
||||
self.outputs = {}
|
||||
self.stream = cuda.Stream()
|
||||
self.initialize_engines()
|
||||
|
||||
def load_plugins(self, logger: trt.Logger):
|
||||
ctypes.CDLL(self.plugin_path, mode=ctypes.RTLD_GLOBAL)
|
||||
trt.init_libnvinfer_plugins(logger, "")
|
||||
|
||||
def initialize_engines(self):
|
||||
for model_name, model_path in self.model_paths.items():
|
||||
engine = self.load_engine(model_path)
|
||||
if engine is None:
|
||||
raise RuntimeError(f"Failed to load engine for {model_name}")
|
||||
context = engine.create_execution_context()
|
||||
if context is None:
|
||||
raise RuntimeError(f"Failed to create execution context for {model_name}")
|
||||
bindings = [Binding(engine, i) for i in range(engine.num_io_tensors)]
|
||||
self.engines[model_name] = engine
|
||||
self.contexts[model_name] = context
|
||||
self.bindings[model_name] = bindings
|
||||
self.binding_addresses[model_name] = [b.device_buffer.ptr for b in bindings]
|
||||
self.inputs[model_name] = [b for b in bindings if b.is_input]
|
||||
self.outputs[model_name] = [b for b in bindings if not b.is_input]
|
||||
self.prepare_buffers(model_name)
|
||||
|
||||
@staticmethod
|
||||
def load_engine(engine_file_path):
|
||||
with open(engine_file_path, "rb") as f, trt.Runtime(TRT_LOGGER) as runtime:
|
||||
return runtime.deserialize_cuda_engine(f.read())
|
||||
|
||||
def prepare_buffers(self, model_name):
|
||||
for binding in self.inputs[model_name] + self.outputs[model_name]:
|
||||
_ = binding.device_buffer # Force buffer allocation
|
||||
|
||||
@staticmethod
|
||||
def check_input_validity(input_idx, input_array, input_binding):
|
||||
if input_array.shape != input_binding.shape:
|
||||
if not (input_binding.shape == (1,) and input_array.shape == ()):
|
||||
raise ValueError(
|
||||
f"Wrong shape for input {input_idx}. Expected {input_binding.shape}, got {input_array.shape}.")
|
||||
if input_array.dtype != input_binding.dtype:
|
||||
if input_array.dtype == np.int64 and input_binding.dtype == np.int32:
|
||||
input_array = input_array.astype(np.int32)
|
||||
if not np.array_equal(input_array, input_array.astype(np.int64)):
|
||||
raise TypeError(
|
||||
f"Wrong dtype for input {input_idx}. Expected {input_binding.dtype}, got {input_array.dtype}. Cannot safely cast.")
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Wrong dtype for input {input_idx}. Expected {input_binding.dtype}, got {input_array.dtype}.")
|
||||
return input_array
|
||||
|
||||
def run_sequential_tasks(self, model_name, inputs):
|
||||
if model_name not in self.engines:
|
||||
raise ValueError(f"Model name {model_name} not found in engines.")
|
||||
engine = self.engines[model_name]
|
||||
context = self.contexts[model_name]
|
||||
binding_addresses = self.binding_addresses[model_name]
|
||||
inputs_bindings = self.inputs[model_name]
|
||||
outputs_bindings = self.outputs[model_name]
|
||||
|
||||
if isinstance(inputs, dict):
|
||||
inputs = [inputs[b.name] for b in inputs_bindings]
|
||||
if len(inputs) != len(inputs_bindings):
|
||||
raise ValueError(f"Number of input arrays does not match number of input bindings for model {model_name}.")
|
||||
|
||||
self.cfx.push() # Push CUDA context
|
||||
|
||||
try:
|
||||
for i, (input_array, input_binding) in enumerate(zip(inputs, inputs_bindings)):
|
||||
input_array = self.check_input_validity(i, input_array, input_binding)
|
||||
input_array = np.ascontiguousarray(input_array) # Ensure the input array is contiguous
|
||||
cuda.memcpy_htod(input_binding.device_buffer.ptr, input_array)
|
||||
|
||||
for i in range(engine.num_io_tensors):
|
||||
tensor_name = engine.get_tensor_name(i)
|
||||
if i < len(inputs) and engine.is_shape_inference_io(tensor_name):
|
||||
context.set_tensor_address(tensor_name, inputs[i].ctypes.data)
|
||||
else:
|
||||
context.set_tensor_address(tensor_name, binding_addresses[i])
|
||||
|
||||
context.execute_async_v3(self.stream.handle)
|
||||
self.stream.synchronize()
|
||||
|
||||
outputs = []
|
||||
for output in outputs_bindings:
|
||||
host_output = np.empty(output.shape, dtype=output.dtype)
|
||||
cuda.memcpy_dtoh(host_output, output.device_buffer.ptr)
|
||||
outputs.append(host_output)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during inference for model {model_name}: {e}")
|
||||
outputs = None
|
||||
|
||||
self.cfx.pop() # Pop CUDA context
|
||||
|
||||
return outputs
|
||||
|
||||
def inference_tensorrt(self, task, inputs):
|
||||
if not isinstance(inputs, list):
|
||||
raise TypeError("Inputs should be a list of numpy arrays or tensors.")
|
||||
|
||||
if task not in self.inputs:
|
||||
raise ValueError(f"Task {task} not found in the model inputs.")
|
||||
|
||||
# Ensure all inputs are on the same memory type
|
||||
if isinstance(inputs[0], pycuda.gpuarray.GPUArray):
|
||||
# Ensure all inputs are on GPU
|
||||
inputs = [cuda.to_gpu(input_array) if not isinstance(input_array, pycuda.gpuarray.GPUArray) else input_array
|
||||
for input_array in inputs]
|
||||
else:
|
||||
# Ensure all inputs are on CPU
|
||||
inputs = [input_array.get() if isinstance(input_array, pycuda.gpuarray.GPUArray) else input_array
|
||||
for input_array in inputs]
|
||||
|
||||
inputs = [self.check_input_validity(i, np.array(input_array), self.inputs[task][i])
|
||||
for i, input_array in enumerate(inputs)]
|
||||
|
||||
result = self.run_sequential_tasks(task, inputs)
|
||||
return result
|
||||
|
||||
def __del__(self):
|
||||
del self.engines
|
||||
del self.contexts
|
||||
del self.bindings
|
||||
del self.binding_addresses
|
||||
del self.inputs
|
||||
del self.outputs
|
||||
del self.stream
|
||||
try:
|
||||
if self.cfx is not None:
|
||||
self.cfx.pop()
|
||||
del self.cfx
|
||||
except Exception as e:
|
||||
print(f"Error during cleanup: {e}")
|
||||
|
||||
# Example usage
|
||||
# engine = TensorRTEngine(half=True, F_rt_half="path/to/F_rt_half", M_rt_half="path/to/M_rt_half",
|
||||
# GW_rt_half="path/to/GW_rt_half", S_rt_half="path/to/S_rt_half",
|
||||
# SE_rt_half="path/to/SE_rt_half", SL_rt_half="path/to/SL_rt_half",
|
||||
# grid_sample_3d="path/to/grid_sample_3d.so")
|
||||
@@ -0,0 +1,202 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
utility functions and classes to handle feature extraction and model loading
|
||||
"""
|
||||
import torch
|
||||
import os
|
||||
from glob import glob
|
||||
import os.path as osp
|
||||
import imageio
|
||||
import numpy as np
|
||||
import cv2
|
||||
from rich.progress import track
|
||||
|
||||
cv2.setNumThreads(0)
|
||||
cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
|
||||
def suffix(filename):
|
||||
"""a.jpg -> jpg"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return ""
|
||||
return filename[pos + 1:]
|
||||
|
||||
|
||||
def prefix(filename):
|
||||
"""a.jpg -> a"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return filename
|
||||
return filename[:pos]
|
||||
|
||||
|
||||
def basename(filename):
|
||||
"""a/b/c.jpg -> c"""
|
||||
return prefix(osp.basename(filename))
|
||||
|
||||
|
||||
def is_video(file_path):
|
||||
if file_path.lower().endswith((".mp4", ".mov", ".avi", ".webm")) or osp.isdir(file_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def is_template(file_path):
|
||||
if file_path.endswith(".pkl"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def mkdir(d, log=False):
|
||||
# return self-assined `d`, for one line code
|
||||
if not osp.exists(d):
|
||||
os.makedirs(d, exist_ok=True)
|
||||
if log:
|
||||
print(f"Make dir: {d}")
|
||||
return d
|
||||
|
||||
|
||||
def squeeze_tensor_to_numpy(tensor):
|
||||
out = tensor.data.squeeze(0).cpu().numpy()
|
||||
return out
|
||||
|
||||
|
||||
def dct2cuda(dct: dict, device: str):
|
||||
for key in dct:
|
||||
dct[key] = torch.tensor(dct[key]).to(device)
|
||||
return dct
|
||||
|
||||
|
||||
def concat_feat(kp_source: torch.Tensor, kp_driving: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
kp_source: (bs, k, 3)
|
||||
kp_driving: (bs, k, 3)
|
||||
Return: (bs, 2k*3)
|
||||
"""
|
||||
bs_src = kp_source.shape[0]
|
||||
bs_dri = kp_driving.shape[0]
|
||||
assert bs_src == bs_dri, 'batch size must be equal'
|
||||
|
||||
feat = torch.cat([kp_source.view(bs_src, -1), kp_driving.view(bs_dri, -1)], dim=1)
|
||||
return feat
|
||||
|
||||
|
||||
def load_image_rgb(image_path: str):
|
||||
if not osp.exists(image_path):
|
||||
raise FileNotFoundError(f"Image not found: {image_path}")
|
||||
img = cv2.imread(image_path, cv2.IMREAD_COLOR)
|
||||
return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
|
||||
|
||||
def load_driving_info(driving_info):
|
||||
driving_video_ori = []
|
||||
|
||||
def load_images_from_directory(directory):
|
||||
image_paths = sorted(glob(osp.join(directory, '*.png')) + glob(osp.join(directory, '*.jpg')))
|
||||
return [load_image_rgb(im_path) for im_path in image_paths]
|
||||
|
||||
def load_images_from_video(file_path):
|
||||
reader = imageio.get_reader(file_path)
|
||||
return [image for idx, image in enumerate(reader)]
|
||||
|
||||
if osp.isdir(driving_info):
|
||||
driving_video_ori = load_images_from_directory(driving_info)
|
||||
elif osp.isfile(driving_info):
|
||||
driving_video_ori = load_images_from_video(driving_info)
|
||||
|
||||
return driving_video_ori
|
||||
|
||||
|
||||
def contiguous(obj):
|
||||
if not obj.flags.c_contiguous:
|
||||
obj = obj.copy(order="C")
|
||||
return obj
|
||||
|
||||
|
||||
def resize_to_limit(img: np.ndarray, max_dim=1920, n=2):
|
||||
"""
|
||||
ajust the size of the image so that the maximum dimension does not exceed max_dim, and the width and the height of the image are multiples of n.
|
||||
:param img: the image to be processed.
|
||||
:param max_dim: the maximum dimension constraint.
|
||||
:param n: the number that needs to be multiples of.
|
||||
:return: the adjusted image.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
|
||||
# ajust the size of the image according to the maximum dimension
|
||||
if max_dim > 0 and max(h, w) > max_dim:
|
||||
if h > w:
|
||||
new_h = max_dim
|
||||
new_w = int(w * (max_dim / h))
|
||||
else:
|
||||
new_w = max_dim
|
||||
new_h = int(h * (max_dim / w))
|
||||
img = cv2.resize(img, (new_w, new_h))
|
||||
|
||||
# ensure that the image dimensions are multiples of n
|
||||
n = max(n, 1)
|
||||
new_h = img.shape[0] - (img.shape[0] % n)
|
||||
new_w = img.shape[1] - (img.shape[1] % n)
|
||||
|
||||
if new_h == 0 or new_w == 0:
|
||||
# when the width or height is less than n, no need to process
|
||||
return img
|
||||
|
||||
if new_h != img.shape[0] or new_w != img.shape[1]:
|
||||
img = img[:new_h, :new_w]
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def load_img_online(obj, mode="bgr", **kwargs):
|
||||
max_dim = kwargs.get("max_dim", 1920)
|
||||
n = kwargs.get("n", 2)
|
||||
if isinstance(obj, str):
|
||||
if mode.lower() == "gray":
|
||||
img = cv2.imread(obj, cv2.IMREAD_GRAYSCALE)
|
||||
else:
|
||||
img = cv2.imread(obj, cv2.IMREAD_COLOR)
|
||||
else:
|
||||
img = obj
|
||||
|
||||
# Resize image to satisfy constraints
|
||||
img = resize_to_limit(img, max_dim=max_dim, n=n)
|
||||
|
||||
if mode.lower() == "bgr":
|
||||
return contiguous(img)
|
||||
elif mode.lower() == "rgb":
|
||||
return contiguous(img[..., ::-1])
|
||||
else:
|
||||
raise Exception(f"Unknown mode {mode}")
|
||||
|
||||
|
||||
def images2video(images, wfp, **kwargs):
|
||||
fps = kwargs.get('fps', 30)
|
||||
video_format = kwargs.get('format', 'mp4') # default is mp4 format
|
||||
codec = kwargs.get('codec', 'libx264') # default is libx264 encoding
|
||||
quality = kwargs.get('quality') # video quality
|
||||
pixelformat = kwargs.get('pixelformat', 'yuv420p') # video pixel format
|
||||
image_mode = kwargs.get('image_mode', 'rgb')
|
||||
macro_block_size = kwargs.get('macro_block_size', 2)
|
||||
ffmpeg_params = ['-crf', str(kwargs.get('crf', 18))]
|
||||
|
||||
writer = imageio.get_writer(
|
||||
wfp, fps=fps, format=video_format,
|
||||
codec=codec, quality=quality, ffmpeg_params=ffmpeg_params, pixelformat=pixelformat,
|
||||
macro_block_size=macro_block_size
|
||||
)
|
||||
|
||||
n = len(images)
|
||||
for i in track(range(n), description='writing', transient=True):
|
||||
if image_mode.lower() == 'bgr':
|
||||
writer.append_data(images[i][..., ::-1])
|
||||
else:
|
||||
writer.append_data(images[i])
|
||||
|
||||
writer.close()
|
||||
|
||||
# print(f':smiley: Dump to {wfp}\n', style="bold green")
|
||||
print(f'Dump to {wfp}\n')
|
||||
return wfp
|
||||
@@ -4,184 +4,287 @@
|
||||
Pipeline of LivePortrait
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import os.path as osp
|
||||
from rich.progress import track
|
||||
|
||||
from .config.inference_config import InferenceConfig
|
||||
|
||||
#from .utils.cropper import Cropper
|
||||
from .utils.camera import get_rotation_matrix
|
||||
#from .utils.video import images2video, concat_frames
|
||||
from .utils.crop import _transform_img
|
||||
#from .utils.retargeting_utils import calc_lip_close_ratio
|
||||
#from .utils.io import load_image_rgb, load_driving_info
|
||||
#from .utils.helper import mkdir, basename, dct2cuda, is_video, is_template, resize_to_limit
|
||||
from .utils.helper import resize_to_limit
|
||||
#from .utils.rprint import rlog as log
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
|
||||
import comfy.utils
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
from .config.inference_config import InferenceConfig
|
||||
from .utils.camera import get_rotation_matrix
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
||||
from .utils.filter import smooth
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
import os
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
class LivePortraitPipeline(object):
|
||||
|
||||
def __init__(self, appearance_feature_extractor, motion_extractor, warping_module,
|
||||
spade_generator, stitching_retargeting_module, inference_cfg: InferenceConfig):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
appearance_feature_extractor,
|
||||
motion_extractor,
|
||||
warping_module,
|
||||
spade_generator,
|
||||
stitching_retargeting_module,
|
||||
inference_cfg: InferenceConfig,
|
||||
):
|
||||
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(
|
||||
appearance_feature_extractor, motion_extractor, warping_module,
|
||||
spade_generator, stitching_retargeting_module, cfg=inference_cfg)
|
||||
appearance_feature_extractor,
|
||||
motion_extractor,
|
||||
warping_module,
|
||||
spade_generator,
|
||||
stitching_retargeting_module,
|
||||
cfg=inference_cfg,
|
||||
)
|
||||
|
||||
def execute(self, img_rgb, driving_images_np):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg # for convenience
|
||||
######## process reference portrait ########
|
||||
#img_rgb = load_image_rgb(args.source_image)
|
||||
img_rgb = resize_to_limit(img_rgb, inference_cfg.ref_max_shape, inference_cfg.ref_shape_n)
|
||||
#log(f"Load source image from {args.source_image}")
|
||||
crop_info = self.cropper.crop_single_image(img_rgb)
|
||||
source_lmk = crop_info['lmk_crop']
|
||||
_, img_crop_256x256 = crop_info['img_crop'], crop_info['img_crop_256x256']
|
||||
if inference_cfg.flag_do_crop:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
|
||||
else:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_rgb)
|
||||
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
|
||||
x_c_s = x_s_info['kp']
|
||||
R_s = get_rotation_matrix(x_s_info['pitch'], x_s_info['yaw'], x_s_info['roll'])
|
||||
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
|
||||
x_s = self.live_portrait_wrapper.transform_keypoint(x_s_info)
|
||||
def _get_source_frame(self, source_np, idx, method):
|
||||
if source_np.shape[0] == 1:
|
||||
return source_np[0]
|
||||
|
||||
if inference_cfg.flag_lip_zero:
|
||||
# let lip-open scalar to be 0 at first
|
||||
c_d_lip_before_animation = [0.]
|
||||
combined_lip_ratio_tensor_before_animation = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk)
|
||||
if combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold:
|
||||
inference_cfg.flag_lip_zero = False
|
||||
else:
|
||||
lip_delta_before_animation = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation)
|
||||
############################################
|
||||
if method == "constant":
|
||||
return source_np[min(idx, source_np.shape[0] - 1)]
|
||||
elif method == "cycle":
|
||||
return source_np[idx % source_np.shape[0]]
|
||||
elif method == "mirror":
|
||||
cycle_length = 2 * source_np.shape[0] - 2
|
||||
mirror_idx = idx % cycle_length
|
||||
if mirror_idx >= source_np.shape[0]:
|
||||
mirror_idx = cycle_length - mirror_idx
|
||||
return source_np[mirror_idx]
|
||||
|
||||
######## process driving info ########
|
||||
#if is_video(args.driving_info):
|
||||
#log(f"Load from video file (mp4 mov avi etc...): {args.driving_info}")
|
||||
# TODO: 这里track一下驱动视频 -> 构建模板
|
||||
#driving_rgb_lst = load_driving_info(args.driving_info)
|
||||
def execute(
|
||||
self, source_np, driving_images, crop_info, driving_landmarks, delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, mismatch_method="constant",
|
||||
):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg
|
||||
device = inference_cfg.device_id
|
||||
|
||||
driving_rgb_lst = driving_images_np
|
||||
|
||||
driving_rgb_lst_256 = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
|
||||
I_d_lst = self.live_portrait_wrapper.prepare_driving_videos(driving_rgb_lst_256)
|
||||
n_frames = I_d_lst.shape[0]
|
||||
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
|
||||
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_rgb_lst)
|
||||
input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
|
||||
|
||||
# elif is_template(args.driving_info):
|
||||
# log(f"Load from video templates {args.driving_info}")
|
||||
# with open(args.driving_info, 'rb') as f:
|
||||
# template_lst, driving_lmk_lst = pickle.load(f)
|
||||
# n_frames = template_lst[0]['n_frames']
|
||||
# input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
|
||||
# else:
|
||||
# raise Exception("Unsupported driving types!")
|
||||
#########################################
|
||||
|
||||
######## prepare for pasteback ########
|
||||
if inference_cfg.flag_pasteback:
|
||||
if inference_cfg.mask_crop is None:
|
||||
inference_cfg.mask_crop = cv2.imread(make_abs_path('./utils/resources/mask_template.png'), cv2.IMREAD_COLOR)
|
||||
mask_ori = _transform_img(inference_cfg.mask_crop, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
|
||||
mask_ori = mask_ori.astype(np.float32) / 255.
|
||||
I_p_paste_lst = []
|
||||
#########################################
|
||||
|
||||
I_p_lst = []
|
||||
out_list = []
|
||||
R_d_0, x_d_0_info = None, None
|
||||
pbar = comfy.utils.ProgressBar(n_frames)
|
||||
for i in track(range(n_frames), description='Animating...', total=n_frames):
|
||||
#if is_video(args.driving_info):
|
||||
# extract kp info by M
|
||||
I_d_i = I_d_lst[i]
|
||||
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
|
||||
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
|
||||
# else:
|
||||
# # from template
|
||||
# x_d_i_info = template_lst[i]
|
||||
# x_d_i_info = dct2cuda(x_d_i_info, inference_cfg.device_id)
|
||||
# R_d_i = x_d_i_info['R_d']
|
||||
|
||||
if mismatch_method == "cut" or relative_motion_mode == "source_video_smoothed":
|
||||
total_frames = source_np.shape[0]
|
||||
else:
|
||||
total_frames = driving_images.shape[0]
|
||||
|
||||
|
||||
disable_progress_bar = True if relative_motion_mode == "single_frame" else False
|
||||
|
||||
source_info = crop_info["source_info"]
|
||||
source_rot_list = crop_info["source_rot_list"]
|
||||
f_s_list = crop_info["f_s_list"]
|
||||
x_s_list = crop_info["x_s_list"]
|
||||
|
||||
driving_info = []
|
||||
driving_exp_list = []
|
||||
driving_rot_list = []
|
||||
|
||||
for i in tqdm(range(driving_images.shape[0]), desc='Processing driving images...', total=driving_images.shape[0], disable=disable_progress_bar):
|
||||
#get driving keypoints info
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
if crop_info["crop_info_list"][safe_index] is None:
|
||||
driving_info.append(None)
|
||||
driving_rot_list.append(None)
|
||||
driving_exp_list.append(None)
|
||||
continue
|
||||
x_d_info = self.live_portrait_wrapper.get_kp_info(driving_images[i].unsqueeze(0).to(device))
|
||||
|
||||
if i == 0:
|
||||
R_d_0 = R_d_i
|
||||
x_d_0_info = x_d_i_info
|
||||
first = x_d_info
|
||||
|
||||
if inference_cfg.flag_relative:
|
||||
R_new = (R_d_i @ R_d_0.permute(0, 2, 1)) @ R_s
|
||||
delta_new = x_s_info['exp'] + (x_d_i_info['exp'] - x_d_0_info['exp'])
|
||||
scale_new = x_s_info['scale'] * (x_d_i_info['scale'] / x_d_0_info['scale'])
|
||||
t_new = x_s_info['t'] + (x_d_i_info['t'] - x_d_0_info['t'])
|
||||
driving_info.append(x_d_info)
|
||||
|
||||
driving_exp = source_info[safe_index]["exp"] + x_d_info["exp"] - first["exp"]
|
||||
driving_exp_list.append(driving_exp.cpu())
|
||||
|
||||
R_d = get_rotation_matrix(
|
||||
x_d_info["pitch"], x_d_info["yaw"], x_d_info["roll"]
|
||||
)
|
||||
driving_rot_list.append(R_d)
|
||||
|
||||
if relative_motion_mode == "source_video_smoothed":
|
||||
x_d_r_lst = []
|
||||
first_driving_rot = driving_rot_list[0].cpu().numpy().astype(np.float32).transpose(0, 2, 1)
|
||||
for i in tqdm(range(source_np.shape[0]), desc='Smoothing...', total=source_np.shape[0]):
|
||||
if driving_rot_list[i] is None:
|
||||
x_d_r_lst.append(None)
|
||||
continue
|
||||
driving_rot = driving_rot_list[i].cpu().numpy().astype(np.float32)
|
||||
source_rot = source_rot_list[i].cpu().numpy().astype(np.float32)
|
||||
dot = np.dot(driving_rot, first_driving_rot) @ source_rot
|
||||
x_d_r_lst.append(dot)
|
||||
|
||||
driving_exp_list_smooth = smooth(driving_exp_list, source_info[0]["exp"].shape, device, observation_variance=driving_smooth_observation_variance)
|
||||
driving_rot_list_smooth = smooth(x_d_r_lst, source_rot_list[0].shape, device, observation_variance=driving_smooth_observation_variance)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(total_frames)
|
||||
|
||||
for i in tqdm(range(total_frames), desc='Animating...', total=total_frames, disable=disable_progress_bar):
|
||||
|
||||
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
|
||||
# skip and return empty frames if no crop due to no face detected
|
||||
if crop_info["crop_info_list"][safe_index] is None:
|
||||
out_list.append({})
|
||||
pbar.update(1)
|
||||
continue
|
||||
|
||||
source_lmk = crop_info["crop_info_list"][safe_index]["lmk_crop"]
|
||||
|
||||
x_d_info = driving_info[i]
|
||||
R_d = driving_rot_list[i]
|
||||
|
||||
x_s_info = source_info[safe_index]
|
||||
R_s = source_rot_list[safe_index]
|
||||
f_s = f_s_list[safe_index]
|
||||
x_s = x_s_list[safe_index]
|
||||
|
||||
x_c_s = x_s_info["kp"]
|
||||
|
||||
#lip zero
|
||||
if inference_cfg.flag_lip_zero:
|
||||
c_d_lip_before_animation = [0.0]
|
||||
combined_lip_ratio_tensor_before_animation = (self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk))
|
||||
|
||||
if (combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold):
|
||||
inference_cfg.flag_lip_zero = False
|
||||
else:
|
||||
lip_delta_before_animation = (self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation))
|
||||
|
||||
if relative_motion_mode == "relative":
|
||||
if i == 0:
|
||||
R_d_0 = R_d
|
||||
x_d_0_info = x_d_info
|
||||
R_new = (R_d @ R_d_0.permute(0, 2, 1)) @ R_s
|
||||
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_d_0_info["exp"])
|
||||
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_d_0_info["scale"])
|
||||
t_new = x_s_info["t"] + (x_d_info["t"] - x_d_0_info["t"])
|
||||
elif relative_motion_mode == "source_video_smoothed":
|
||||
R_new = driving_rot_list_smooth[i]
|
||||
delta_new = driving_exp_list_smooth[i]
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
elif relative_motion_mode == "relative_rotation_only":
|
||||
R_new = R_s
|
||||
delta_new = x_s_info['exp']
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
elif relative_motion_mode == "single_frame":
|
||||
R_new = R_d
|
||||
delta_new = x_d_info['exp']
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
else:
|
||||
R_new = R_d_i
|
||||
delta_new = x_d_i_info['exp']
|
||||
scale_new = x_s_info['scale']
|
||||
t_new = x_d_i_info['t']
|
||||
R_new = R_d
|
||||
delta_new = x_s_info['exp']
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
|
||||
t_new[..., 2].fill_(0) # zero tz
|
||||
t_new[..., 2].fill_(0) # zero tz
|
||||
|
||||
delta_new = delta_new * delta_multiplier
|
||||
|
||||
x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
|
||||
|
||||
# Algorithm 1:
|
||||
if not inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
|
||||
if (
|
||||
not inference_cfg.flag_stitching
|
||||
and not inference_cfg.flag_eye_retargeting
|
||||
and not inference_cfg.flag_lip_retargeting
|
||||
):
|
||||
# without stitching or retargeting
|
||||
if inference_cfg.flag_lip_zero:
|
||||
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
else:
|
||||
pass
|
||||
elif inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
|
||||
elif (
|
||||
inference_cfg.flag_stitching
|
||||
and not inference_cfg.flag_eye_retargeting
|
||||
and not inference_cfg.flag_lip_retargeting
|
||||
):
|
||||
# with stitching and without retargeting
|
||||
if inference_cfg.flag_lip_zero:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(
|
||||
x_s, x_d_i_new
|
||||
) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
else:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
else:
|
||||
eyes_delta, lip_delta = None, None
|
||||
if inference_cfg.flag_eye_retargeting:
|
||||
c_d_eyes_i = input_eye_ratio_lst[i]
|
||||
combined_eye_ratio_tensor = self.live_portrait_wrapper.calc_combined_eye_ratio(c_d_eyes_i, source_lmk)
|
||||
combined_eye_ratio_tensor = combined_eye_ratio_tensor * inference_cfg.eyes_retargeting_multiplier
|
||||
c_d_eyes_i = calc_eye_close_ratio(driving_landmarks[i][None])
|
||||
combined_eye_ratio_tensor = (
|
||||
self.live_portrait_wrapper.calc_combined_eye_ratio(
|
||||
c_d_eyes_i, source_lmk
|
||||
)
|
||||
)
|
||||
combined_eye_ratio_tensor = (
|
||||
combined_eye_ratio_tensor
|
||||
* inference_cfg.eyes_retargeting_multiplier
|
||||
)
|
||||
# ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
|
||||
eyes_delta = self.live_portrait_wrapper.retarget_eye(x_s, combined_eye_ratio_tensor)
|
||||
eyes_delta = self.live_portrait_wrapper.retarget_eye(
|
||||
x_s, combined_eye_ratio_tensor
|
||||
)
|
||||
if inference_cfg.flag_lip_retargeting:
|
||||
c_d_lip_i = input_lip_ratio_lst[i]
|
||||
combined_lip_ratio_tensor = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_i, source_lmk)
|
||||
combined_lip_ratio_tensor = combined_lip_ratio_tensor * inference_cfg.lip_retargeting_multiplier
|
||||
c_d_lip_i = calc_lip_close_ratio(driving_landmarks[i][None])
|
||||
combined_lip_ratio_tensor = (
|
||||
self.live_portrait_wrapper.calc_combined_lip_ratio(
|
||||
c_d_lip_i, source_lmk
|
||||
)
|
||||
)
|
||||
combined_lip_ratio_tensor = (
|
||||
combined_lip_ratio_tensor
|
||||
* inference_cfg.lip_retargeting_multiplier
|
||||
)
|
||||
# ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
|
||||
lip_delta = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor)
|
||||
lip_delta = self.live_portrait_wrapper.retarget_lip(
|
||||
x_s, combined_lip_ratio_tensor
|
||||
)
|
||||
|
||||
if inference_cfg.flag_relative: # use x_s
|
||||
x_d_i_new = x_s + \
|
||||
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
|
||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
||||
x_d_i_new = (
|
||||
x_s
|
||||
+ (
|
||||
eyes_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if eyes_delta is not None
|
||||
else 0
|
||||
)
|
||||
+ (
|
||||
lip_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if lip_delta is not None
|
||||
else 0
|
||||
)
|
||||
)
|
||||
else: # use x_d,i
|
||||
x_d_i_new = x_d_i_new + \
|
||||
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
|
||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
||||
x_d_i_new = (
|
||||
x_d_i_new
|
||||
+ (
|
||||
eyes_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if eyes_delta is not None
|
||||
else 0
|
||||
)
|
||||
+ (
|
||||
lip_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if lip_delta is not None
|
||||
else 0
|
||||
)
|
||||
)
|
||||
|
||||
if inference_cfg.flag_stitching:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
|
||||
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
|
||||
I_p_i = self.live_portrait_wrapper.parse_output(out['out'])[0]
|
||||
I_p_lst.append(I_p_i)
|
||||
if inference_cfg.flag_stitching:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
|
||||
out = self.live_portrait_wrapper.warp_decode_tensorrt(f_s, x_s, x_d_i_new)
|
||||
#out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
|
||||
|
||||
out_list.append(out)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
#if inference_cfg.flag_pasteback:
|
||||
I_p_i_to_ori = _transform_img(I_p_i, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
|
||||
I_p_i_to_ori_blend = np.clip(mask_ori * I_p_i_to_ori + (1 - mask_ori) * img_rgb, 0, 255).astype(np.uint8)
|
||||
out = np.hstack([I_p_i_to_ori, I_p_i_to_ori_blend])
|
||||
I_p_paste_lst.append(I_p_i_to_ori_blend)
|
||||
|
||||
return I_p_lst, I_p_paste_lst
|
||||
out_dict = {
|
||||
"out_list": out_list,
|
||||
"crop_info": crop_info,
|
||||
"mismatch_method": mismatch_method,
|
||||
}
|
||||
|
||||
return out_dict
|
||||
|
||||
@@ -13,6 +13,9 @@ from .utils.retargeting_utils import compute_eye_delta, compute_lip_delta
|
||||
from .utils.camera import headpose_pred_to_degree, get_rotation_matrix
|
||||
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
||||
from .config.inference_config import InferenceConfig
|
||||
from contextlib import nullcontext
|
||||
|
||||
from .efficient import EfficientLivePortraitPredictor
|
||||
|
||||
from comfy.model_management import get_autocast_device
|
||||
|
||||
@@ -31,10 +34,7 @@ class LivePortraitWrapper(object):
|
||||
self.device_id = cfg.device_id
|
||||
self.timer = Timer()
|
||||
|
||||
def update_config(self, user_args):
|
||||
for k, v in user_args.items():
|
||||
if hasattr(self.cfg, k):
|
||||
setattr(self.cfg, k, v)
|
||||
self.predictor = EfficientLivePortraitPredictor(use_tensorrt = True, half = True)
|
||||
|
||||
def prepare_source(self, img: np.ndarray) -> torch.Tensor:
|
||||
""" construct the input as standard
|
||||
@@ -57,31 +57,12 @@ class LivePortraitWrapper(object):
|
||||
x = x.to(self.device_id)
|
||||
return x
|
||||
|
||||
def prepare_driving_videos(self, imgs) -> torch.Tensor:
|
||||
""" construct the input as standard
|
||||
imgs: NxBxHxWx3, uint8
|
||||
"""
|
||||
if isinstance(imgs, list):
|
||||
_imgs = np.array(imgs)[..., np.newaxis] # TxHxWx3x1
|
||||
elif isinstance(imgs, np.ndarray):
|
||||
_imgs = imgs
|
||||
else:
|
||||
raise ValueError(f'imgs type error: {type(imgs)}')
|
||||
|
||||
y = _imgs.astype(np.float32) / 255.
|
||||
y = np.clip(y, 0, 1) # clip to 0~1
|
||||
y = torch.from_numpy(y).permute(0, 4, 3, 1, 2) # TxHxWx3x1 -> Tx1x3xHxW
|
||||
y = y.to(self.device_id)
|
||||
|
||||
return y
|
||||
|
||||
def extract_feature_3d(self, x: torch.Tensor) -> torch.Tensor:
|
||||
""" get the appearance feature of the image by F
|
||||
x: Bx3xHxW, normalized to 0~1
|
||||
"""
|
||||
with torch.no_grad():
|
||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
||||
feature_3d = self.appearance_feature_extractor(x)
|
||||
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||
feature_3d = self.appearance_feature_extractor(x)
|
||||
|
||||
return feature_3d.float()
|
||||
|
||||
@@ -91,15 +72,14 @@ class LivePortraitWrapper(object):
|
||||
flag_refine_info: whether to trandform the pose to degrees and the dimention of the reshape
|
||||
return: A dict contains keys: 'pitch', 'yaw', 'roll', 't', 'exp', 'scale', 'kp'
|
||||
"""
|
||||
with torch.no_grad():
|
||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
||||
kp_info = self.motion_extractor(x)
|
||||
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||
kp_info = self.motion_extractor(x)
|
||||
|
||||
if self.cfg.flag_use_half_precision:
|
||||
# float the dict
|
||||
for k, v in kp_info.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
kp_info[k] = v.float()
|
||||
if self.cfg.flag_use_half_precision:
|
||||
# float the dict
|
||||
for k, v in kp_info.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
kp_info[k] = v.float()
|
||||
|
||||
flag_refine_info: bool = kwargs.get('flag_refine_info', True)
|
||||
if flag_refine_info:
|
||||
@@ -265,32 +245,37 @@ class LivePortraitWrapper(object):
|
||||
kp_source: BxNx3
|
||||
kp_driving: BxNx3
|
||||
"""
|
||||
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i))
|
||||
with torch.no_grad():
|
||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
||||
# get decoder input
|
||||
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
|
||||
# decode
|
||||
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
|
||||
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i)
|
||||
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||
# get decoder input
|
||||
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
|
||||
# decode
|
||||
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
|
||||
|
||||
# float the dict
|
||||
if self.cfg.flag_use_half_precision:
|
||||
for k, v in ret_dct.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
ret_dct[k] = v.float()
|
||||
# float the dict
|
||||
if self.cfg.flag_use_half_precision:
|
||||
for k, v in ret_dct.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
ret_dct[k] = v.float()
|
||||
|
||||
return ret_dct
|
||||
|
||||
def parse_output(self, out: torch.Tensor) -> np.ndarray:
|
||||
""" construct the output as standard
|
||||
return: 1xHxWx3, uint8
|
||||
"""
|
||||
out = np.transpose(out.data.cpu().numpy(), [0, 2, 3, 1]) # 1x3xHxW -> 1xHxWx3
|
||||
|
||||
def warp_decode_tensorrt(self, feature_3d, kp_source, kp_driving):
|
||||
inputs = {
|
||||
'feature_3d': np.array(feature_3d.cpu()),
|
||||
'kp_driving': np.array(kp_driving.cpu()),
|
||||
'kp_source': np.array(kp_source.cpu())
|
||||
}
|
||||
generator = self.predictor.run_time(engine_name='generator', task='gw_session',
|
||||
inputs_onnx=inputs, inputs_tensorrt=[feature_3d.cpu(), kp_driving.cpu(), kp_source.cpu()])
|
||||
|
||||
out = np.transpose(generator[0], [0, 2, 3, 1]) # 1x3xHxW -> 1xHxWx3
|
||||
out = np.clip(out, 0, 1) # clip to 0~1
|
||||
out = np.clip(out * 255, 0, 255).astype(np.uint8) # 0~1 -> 0~255
|
||||
|
||||
return out
|
||||
|
||||
out = torch.from_numpy(out).permute(0, 3, 1, 2) / 255
|
||||
|
||||
return {'out': out}
|
||||
|
||||
def calc_retargeting_ratio(self, source_lmk, driving_lmk_lst):
|
||||
input_eye_ratio_lst = []
|
||||
input_lip_ratio_lst = []
|
||||
@@ -303,17 +288,19 @@ class LivePortraitWrapper(object):
|
||||
|
||||
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk):
|
||||
eye_close_ratio = calc_eye_close_ratio(source_lmk[None])
|
||||
eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float().to(self.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).to(self.device_id)
|
||||
eye_close_ratios_tensor = torch.from_numpy(eye_close_ratio).float().to(self.device_id)
|
||||
input_eye_ratio_array = np.array(input_eye_ratio[0][0]).reshape(1, 1)
|
||||
input_eye_ratio_tensor = torch.from_numpy(input_eye_ratio_array).float().to(self.device_id)
|
||||
# [c_s,eyes, c_d,eyes,i]
|
||||
combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1)
|
||||
return combined_eye_ratio_tensor
|
||||
combined_eye_ratios_tensor = torch.cat([eye_close_ratios_tensor, input_eye_ratio_tensor], dim=1)
|
||||
return combined_eye_ratios_tensor
|
||||
|
||||
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk):
|
||||
lip_close_ratio = calc_lip_close_ratio(source_lmk[None])
|
||||
lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().to(self.device_id)
|
||||
# [c_s,lip, c_d,lip,i]
|
||||
input_lip_ratio_tensor = torch.Tensor([input_lip_ratio[0]]).to(self.device_id)
|
||||
input_lip_ratio_array = np.array([input_lip_ratio[0]])
|
||||
input_lip_ratio_tensor = torch.from_numpy(input_lip_ratio_array).float().to(self.device_id)
|
||||
if input_lip_ratio_tensor.shape != [1, 1]:
|
||||
input_lip_ratio_tensor = input_lip_ratio_tensor.reshape(1, 1)
|
||||
combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1)
|
||||
|
||||
@@ -47,7 +47,13 @@ class DenseMotionNetwork(nn.Module):
|
||||
feature_repeat = feature.unsqueeze(1).unsqueeze(1).repeat(1, self.num_kp+1, 1, 1, 1, 1, 1) # (bs, num_kp+1, 1, c, d, h, w)
|
||||
feature_repeat = feature_repeat.view(bs * (self.num_kp+1), -1, d, h, w) # (bs*(num_kp+1), c, d, h, w)
|
||||
sparse_motions = sparse_motions.view((bs * (self.num_kp+1), d, h, w, -1)) # (bs*(num_kp+1), d, h, w, 3)
|
||||
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False)
|
||||
try:
|
||||
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False)
|
||||
except NotImplementedError: #MPS fallback
|
||||
out_device = feature_repeat.device # Store input device
|
||||
feature_repeat = feature_repeat.to('cpu')
|
||||
sparse_motions = sparse_motions.to('cpu')
|
||||
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False).to(out_device)
|
||||
sparse_deformed = sparse_deformed.view((bs, self.num_kp+1, -1, d, h, w)) # (bs, num_kp+1, c, d, h, w)
|
||||
|
||||
return sparse_deformed
|
||||
@@ -61,7 +67,7 @@ class DenseMotionNetwork(nn.Module):
|
||||
# adding background feature
|
||||
try:
|
||||
zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).type(heatmap.type()).to(heatmap.device)
|
||||
except:
|
||||
except ValueError:
|
||||
zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).to(heatmap.device)
|
||||
heatmap = torch.cat([zeros, heatmap], dim=1)
|
||||
heatmap = heatmap.unsqueeze(2) # (bs, 1+num_kp, 1, d, h, w)
|
||||
|
||||
@@ -158,7 +158,11 @@ class DownBlock3d(nn.Module):
|
||||
out = self.conv(x)
|
||||
out = self.norm(out)
|
||||
out = F.relu(out)
|
||||
out = self.pool(out)
|
||||
try:
|
||||
out = self.pool(out)
|
||||
except NotImplementedError:
|
||||
out_device = out.device # Store input device
|
||||
out = self.pool(out.to('cpu')).to(out_device)
|
||||
return out
|
||||
|
||||
|
||||
|
||||
@@ -44,7 +44,11 @@ class WarpingNetwork(nn.Module):
|
||||
self.estimate_occlusion_map = estimate_occlusion_map
|
||||
|
||||
def deform_input(self, inp, deformation):
|
||||
return F.grid_sample(inp, deformation, align_corners=False)
|
||||
try:
|
||||
return F.grid_sample(inp, deformation, align_corners=False)
|
||||
except NotImplementedError:
|
||||
out_device = inp.device # Store input device
|
||||
return F.grid_sample(inp.to('cpu'), deformation.to('cpu'), align_corners=False).to(out_device)
|
||||
|
||||
def forward(self, feature_3d, kp_driving, kp_source):
|
||||
if self.dense_motion_network is not None:
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Make video template
|
||||
"""
|
||||
|
||||
import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pickle
|
||||
from rich.progress import track
|
||||
from .utils.cropper import Cropper
|
||||
|
||||
from .utils.io import load_driving_info
|
||||
from .utils.camera import get_rotation_matrix
|
||||
from .utils.helper import mkdir, basename
|
||||
from .utils.rprint import rlog as log
|
||||
from .config.crop_config import CropConfig
|
||||
from .config.inference_config import InferenceConfig
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
|
||||
class TemplateMaker:
|
||||
|
||||
def __init__(self, inference_cfg: InferenceConfig, crop_cfg: CropConfig):
|
||||
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(cfg=inference_cfg)
|
||||
self.cropper = Cropper(crop_cfg=crop_cfg)
|
||||
|
||||
def make_motion_template(self, video_fp: str, output_path: str, **kwargs):
|
||||
""" make video template (.pkl format)
|
||||
video_fp: driving video file path
|
||||
output_path: where to save the pickle file
|
||||
"""
|
||||
|
||||
driving_rgb_lst = load_driving_info(video_fp)
|
||||
driving_rgb_lst = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
|
||||
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_rgb_lst)
|
||||
I_d_lst = self.live_portrait_wrapper.prepare_driving_videos(driving_rgb_lst)
|
||||
|
||||
n_frames = I_d_lst.shape[0]
|
||||
|
||||
templates = []
|
||||
|
||||
|
||||
for i in track(range(n_frames), description='Making templates...', total=n_frames):
|
||||
I_d_i = I_d_lst[i]
|
||||
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
|
||||
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
|
||||
# collect s_d, R_d, δ_d and t_d for inference
|
||||
template_dct = {
|
||||
'n_frames': n_frames,
|
||||
'frames_index': i,
|
||||
}
|
||||
template_dct['scale'] = x_d_i_info['scale'].cpu().numpy().astype(np.float32)
|
||||
template_dct['R_d'] = R_d_i.cpu().numpy().astype(np.float32)
|
||||
template_dct['exp'] = x_d_i_info['exp'].cpu().numpy().astype(np.float32)
|
||||
template_dct['t'] = x_d_i_info['t'].cpu().numpy().astype(np.float32)
|
||||
|
||||
templates.append(template_dct)
|
||||
|
||||
mkdir(output_path)
|
||||
# Save the dictionary as a pickle file
|
||||
pickle_fp = os.path.join(output_path, f'{basename(video_fp)}.pkl')
|
||||
with open(pickle_fp, 'wb') as f:
|
||||
pickle.dump([templates, driving_lmk_lst], f)
|
||||
log(f"Template saved at {pickle_fp}")
|
||||
+91
-22
@@ -4,14 +4,12 @@
|
||||
cropping function and the related preprocess functions for cropping
|
||||
"""
|
||||
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import numpy as np
|
||||
from .rprint import rprint as print
|
||||
from math import sin, cos, acos, degrees
|
||||
|
||||
DTYPE = np.float32
|
||||
CV2_INTERP = cv2.INTER_LINEAR
|
||||
|
||||
import comfy.model_management as mm
|
||||
|
||||
def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
||||
""" conduct similarity or affine transformation to the image, do not do border operation!
|
||||
@@ -29,6 +27,43 @@ def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
||||
else:
|
||||
return cv2.warpAffine(img, M[:2, :], dsize=_dsize, flags=flags)
|
||||
|
||||
import torch
|
||||
import kornia.geometry.transform as KGT
|
||||
|
||||
def _transform_img_kornia(img, M, dsize, device, flags='bilinear', borderMode='zeros'):
|
||||
"""Conduct similarity or affine transformation to the image using Kornia.
|
||||
|
||||
img: Input image as a PyTorch tensor of shape (C, H, W).
|
||||
M: 2x3 transformation matrix as a PyTorch tensor.
|
||||
dsize: Target shape (width, height).
|
||||
"""
|
||||
|
||||
# Convert dsize to tensor shape (H, W)
|
||||
_dsize = torch.tensor([dsize[1], dsize[0]]) # Kornia expects (H, W)
|
||||
|
||||
# Convert M from numpy.ndarray to PyTorch tensor
|
||||
M = torch.from_numpy(M).float().to(device)
|
||||
if M.shape == (3, 3):
|
||||
M = M[:2, :].unsqueeze(0) # Adjust M to the expected shape Bx2x3
|
||||
elif M.shape == (2, 3):
|
||||
M = M.unsqueeze(0) # Add batch dimension if not present
|
||||
|
||||
# Reshape M for Kornia (1, 2, 3) and upscale to 3D affine matrix if not already
|
||||
if M.shape == (2, 3):
|
||||
M = M.unsqueeze(0) # Add batch dimension
|
||||
|
||||
# Convert image to floating point tensor if not already
|
||||
if img.dtype != torch.float32:
|
||||
img = img.float()
|
||||
img = img.to(device)
|
||||
|
||||
# Reshape img for Kornia (B, C, H, W)
|
||||
img = img.permute(0, 3, 1, 2)
|
||||
|
||||
# Apply the affine transformation
|
||||
img_warped = KGT.warp_affine(img, M, _dsize, mode=flags, padding_mode=borderMode)
|
||||
|
||||
return img_warped
|
||||
|
||||
def _transform_pts(pts, M):
|
||||
""" conduct similarity or affine transformation to the pts
|
||||
@@ -89,29 +124,60 @@ def parse_pt2_from_pt203(pt203, use_lip=True):
|
||||
pt2 = np.stack([pt_left_eye, pt_right_eye], axis=0)
|
||||
return pt2
|
||||
|
||||
def parse_pt2_from_pt9(pt9, use_lip=True):
|
||||
'''
|
||||
animal_face = {"keypoints": ['right eye right', 'right eye left', 'left eye right', 'left eye left', 'nose tip', 'lip right', 'lip left', 'upper lip', 'lower lip'], "skeleton": []}
|
||||
|
||||
def parse_pt2_from_pt68(pt68, use_lip=True):
|
||||
"""
|
||||
parsing the 2 points according to the 68 points, which cancels the roll
|
||||
"""
|
||||
lm_idx = np.array([31, 37, 40, 43, 46, 49, 55], dtype=np.int32) - 1
|
||||
|
||||
'''
|
||||
if use_lip:
|
||||
pt5 = np.stack([
|
||||
np.mean(pt68[lm_idx[[1, 2]], :], 0), # left eye
|
||||
np.mean(pt68[lm_idx[[3, 4]], :], 0), # right eye
|
||||
pt68[lm_idx[0], :], # nose
|
||||
pt68[lm_idx[5], :], # lip
|
||||
pt68[lm_idx[6], :] # lip
|
||||
pt9 = np.stack([
|
||||
(pt9[2]+pt9[3])/2, # left eye
|
||||
(pt9[0]+pt9[1])/2, # right eye
|
||||
pt9[4],
|
||||
# (pt9[5]+pt9[6]+pt9[7]+pt9[8])/4 # lip
|
||||
(pt9[5] + pt9[6] ) / 2 # lip
|
||||
], axis=0)
|
||||
|
||||
pt2 = np.stack([
|
||||
(pt5[0] + pt5[1]) / 2,
|
||||
(pt5[3] + pt5[4]) / 2
|
||||
(pt9[0] + pt9[1]) / 2, # eye
|
||||
pt9[3] # lip
|
||||
], axis=0)
|
||||
else:
|
||||
pt2 = np.stack([
|
||||
np.mean(pt68[lm_idx[[1, 2]], :], 0), # left eye
|
||||
np.mean(pt68[lm_idx[[3, 4]], :], 0), # right eye
|
||||
(pt9[2] + pt9[3]) / 2,
|
||||
(pt9[0] + pt9[1]) / 2,
|
||||
], axis=0)
|
||||
|
||||
return pt2
|
||||
|
||||
def parse_pt2_from_pt68(pt68, use_lip=True):
|
||||
'''
|
||||
face = {"keypoints": ['right cheekbone 1', 'right cheekbone 2', 'right cheek 1', 'right cheek 2', 'right cheek 3', 'right cheek 4', 'right cheek 5', 'right chin', 'chin center',
|
||||
'left chin', 'left cheek 5', 'left cheek 4', 'left cheek 3', 'left cheek 2', 'left cheek 1', 'left cheekbone 2', 'left cheekbone 1', 'right eyebrow 1', 'right eyebrow 2', 'right eyebrow 3',
|
||||
'right eyebrow 4', 'right eyebrow 5', 'left eyebrow 1', 'left eyebrow 2', 'left eyebrow 3', 'left eyebrow 4', 'left eyebrow 5', 'nasal bridge 1', 'nasal bridge 2', 'nasal bridge 3', 'nasal bridge 4',
|
||||
'right nasal wing 1', 'right nasal wing 2', 'nasal wing center', 'left nasal wing 1', 'left nasal wing 2', 'right eye eye corner 1', 'right eye upper eyelid 1', 'right eye upper eyelid 2',
|
||||
'right eye eye corner 2', 'right eye lower eyelid 2', 'right eye lower eyelid 1', 'left eye eye corner 1', 'left eye upper eyelid 1', 'left eye upper eyelid 2', 'left eye eye corner 2', 'left eye lower eyelid 2',
|
||||
'left eye lower eyelid 1', 'right mouth corner', 'upper lip outer edge 1', 'upper lip outer edge 2', 'upper lip outer edge 3', 'upper lip outer edge 4', 'upper lip outer edge 5', 'left mouth corner',
|
||||
'lower lip outer edge 5', 'lower lip outer edge 4', 'lower lip outer edge 3', 'lower lip outer edge 2', 'lower lip outer edge 1', 'upper lip inter edge 1', 'upper lip inter edge 2', 'upper lip inter edge 3',
|
||||
'upper lip inter edge 4', 'upper lip inter edge 5', 'lower lip inter edge 3', 'lower lip inter edge 2', 'lower lip inter edge 1'], "skeleton": []}
|
||||
|
||||
|
||||
'''
|
||||
if use_lip:
|
||||
pt68 = np.stack([
|
||||
(pt68[42] + pt68[43] + pt68[44] + pt68[45] + pt68[46]+ pt68[47])/6, # left eye
|
||||
(pt68[36] + pt68[37] + pt68[38] + pt68[39] + pt68[40] + pt68[41]) / 6, # right eye
|
||||
(pt68[48] + pt68[54])/2
|
||||
|
||||
], axis=0)
|
||||
pt2 = np.stack([
|
||||
(pt68[0] + pt68[1]) / 2,
|
||||
pt68[2]
|
||||
], axis=0)
|
||||
else:
|
||||
pt2 = np.stack([
|
||||
(pt68[42] + pt68[43] + pt68[44] + pt68[45] + pt68[46] + pt68[47]) / 6, # left eye
|
||||
(pt68[36] + pt68[37] + pt68[38] + pt68[39] + pt68[40] + pt68[41]) / 6, # right eye
|
||||
], axis=0)
|
||||
|
||||
return pt2
|
||||
@@ -148,6 +214,8 @@ def parse_pt2_from_pt_x(pts, use_lip=True):
|
||||
elif pts.shape[0] > 101:
|
||||
# take the first 101 points
|
||||
pt2 = parse_pt2_from_pt101(pts[:101], use_lip=use_lip)
|
||||
elif pts.shape[0] == 9:
|
||||
pt2 = parse_pt2_from_pt9(pts, use_lip=use_lip)
|
||||
else:
|
||||
raise Exception(f'Unknow shape: {pts.shape}')
|
||||
|
||||
@@ -350,13 +418,15 @@ def crop_image(img, pts: np.ndarray, **kwargs):
|
||||
dsize = kwargs.get('dsize', 224)
|
||||
scale = kwargs.get('scale', 1.5) # 1.5 | 1.6
|
||||
vy_ratio = kwargs.get('vy_ratio', -0.1) # -0.0625 | -0.1
|
||||
vx_ratio = kwargs.get('vx_ratio', 0)
|
||||
|
||||
M_INV, _ = _estimate_similar_transform_from_pts(
|
||||
pts,
|
||||
dsize=dsize,
|
||||
scale=scale,
|
||||
vy_ratio=vy_ratio,
|
||||
flag_do_rot=kwargs.get('flag_do_rot', True),
|
||||
vx_ratio=vx_ratio,
|
||||
flag_do_rot=kwargs.get('rotate', True),
|
||||
)
|
||||
|
||||
if img is None:
|
||||
@@ -390,4 +460,3 @@ def average_bbox_lst(bbox_lst):
|
||||
return None
|
||||
bbox_arr = np.array(bbox_lst)
|
||||
return np.mean(bbox_arr, axis=0).tolist()
|
||||
|
||||
|
||||
@@ -1,32 +1,22 @@
|
||||
# coding: utf-8
|
||||
|
||||
import numpy as np
|
||||
import os.path as osp
|
||||
from typing import List, Union, Tuple
|
||||
from dataclasses import dataclass, field
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
from .landmark_runner import LandmarkRunner
|
||||
from .face_analysis_diy import FaceAnalysisDIY
|
||||
#from .helper import prefix
|
||||
from .crop import crop_image, crop_image_by_bbox, parse_bbox_from_landmark, average_bbox_lst
|
||||
#from .timer import Timer
|
||||
from .rprint import rlog as log
|
||||
from .io import load_image_rgb
|
||||
#from .video import VideoWriter, get_fps, change_video_fps
|
||||
from .crop import crop_image
|
||||
|
||||
import folder_paths
|
||||
import os
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Trajectory:
|
||||
start: int = -1 # 起始帧 闭区间
|
||||
end: int = -1 # 结束帧 闭区间
|
||||
start: int = -1
|
||||
end: int = -1
|
||||
lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list
|
||||
bbox_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # bbox list
|
||||
frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list
|
||||
@@ -34,10 +24,10 @@ class Trajectory:
|
||||
|
||||
|
||||
class Cropper(object):
|
||||
def __init__(self, provider, **kwargs) -> None:
|
||||
def __init__(self, **kwargs) -> None:
|
||||
device_id = kwargs.get('device_id', 0)
|
||||
provider = kwargs.get('onnx_device', 'CPU')
|
||||
self.landmark_runner = LandmarkRunner(
|
||||
#ckpt_path=make_abs_path('../../pretrained_weights/liveportrait/landmark.onnx'),
|
||||
ckpt_path=os.path.join(folder_paths.models_dir, 'liveportrait', 'landmark.onnx'),
|
||||
onnx_provider=provider,
|
||||
device_id=device_id
|
||||
@@ -52,21 +42,8 @@ class Cropper(object):
|
||||
self.face_analysis_wrapper.prepare(ctx_id=device_id, det_size=(512, 512))
|
||||
self.face_analysis_wrapper.warmup()
|
||||
|
||||
self.crop_cfg = kwargs.get('crop_cfg', None)
|
||||
|
||||
def update_config(self, user_args):
|
||||
for k, v in user_args.items():
|
||||
if hasattr(self.crop_cfg, k):
|
||||
setattr(self.crop_cfg, k, v)
|
||||
|
||||
def crop_single_image(self, obj, **kwargs):
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
|
||||
# crop and align a single image
|
||||
if isinstance(obj, str):
|
||||
img_rgb = load_image_rgb(obj)
|
||||
elif isinstance(obj, np.ndarray):
|
||||
img_rgb = obj
|
||||
def crop_single_image(self, img_rgb, dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate):
|
||||
direction = face_index_order
|
||||
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
img_rgb,
|
||||
@@ -75,73 +52,34 @@ class Cropper(object):
|
||||
)
|
||||
|
||||
if len(src_face) == 0:
|
||||
log('No face detected in the source image.')
|
||||
raise Exception("No face detected in the source image!")
|
||||
elif len(src_face) > 1:
|
||||
log(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
||||
ret_dct = {}
|
||||
return ret_dct
|
||||
#raise Exception("No face detected in the source image!")
|
||||
#elif len(src_face) > 1:
|
||||
# print(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
||||
|
||||
src_face = src_face[0]
|
||||
src_face = src_face[face_index] # choose the index if multiple faces detected
|
||||
pts = src_face.landmark_2d_106
|
||||
|
||||
|
||||
# crop the face
|
||||
ret_dct = crop_image(
|
||||
img_rgb, # ndarray
|
||||
pts, # 106x2 or Nx2
|
||||
dsize=kwargs.get('dsize', 512),
|
||||
scale=kwargs.get('scale', 2.3),
|
||||
vy_ratio=kwargs.get('vy_ratio', -0.15),
|
||||
dsize=dsize,
|
||||
scale=scale,
|
||||
vy_ratio=vy_ratio,
|
||||
vx_ratio=vx_ratio,
|
||||
rotate=rotate
|
||||
)
|
||||
# update a 256x256 version for network input or else
|
||||
ret_dct['img_crop_256x256'] = cv2.resize(ret_dct['img_crop'], (256, 256), interpolation=cv2.INTER_AREA)
|
||||
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / kwargs.get('dsize', 512)
|
||||
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / dsize
|
||||
|
||||
input_image_size = img_rgb.shape[:2]
|
||||
ret_dct['input_image_size'] = input_image_size
|
||||
|
||||
recon_ret = self.landmark_runner.run(img_rgb, pts)
|
||||
lmk = recon_ret['pts']
|
||||
ret_dct['lmk_crop'] = lmk
|
||||
|
||||
return ret_dct
|
||||
|
||||
def get_retargeting_lmk_info(self, driving_rgb_lst):
|
||||
# TODO: implement a tracking-based version
|
||||
driving_lmk_lst = []
|
||||
for driving_image in driving_rgb_lst:
|
||||
ret_dct = self.crop_single_image(driving_image)
|
||||
driving_lmk_lst.append(ret_dct['lmk_crop'])
|
||||
return driving_lmk_lst
|
||||
|
||||
def make_video_clip(self, driving_rgb_lst, output_path, output_fps=30, **kwargs):
|
||||
trajectory = Trajectory()
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
for idx, driving_image in enumerate(driving_rgb_lst):
|
||||
if idx == 0 or trajectory.start == -1:
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
driving_image,
|
||||
flag_do_landmark_2d_106=True,
|
||||
direction=direction
|
||||
)
|
||||
if len(src_face) == 0:
|
||||
# No face detected in the driving_image
|
||||
continue
|
||||
elif len(src_face) > 1:
|
||||
log(f'More than one face detected in the driving frame_{idx}, only pick one face by rule {direction}.')
|
||||
src_face = src_face[0]
|
||||
pts = src_face.landmark_2d_106
|
||||
lmk_203 = self.landmark_runner(driving_image, pts)['pts']
|
||||
trajectory.start, trajectory.end = idx, idx
|
||||
else:
|
||||
lmk_203 = self.face_recon_wrapper(driving_image, trajectory.lmk_lst[-1])['pts']
|
||||
trajectory.end = idx
|
||||
|
||||
trajectory.lmk_lst.append(lmk_203)
|
||||
ret_bbox = parse_bbox_from_landmark(lmk_203, scale=self.crop_cfg.globalscale, vy_ratio=elf.crop_cfg.vy_ratio)['bbox']
|
||||
bbox = [ret_bbox[0, 0], ret_bbox[0, 1], ret_bbox[2, 0], ret_bbox[2, 1]] # 4,
|
||||
trajectory.bbox_lst.append(bbox) # bbox
|
||||
trajectory.frame_rgb_lst.append(driving_image)
|
||||
|
||||
global_bbox = average_bbox_lst(trajectory.bbox_lst)
|
||||
for idx, (frame_rgb, lmk) in enumerate(zip(trajectory.frame_rgb_lst, trajectory.lmk_lst)):
|
||||
ret_dct = crop_image_by_bbox(
|
||||
frame_rgb, global_bbox, lmk=lmk,
|
||||
dsize=self.video_crop_cfg.dsize, flag_rot=self.video_crop_cfg.flag_rot, borderValue=self.video_crop_cfg.borderValue
|
||||
)
|
||||
frame_rgb_crop = ret_dct['img_crop']
|
||||
return ret_dct
|
||||
@@ -1,16 +1,30 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
face detectoin and alignment using InsightFace
|
||||
face detection and alignment using InsightFace
|
||||
"""
|
||||
from insightface.utils import transform
|
||||
|
||||
#patch Insightface function to get rid of the annoying warnings
|
||||
def patched_estimate_affine_matrix_3d23d(X, Y):
|
||||
''' Using least-squares solution
|
||||
Args:
|
||||
X: [n, 3]. 3d points(fixed)
|
||||
Y: [n, 3]. corresponding 3d points(moving). Y = PX
|
||||
Returns:
|
||||
P_Affine: (3, 4). Affine camera matrix (the third row is [0, 0, 0, 1]).
|
||||
'''
|
||||
X_homo = np.hstack((X, np.ones([X.shape[0],1]))) # n x 4
|
||||
P = np.linalg.lstsq(X_homo, Y, rcond=None)[0].T # Affine matrix. 3 x 4
|
||||
return P
|
||||
|
||||
transform.estimate_affine_matrix_3d23d = patched_estimate_affine_matrix_3d23d
|
||||
|
||||
import numpy as np
|
||||
from .rprint import rlog as log
|
||||
from insightface.app import FaceAnalysis
|
||||
from insightface.app.common import Face
|
||||
from .timer import Timer
|
||||
|
||||
|
||||
def sort_by_direction(faces, direction: str = 'large-small', face_center=None):
|
||||
if len(faces) <= 0:
|
||||
return faces
|
||||
@@ -76,4 +90,4 @@ class FaceAnalysisDIY(FaceAnalysis):
|
||||
self.get(img_bgr)
|
||||
|
||||
elapse = self.timer.toc()
|
||||
log(f'FaceAnalysisDIY warmup time: {elapse:.3f}s')
|
||||
print(f'FaceAnalysisDIY warmup time: {elapse:.3f}s')
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from pykalman import KalmanFilter
|
||||
|
||||
|
||||
def smooth(x_d_lst, shape, device, observation_variance=3e-6, process_variance=1e-5):
|
||||
# Reshape x_d_lst, skipping None values
|
||||
x_d_lst_reshape = [x.reshape(-1) for x in x_d_lst if x is not None]
|
||||
|
||||
if not x_d_lst_reshape: # Check if x_d_lst_reshape is empty after filtering
|
||||
return [None] * len(x_d_lst) # Return a list of Nones with the same length as x_d_lst
|
||||
|
||||
x_d_stacked = np.vstack(x_d_lst_reshape)
|
||||
|
||||
kf = KalmanFilter(
|
||||
initial_state_mean=x_d_stacked[0],
|
||||
n_dim_obs=x_d_stacked.shape[1],
|
||||
transition_covariance=process_variance * np.eye(x_d_stacked.shape[1]),
|
||||
observation_covariance=observation_variance * np.eye(x_d_stacked.shape[1])
|
||||
)
|
||||
|
||||
smoothed_state_means, _ = kf.smooth(x_d_stacked)
|
||||
|
||||
# Initialize an iterator for smoothed_state_means
|
||||
smoothed_states_iter = iter(smoothed_state_means)
|
||||
|
||||
# Create x_d_lst_smooth, inserting None for each None encountered in the original list
|
||||
x_d_lst_smooth = [torch.tensor(next(smoothed_states_iter).reshape(shape[-2:]), dtype=torch.float32, device=device) if x is not None else None for x in x_d_lst]
|
||||
|
||||
return x_d_lst_smooth
|
||||
@@ -4,58 +4,15 @@
|
||||
utility functions and classes to handle feature extraction and model loading
|
||||
"""
|
||||
|
||||
import os
|
||||
import os.path as osp
|
||||
import cv2
|
||||
import torch
|
||||
from collections import OrderedDict
|
||||
|
||||
def suffix(filename):
|
||||
"""a.jpg -> jpg"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return ""
|
||||
return filename[pos + 1:]
|
||||
|
||||
|
||||
def prefix(filename):
|
||||
"""a.jpg -> a"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return filename
|
||||
return filename[:pos]
|
||||
|
||||
|
||||
def basename(filename):
|
||||
"""a/b/c.jpg -> c"""
|
||||
return prefix(osp.basename(filename))
|
||||
|
||||
|
||||
def is_video(file_path):
|
||||
if file_path.lower().endswith((".mp4", ".mov", ".avi", ".webm")) or osp.isdir(file_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
def is_template(file_path):
|
||||
if file_path.endswith(".pkl"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def mkdir(d, log=False):
|
||||
# return self-assined `d`, for one line code
|
||||
if not osp.exists(d):
|
||||
os.makedirs(d, exist_ok=True)
|
||||
if log:
|
||||
print(f"Make dir: {d}")
|
||||
return d
|
||||
|
||||
|
||||
def squeeze_tensor_to_numpy(tensor):
|
||||
out = tensor.data.squeeze(0).cpu().numpy()
|
||||
return out
|
||||
|
||||
|
||||
def dct2cuda(dct: dict, device_id: int):
|
||||
for key in dct:
|
||||
dct[key] = torch.tensor(dct[key]).to(device_id)
|
||||
@@ -95,11 +52,6 @@ def calculate_transformation(config, s_kp_info, t_0_kp_info, t_i_kp_info, R_s, R
|
||||
new_scale = s_kp_info['scale'] * (t_i_kp_info['scale'] / t_0_kp_info['scale'])
|
||||
return new_rotation, new_expression, new_translation, new_scale
|
||||
|
||||
def load_description(fp):
|
||||
with open(fp, 'r', encoding='utf-8') as f:
|
||||
content = f.read()
|
||||
return content
|
||||
|
||||
|
||||
def resize_to_limit(img, max_dim=1280, n=2):
|
||||
h, w = img.shape[:2]
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
import os
|
||||
from glob import glob
|
||||
import os.path as osp
|
||||
import imageio
|
||||
import numpy as np
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
|
||||
def load_image_rgb(image_path: str):
|
||||
if not osp.exists(image_path):
|
||||
raise FileNotFoundError(f"Image not found: {image_path}")
|
||||
img = cv2.imread(image_path, cv2.IMREAD_COLOR)
|
||||
return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
|
||||
|
||||
def load_driving_info(driving_info):
|
||||
driving_video_ori = []
|
||||
|
||||
def load_images_from_directory(directory):
|
||||
image_paths = sorted(glob(osp.join(directory, '*.png')) + glob(osp.join(directory, '*.jpg')))
|
||||
return [load_image_rgb(im_path) for im_path in image_paths]
|
||||
|
||||
def load_images_from_video(file_path):
|
||||
reader = imageio.get_reader(file_path)
|
||||
return [image for idx, image in enumerate(reader)]
|
||||
|
||||
if osp.isdir(driving_info):
|
||||
driving_video_ori = load_images_from_directory(driving_info)
|
||||
elif osp.isfile(driving_info):
|
||||
driving_video_ori = load_images_from_video(driving_info)
|
||||
|
||||
return driving_video_ori
|
||||
|
||||
|
||||
def contiguous(obj):
|
||||
if not obj.flags.c_contiguous:
|
||||
obj = obj.copy(order="C")
|
||||
return obj
|
||||
|
||||
|
||||
def _resize_to_limit(img: np.ndarray, max_dim=1920, n=2):
|
||||
"""
|
||||
ajust the size of the image so that the maximum dimension does not exceed max_dim, and the width and the height of the image are multiples of n.
|
||||
:param img: the image to be processed.
|
||||
:param max_dim: the maximum dimension constraint.
|
||||
:param n: the number that needs to be multiples of.
|
||||
:return: the adjusted image.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
|
||||
# ajust the size of the image according to the maximum dimension
|
||||
if max_dim > 0 and max(h, w) > max_dim:
|
||||
if h > w:
|
||||
new_h = max_dim
|
||||
new_w = int(w * (max_dim / h))
|
||||
else:
|
||||
new_w = max_dim
|
||||
new_h = int(h * (max_dim / w))
|
||||
img = cv2.resize(img, (new_w, new_h))
|
||||
|
||||
# ensure that the image dimensions are multiples of n
|
||||
n = max(n, 1)
|
||||
new_h = img.shape[0] - (img.shape[0] % n)
|
||||
new_w = img.shape[1] - (img.shape[1] % n)
|
||||
|
||||
if new_h == 0 or new_w == 0:
|
||||
# when the width or height is less than n, no need to process
|
||||
return img
|
||||
|
||||
if new_h != img.shape[0] or new_w != img.shape[1]:
|
||||
img = img[:new_h, :new_w]
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def load_img_online(obj, mode="bgr", **kwargs):
|
||||
max_dim = kwargs.get("max_dim", 1920)
|
||||
n = kwargs.get("n", 2)
|
||||
if isinstance(obj, str):
|
||||
if mode.lower() == "gray":
|
||||
img = cv2.imread(obj, cv2.IMREAD_GRAYSCALE)
|
||||
else:
|
||||
img = cv2.imread(obj, cv2.IMREAD_COLOR)
|
||||
else:
|
||||
img = obj
|
||||
|
||||
# Resize image to satisfy constraints
|
||||
img = _resize_to_limit(img, max_dim=max_dim, n=n)
|
||||
|
||||
if mode.lower() == "bgr":
|
||||
return contiguous(img)
|
||||
elif mode.lower() == "rgb":
|
||||
return contiguous(img[..., ::-1])
|
||||
else:
|
||||
raise Exception(f"Unknown mode {mode}")
|
||||
@@ -1,19 +1,12 @@
|
||||
# coding: utf-8
|
||||
|
||||
import os.path as osp
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import torch
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
from .timer import Timer
|
||||
from .rprint import rlog
|
||||
from .crop import crop_image, _transform_pts
|
||||
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
|
||||
def to_ndarray(obj):
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return obj.cpu().numpy()
|
||||
@@ -22,12 +15,11 @@ def to_ndarray(obj):
|
||||
else:
|
||||
return np.array(obj)
|
||||
|
||||
|
||||
class LandmarkRunner(object):
|
||||
"""landmark runner"""
|
||||
def __init__(self, **kwargs):
|
||||
ckpt_path = kwargs.get('ckpt_path')
|
||||
onnx_provider = kwargs.get('onnx_provider', 'cuda') # 默认用cuda
|
||||
onnx_provider = kwargs.get('onnx_provider', 'cuda')
|
||||
device_id = kwargs.get('device_id', 0)
|
||||
self.dsize = kwargs.get('dsize', 224)
|
||||
self.timer = Timer()
|
||||
@@ -40,7 +32,7 @@ class LandmarkRunner(object):
|
||||
)
|
||||
else:
|
||||
opts = onnxruntime.SessionOptions()
|
||||
opts.intra_op_num_threads = 4 # 默认线程数为 4
|
||||
opts.intra_op_num_threads = 4
|
||||
self.session = onnxruntime.InferenceSession(
|
||||
ckpt_path, providers=['CPUExecutionProvider'],
|
||||
sess_options=opts
|
||||
@@ -78,7 +70,6 @@ class LandmarkRunner(object):
|
||||
}
|
||||
|
||||
def warmup(self):
|
||||
# 构造dummy image进行warmup
|
||||
self.timer.tic()
|
||||
|
||||
dummy_image = np.zeros((1, 3, self.dsize, self.dsize), dtype=np.float32)
|
||||
@@ -86,4 +77,4 @@ class LandmarkRunner(object):
|
||||
_ = self._run(dummy_image)
|
||||
|
||||
elapse = self.timer.toc()
|
||||
rlog(f'LandmarkRunner warmup time: {elapse:.3f}s')
|
||||
print(f'LandmarkRunner warmup time: {elapse:.3f}s')
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
custom print and log functions
|
||||
"""
|
||||
|
||||
__all__ = ['rprint', 'rlog']
|
||||
|
||||
try:
|
||||
from rich.console import Console
|
||||
console = Console()
|
||||
rprint = console.print
|
||||
rlog = console.log
|
||||
except:
|
||||
rprint = print
|
||||
rlog = print
|
||||
@@ -1,139 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
functions for processing video
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
import numpy as np
|
||||
import subprocess
|
||||
import imageio
|
||||
import cv2
|
||||
|
||||
from rich.progress import track
|
||||
from .helper import prefix
|
||||
from .rprint import rprint as print
|
||||
|
||||
|
||||
def exec_cmd(cmd):
|
||||
subprocess.run(cmd, shell=True, check=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
|
||||
|
||||
def images2video(images, wfp, **kwargs):
|
||||
fps = kwargs.get('fps', 30)
|
||||
video_format = kwargs.get('format', 'mp4') # default is mp4 format
|
||||
codec = kwargs.get('codec', 'libx264') # default is libx264 encoding
|
||||
quality = kwargs.get('quality') # video quality
|
||||
pixelformat = kwargs.get('pixelformat', 'yuv420p') # video pixel format
|
||||
image_mode = kwargs.get('image_mode', 'rgb')
|
||||
macro_block_size = kwargs.get('macro_block_size', 2)
|
||||
ffmpeg_params = ['-crf', str(kwargs.get('crf', 18))]
|
||||
|
||||
writer = imageio.get_writer(
|
||||
wfp, fps=fps, format=video_format,
|
||||
codec=codec, quality=quality, ffmpeg_params=ffmpeg_params, pixelformat=pixelformat, macro_block_size=macro_block_size
|
||||
)
|
||||
|
||||
n = len(images)
|
||||
for i in track(range(n), description='writing', transient=True):
|
||||
if image_mode.lower() == 'bgr':
|
||||
writer.append_data(images[i][..., ::-1])
|
||||
else:
|
||||
writer.append_data(images[i])
|
||||
|
||||
writer.close()
|
||||
|
||||
# print(f':smiley: Dump to {wfp}\n', style="bold green")
|
||||
print(f'Dump to {wfp}\n')
|
||||
|
||||
|
||||
def video2gif(video_fp, fps=30, size=256):
|
||||
if osp.exists(video_fp):
|
||||
d = osp.split(video_fp)[0]
|
||||
fn = prefix(osp.basename(video_fp))
|
||||
palette_wfp = osp.join(d, 'palette.png')
|
||||
gif_wfp = osp.join(d, f'{fn}.gif')
|
||||
# generate the palette
|
||||
cmd = f'ffmpeg -i {video_fp} -vf "fps={fps},scale={size}:-1:flags=lanczos,palettegen" {palette_wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
# use the palette to generate the gif
|
||||
cmd = f'ffmpeg -i {video_fp} -i {palette_wfp} -filter_complex "fps={fps},scale={size}:-1:flags=lanczos[x];[x][1:v]paletteuse" {gif_wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
else:
|
||||
print(f'video_fp: {video_fp} not exists!')
|
||||
|
||||
|
||||
def merge_audio_video(video_fp, audio_fp, wfp):
|
||||
if osp.exists(video_fp) and osp.exists(audio_fp):
|
||||
cmd = f'ffmpeg -i {video_fp} -i {audio_fp} -c:v copy -c:a aac {wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
print(f'merge {video_fp} and {audio_fp} to {wfp}')
|
||||
else:
|
||||
print(f'video_fp: {video_fp} or audio_fp: {audio_fp} not exists!')
|
||||
|
||||
|
||||
def blend(img: np.ndarray, mask: np.ndarray, background_color=(255, 255, 255)):
|
||||
mask_float = mask.astype(np.float32) / 255.
|
||||
background_color = np.array(background_color).reshape([1, 1, 3])
|
||||
bg = np.ones_like(img) * background_color
|
||||
img = np.clip(mask_float * img + (1 - mask_float) * bg, 0, 255).astype(np.uint8)
|
||||
return img
|
||||
|
||||
|
||||
def concat_frames(I_p_lst, driving_rgb_lst, img_rgb):
|
||||
# TODO: add more concat style, e.g., left-down corner driving
|
||||
out_lst = []
|
||||
for idx, _ in track(enumerate(I_p_lst), total=len(I_p_lst), description='Concatenating result...'):
|
||||
source_image_drived = I_p_lst[idx]
|
||||
image_drive = driving_rgb_lst[idx]
|
||||
|
||||
# resize images to match source_image_drived shape
|
||||
h, w, _ = source_image_drived.shape
|
||||
image_drive_resized = cv2.resize(image_drive, (w, h))
|
||||
img_rgb_resized = cv2.resize(img_rgb, (w, h))
|
||||
|
||||
# concatenate images horizontally
|
||||
frame = np.concatenate((image_drive_resized, img_rgb_resized, source_image_drived), axis=1)
|
||||
out_lst.append(frame)
|
||||
return out_lst
|
||||
|
||||
|
||||
class VideoWriter:
|
||||
def __init__(self, **kwargs):
|
||||
self.fps = kwargs.get('fps', 30)
|
||||
self.wfp = kwargs.get('wfp', 'video.mp4')
|
||||
self.video_format = kwargs.get('format', 'mp4')
|
||||
self.codec = kwargs.get('codec', 'libx264')
|
||||
self.quality = kwargs.get('quality')
|
||||
self.pixelformat = kwargs.get('pixelformat', 'yuv420p')
|
||||
self.image_mode = kwargs.get('image_mode', 'rgb')
|
||||
self.ffmpeg_params = kwargs.get('ffmpeg_params')
|
||||
|
||||
self.writer = imageio.get_writer(
|
||||
self.wfp, fps=self.fps, format=self.video_format,
|
||||
codec=self.codec, quality=self.quality,
|
||||
ffmpeg_params=self.ffmpeg_params, pixelformat=self.pixelformat
|
||||
)
|
||||
|
||||
def write(self, image):
|
||||
if self.image_mode.lower() == 'bgr':
|
||||
self.writer.append_data(image[..., ::-1])
|
||||
else:
|
||||
self.writer.append_data(image)
|
||||
|
||||
def close(self):
|
||||
if self.writer is not None:
|
||||
self.writer.close()
|
||||
|
||||
|
||||
def change_video_fps(input_file, output_file, fps=20, codec='libx264', crf=5):
|
||||
cmd = f"ffmpeg -i {input_file} -c:v {codec} -crf {crf} -r {fps} {output_file} -y"
|
||||
exec_cmd(cmd)
|
||||
|
||||
|
||||
def get_fps(filepath):
|
||||
import ffmpeg
|
||||
probe = ffmpeg.probe(filepath)
|
||||
video_stream = next((stream for stream in probe['streams'] if stream['codec_type'] == 'video'), None)
|
||||
fps = eval(video_stream['avg_frame_rate'])
|
||||
return fps
|
||||
@@ -4,37 +4,45 @@ import yaml
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import cv2
|
||||
from tqdm import tqdm
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
from .liveportrait.config.argument_config import ArgumentConfig
|
||||
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
|
||||
from .liveportrait.utils.cropper import Cropper
|
||||
from .liveportrait.modules.spade_generator import SPADEDecoder
|
||||
from .liveportrait.modules.warping_network import WarpingNetwork
|
||||
from .liveportrait.modules.motion_extractor import MotionExtractor
|
||||
from .liveportrait.modules.appearance_feature_extractor import AppearanceFeatureExtractor
|
||||
from .liveportrait.modules.stitching_retargeting_network import StitchingRetargetingNetwork
|
||||
from .liveportrait.modules.appearance_feature_extractor import (
|
||||
AppearanceFeatureExtractor,
|
||||
)
|
||||
from .liveportrait.modules.stitching_retargeting_network import (
|
||||
StitchingRetargetingNetwork,
|
||||
)
|
||||
from .liveportrait.utils.camera import get_rotation_matrix
|
||||
from .liveportrait.utils.crop import _transform_img_kornia
|
||||
|
||||
import logging
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
class InferenceConfig:
|
||||
def __init__(self,
|
||||
mask_crop = None,
|
||||
flag_use_half_precision=True,
|
||||
flag_lip_zero=True,
|
||||
lip_zero_threshold=0.03,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
anchor_frame=0,
|
||||
input_shape=(256, 256),
|
||||
flag_write_result=True,
|
||||
flag_pasteback=True,
|
||||
ref_max_shape=1280,
|
||||
ref_shape_n=2,
|
||||
device_id=0,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True):
|
||||
def __init__(
|
||||
self,
|
||||
flag_use_half_precision=True,
|
||||
flag_lip_zero=True,
|
||||
lip_zero_threshold=0.03,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
input_shape=(256, 256),
|
||||
device_id=0,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True,
|
||||
):
|
||||
self.flag_use_half_precision = flag_use_half_precision
|
||||
self.flag_lip_zero = flag_lip_zero
|
||||
self.lip_zero_threshold = lip_zero_threshold
|
||||
@@ -42,68 +50,26 @@ class InferenceConfig:
|
||||
self.flag_lip_retargeting = flag_lip_retargeting
|
||||
self.flag_stitching = flag_stitching
|
||||
self.flag_relative = flag_relative
|
||||
self.anchor_frame = anchor_frame
|
||||
self.input_shape = input_shape
|
||||
self.flag_write_result = flag_write_result
|
||||
self.flag_pasteback = flag_pasteback
|
||||
self.ref_max_shape = ref_max_shape
|
||||
self.ref_shape_n = ref_shape_n
|
||||
self.device_id = device_id
|
||||
self.flag_do_crop = flag_do_crop
|
||||
self.flag_do_rot = flag_do_rot
|
||||
self.mask_crop=mask_crop
|
||||
|
||||
class CropConfig:
|
||||
def __init__(self, dsize=512, scale=2.3, vx_ratio=0, vy_ratio=-0.125):
|
||||
self.dsize = dsize
|
||||
self.scale = scale
|
||||
self.vx_ratio = vx_ratio
|
||||
self.vy_ratio = vy_ratio
|
||||
|
||||
class ArgumentConfig:
|
||||
def __init__(self,
|
||||
device_id=0,
|
||||
flag_lip_zero=True,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
flag_pasteback=True,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True,
|
||||
dsize=512,
|
||||
scale=2.3,
|
||||
vx_ratio=0,
|
||||
vy_ratio=-0.125,
|
||||
):
|
||||
self.device_id = device_id
|
||||
self.flag_lip_zero = flag_lip_zero
|
||||
self.flag_eye_retargeting = flag_eye_retargeting
|
||||
self.flag_lip_retargeting = flag_lip_retargeting
|
||||
self.flag_stitching = flag_stitching
|
||||
self.flag_relative = flag_relative
|
||||
self.flag_pasteback = flag_pasteback
|
||||
self.flag_do_crop = flag_do_crop
|
||||
self.flag_do_rot = flag_do_rot
|
||||
self.dsize = dsize
|
||||
self.scale = scale
|
||||
self.vx_ratio = vx_ratio
|
||||
self.vy_ratio = vy_ratio
|
||||
|
||||
|
||||
class DownloadAndLoadLivePortraitModels:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
},
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"precision": (
|
||||
"precision": (
|
||||
[
|
||||
'fp16',
|
||||
'fp32',
|
||||
], {
|
||||
"default": 'fp16'
|
||||
}),
|
||||
}
|
||||
"fp16",
|
||||
"fp32",
|
||||
"auto",
|
||||
],
|
||||
{"default": "auto"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
|
||||
@@ -111,96 +77,138 @@ class DownloadAndLoadLivePortraitModels:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def loadmodel(self, precision='fp16'):
|
||||
def loadmodel(self, precision="fp16"):
|
||||
device = mm.get_torch_device()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
if precision == 'auto':
|
||||
try:
|
||||
if mm.is_device_mps(device):
|
||||
log.info("LivePortrait using fp32 for MPS")
|
||||
dtype = 'fp32'
|
||||
elif mm.should_use_fp16():
|
||||
log.info("LivePortrait using fp16")
|
||||
dtype = 'fp16'
|
||||
else:
|
||||
log.info("LivePortrait using fp32")
|
||||
dtype = 'fp32'
|
||||
except:
|
||||
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
|
||||
else:
|
||||
dtype = precision
|
||||
log.info(f"LivePortrait using {dtype}")
|
||||
|
||||
pbar = comfy.utils.ProgressBar(3)
|
||||
|
||||
download_path = os.path.join(folder_paths.models_dir, "liveportrait")
|
||||
model_path = os.path.join(download_path)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading model to: {model_path}")
|
||||
log.info(f"Downloading model to: {model_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="Kijai/LivePortrait_safetensors",
|
||||
local_dir=download_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
model_config_path = os.path.join(script_directory, 'liveportrait', 'config', 'models.yaml')
|
||||
with open(model_config_path, 'r') as file:
|
||||
snapshot_download(
|
||||
repo_id="Kijai/LivePortrait_safetensors",
|
||||
local_dir=download_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
|
||||
model_config_path = os.path.join(
|
||||
script_directory, "liveportrait", "config", "models.yaml"
|
||||
)
|
||||
with open(model_config_path, "r") as file:
|
||||
model_config = yaml.safe_load(file)
|
||||
|
||||
feature_extractor_path = os.path.join(model_path, 'appearance_feature_extractor.safetensors')
|
||||
motion_extractor_path = os.path.join(model_path, 'motion_extractor.safetensors')
|
||||
warping_module_path = os.path.join(model_path, 'warping_module.safetensors')
|
||||
spade_generator_path = os.path.join(model_path, 'spade_generator.safetensors')
|
||||
stitching_retargeting_path = os.path.join(model_path, 'stitching_retargeting_module.safetensors')
|
||||
|
||||
feature_extractor_path = os.path.join(
|
||||
model_path, "appearance_feature_extractor.safetensors"
|
||||
)
|
||||
motion_extractor_path = os.path.join(model_path, "motion_extractor.safetensors")
|
||||
warping_module_path = os.path.join(model_path, "warping_module.safetensors")
|
||||
spade_generator_path = os.path.join(model_path, "spade_generator.safetensors")
|
||||
stitching_retargeting_path = os.path.join(
|
||||
model_path, "stitching_retargeting_module.safetensors"
|
||||
)
|
||||
|
||||
# init F
|
||||
model_params = model_config['model_params']['appearance_feature_extractor_params']
|
||||
self.appearance_feature_extractor = AppearanceFeatureExtractor(**model_params).to(device)
|
||||
self.appearance_feature_extractor.load_state_dict(comfy.utils.load_torch_file(feature_extractor_path))
|
||||
model_params = model_config["model_params"][
|
||||
"appearance_feature_extractor_params"
|
||||
]
|
||||
self.appearance_feature_extractor = AppearanceFeatureExtractor(
|
||||
**model_params
|
||||
).to(device)
|
||||
self.appearance_feature_extractor.load_state_dict(
|
||||
comfy.utils.load_torch_file(feature_extractor_path)
|
||||
)
|
||||
self.appearance_feature_extractor.eval()
|
||||
print('Load appearance_feature_extractor done.')
|
||||
log.info("Load appearance_feature_extractor done.")
|
||||
pbar.update(1)
|
||||
# init M
|
||||
model_params = model_config['model_params']['motion_extractor_params']
|
||||
model_params = model_config["model_params"]["motion_extractor_params"]
|
||||
self.motion_extractor = MotionExtractor(**model_params).to(device)
|
||||
self.motion_extractor.load_state_dict(comfy.utils.load_torch_file(motion_extractor_path))
|
||||
self.motion_extractor.load_state_dict(
|
||||
comfy.utils.load_torch_file(motion_extractor_path)
|
||||
)
|
||||
self.motion_extractor.eval()
|
||||
print('Load motion_extractor done.')
|
||||
log.info("Load motion_extractor done.")
|
||||
pbar.update(1)
|
||||
# init W
|
||||
model_params = model_config['model_params']['warping_module_params']
|
||||
model_params = model_config["model_params"]["warping_module_params"]
|
||||
self.warping_module = WarpingNetwork(**model_params).to(device)
|
||||
self.warping_module.load_state_dict(comfy.utils.load_torch_file(warping_module_path))
|
||||
self.warping_module.load_state_dict(
|
||||
comfy.utils.load_torch_file(warping_module_path)
|
||||
)
|
||||
self.warping_module.eval()
|
||||
print('Load warping_module done.')
|
||||
log.info("Load warping_module done.")
|
||||
pbar.update(1)
|
||||
# init G
|
||||
model_params = model_config['model_params']['spade_generator_params']
|
||||
model_params = model_config["model_params"]["spade_generator_params"]
|
||||
self.spade_generator = SPADEDecoder(**model_params).to(device)
|
||||
self.spade_generator.load_state_dict(comfy.utils.load_torch_file(spade_generator_path))
|
||||
self.spade_generator.load_state_dict(
|
||||
comfy.utils.load_torch_file(spade_generator_path)
|
||||
)
|
||||
self.spade_generator.eval()
|
||||
print('Load spade_generator done.')
|
||||
log.info("Load spade_generator done.")
|
||||
pbar.update(1)
|
||||
|
||||
def filter_checkpoint_for_model(checkpoint, prefix):
|
||||
"""Filter and adjust the checkpoint dictionary for a specific model based on the prefix."""
|
||||
# Create a new dictionary where keys are adjusted by removing the prefix and the model name
|
||||
filtered_checkpoint = {key.replace(prefix + "_module.", ""): value for key, value in checkpoint.items() if key.startswith(prefix)}
|
||||
filtered_checkpoint = {
|
||||
key.replace(prefix + "_module.", ""): value
|
||||
for key, value in checkpoint.items()
|
||||
if key.startswith(prefix)
|
||||
}
|
||||
return filtered_checkpoint
|
||||
|
||||
config = model_config['model_params']['stitching_retargeting_module_params']
|
||||
config = model_config["model_params"]["stitching_retargeting_module_params"]
|
||||
checkpoint = comfy.utils.load_torch_file(stitching_retargeting_path)
|
||||
|
||||
stitcher_prefix = 'retarget_shoulder'
|
||||
stitcher_prefix = "retarget_shoulder"
|
||||
stitcher_checkpoint = filter_checkpoint_for_model(checkpoint, stitcher_prefix)
|
||||
stitcher = StitchingRetargetingNetwork(**config.get('stitching'))
|
||||
stitcher = StitchingRetargetingNetwork(**config.get("stitching"))
|
||||
stitcher.load_state_dict(stitcher_checkpoint)
|
||||
stitcher = stitcher.to(device)
|
||||
stitcher.eval()
|
||||
|
||||
lip_prefix = 'retarget_mouth'
|
||||
lip_prefix = "retarget_mouth"
|
||||
lip_checkpoint = filter_checkpoint_for_model(checkpoint, lip_prefix)
|
||||
retargetor_lip = StitchingRetargetingNetwork(**config.get('lip'))
|
||||
retargetor_lip = StitchingRetargetingNetwork(**config.get("lip"))
|
||||
retargetor_lip.load_state_dict(lip_checkpoint)
|
||||
retargetor_lip = retargetor_lip.to(device)
|
||||
retargetor_lip.eval()
|
||||
|
||||
eye_prefix = 'retarget_eye'
|
||||
eye_prefix = "retarget_eye"
|
||||
eye_checkpoint = filter_checkpoint_for_model(checkpoint, eye_prefix)
|
||||
retargetor_eye = StitchingRetargetingNetwork(**config.get('eye'))
|
||||
retargetor_eye = StitchingRetargetingNetwork(**config.get("eye"))
|
||||
retargetor_eye.load_state_dict(eye_checkpoint)
|
||||
retargetor_eye = retargetor_eye.to(device)
|
||||
retargetor_eye.eval()
|
||||
print('Load stitching_retargeting_module done.')
|
||||
log.info("Load stitching_retargeting_module done.")
|
||||
|
||||
self.stich_retargeting_module = {
|
||||
'stitching': stitcher,
|
||||
'lip': retargetor_lip,
|
||||
'eye': retargetor_eye
|
||||
"stitching": stitcher,
|
||||
"lip": retargetor_lip,
|
||||
"eye": retargetor_eye,
|
||||
}
|
||||
|
||||
pipeline = LivePortraitPipeline(
|
||||
@@ -210,98 +218,526 @@ class DownloadAndLoadLivePortraitModels:
|
||||
self.spade_generator,
|
||||
self.stich_retargeting_module,
|
||||
InferenceConfig(
|
||||
device_id=device,
|
||||
flag_use_half_precision = True if precision == 'fp16' else False
|
||||
)
|
||||
device_id=device,
|
||||
flag_use_half_precision=True if precision == "fp16" else False,
|
||||
),
|
||||
)
|
||||
|
||||
return (pipeline,)
|
||||
|
||||
|
||||
class LivePortraitProcess:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"source_image": ("IMAGE",),
|
||||
"driving_images": ("IMAGE",),
|
||||
"lip_zero": ("BOOLEAN", {"default": False}),
|
||||
"lip_zero_threshold": ("FLOAT", {"default": 0.03, "min": 0.001, "max": 4.0, "step": 0.001}),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"delta_multiplier": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.001}),
|
||||
"mismatch_method": (
|
||||
[
|
||||
"constant",
|
||||
"cycle",
|
||||
"mirror",
|
||||
"cut"
|
||||
],
|
||||
{"default": "constant"},
|
||||
),
|
||||
|
||||
"relative_motion_mode": (
|
||||
[
|
||||
"relative",
|
||||
"source_video_smoothed",
|
||||
"relative_rotation_only",
|
||||
"single_frame",
|
||||
"off"
|
||||
],
|
||||
),
|
||||
"driving_smooth_observation_variance": ("FLOAT", {"default": 3e-6, "min": 1e-11, "max": 1e-2, "step": 1e-11}),
|
||||
},
|
||||
|
||||
"optional": {
|
||||
"opt_retargeting_info": ("RETARGETINGINFO", {"default": None}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"IMAGE",
|
||||
"LP_OUT",
|
||||
)
|
||||
RETURN_NAMES = (
|
||||
"cropped_image",
|
||||
"output",
|
||||
)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(
|
||||
self,
|
||||
source_image: torch.Tensor,
|
||||
driving_images: torch.Tensor,
|
||||
crop_info: dict,
|
||||
pipeline: LivePortraitPipeline,
|
||||
lip_zero: bool,
|
||||
lip_zero_threshold: float,
|
||||
stitching: bool,
|
||||
relative_motion_mode: str,
|
||||
driving_smooth_observation_variance: float,
|
||||
delta_multiplier: float = 1.0,
|
||||
mismatch_method: str = "constant",
|
||||
opt_retargeting_info: dict = None,
|
||||
):
|
||||
if driving_images.shape[0] < source_image.shape[0]:
|
||||
raise ValueError("The number of driving images should be larger than the number of source images.")
|
||||
source_np = (source_image * 255).byte().numpy()
|
||||
|
||||
if opt_retargeting_info is not None:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = opt_retargeting_info["eye_retargeting"]
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = (opt_retargeting_info["eyes_retargeting_multiplier"])
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = opt_retargeting_info["lip_retargeting"]
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = (opt_retargeting_info["lip_retargeting_multiplier"])
|
||||
driving_landmarks = opt_retargeting_info["driving_landmarks"]
|
||||
else:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = False
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = 1.0
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = False
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = 1.0
|
||||
driving_landmarks = None
|
||||
|
||||
pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
|
||||
pipeline.live_portrait_wrapper.cfg.lip_zero_threshold = lip_zero_threshold
|
||||
|
||||
if relative_motion_mode != "off":
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = True
|
||||
else:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = False
|
||||
|
||||
if lip_zero and opt_retargeting_info is not None:
|
||||
log.warning("Warning: lip_zero only has an effect with lip or eye retargeting")
|
||||
|
||||
if driving_images.shape[1] != 256 or driving_images.shape[2] != 256:
|
||||
driving_images_256 = comfy.utils.common_upscale(driving_images.permute(0, 3, 1, 2), 256, 256, "lanczos", "disabled")
|
||||
else:
|
||||
driving_images_256 = driving_images.permute(0, 3, 1, 2)
|
||||
|
||||
if pipeline.live_portrait_wrapper.cfg.flag_use_half_precision:
|
||||
driving_images_256 = driving_images_256.to(torch.float16)
|
||||
|
||||
out = pipeline.execute(
|
||||
source_np,
|
||||
driving_images_256,
|
||||
crop_info,
|
||||
driving_landmarks,
|
||||
delta_multiplier,
|
||||
relative_motion_mode,
|
||||
driving_smooth_observation_variance,
|
||||
mismatch_method
|
||||
)
|
||||
|
||||
total_frames = len(out["out_list"])
|
||||
|
||||
if total_frames > 1:
|
||||
cropped_image_list = []
|
||||
for i in (range(total_frames)):
|
||||
if not out["out_list"][i]:
|
||||
cropped_image_list.append(torch.zeros(1, 512, 512, 3, dtype=torch.float32, device = "cpu"))
|
||||
else:
|
||||
cropped_image = torch.clamp(out["out_list"][i]["out"], 0, 1).permute(0, 2, 3, 1).cpu()
|
||||
cropped_image_list.append(cropped_image)
|
||||
|
||||
cropped_out_tensors = torch.cat(cropped_image_list, dim=0)
|
||||
else:
|
||||
cropped_out_tensors = torch.clamp(out["out_list"][0]["out"], 0, 1).permute(0, 2, 3, 1)
|
||||
|
||||
return (cropped_out_tensors, out,)
|
||||
|
||||
class LivePortraitComposite:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"source_image": ("IMAGE",),
|
||||
"cropped_image": ("IMAGE",),
|
||||
"liveportrait_out": ("LP_OUT", ),
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK", {"default": None}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"IMAGE",
|
||||
"MASK",
|
||||
)
|
||||
RETURN_NAMES = (
|
||||
"full_images",
|
||||
"mask",
|
||||
)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, source_image, cropped_image, liveportrait_out, mask=None):
|
||||
mm.soft_empty_cache()
|
||||
device = mm.get_torch_device()
|
||||
if mm.is_device_mps(device):
|
||||
device = torch.device('cpu') #this function returns NaNs on MPS, defaulting to CPU
|
||||
|
||||
B, H, W, C = source_image.shape
|
||||
source_image = source_image.permute(0, 3, 1, 2) # B,H,W,C -> B,C,H,W
|
||||
cropped_image = cropped_image.permute(0, 3, 1, 2)
|
||||
|
||||
if mask is not None:
|
||||
crop_mask = mask.unsqueeze(-1).expand(-1, -1, -1, 3)
|
||||
else:
|
||||
log.info("Using default mask template")
|
||||
crop_mask = cv2.imread(os.path.join(script_directory, "liveportrait", "utils", "resources", "mask_template.png"), cv2.IMREAD_COLOR)
|
||||
crop_mask = torch.from_numpy(crop_mask)
|
||||
crop_mask = crop_mask.unsqueeze(0).float() / 255.0
|
||||
|
||||
crop_info = liveportrait_out["crop_info"]
|
||||
composited_image_list = []
|
||||
out_mask_list = []
|
||||
|
||||
total_frames = len(liveportrait_out["out_list"])
|
||||
log.info(f"Total frames: {total_frames}")
|
||||
|
||||
pbar = comfy.utils.ProgressBar(total_frames)
|
||||
for i in tqdm(range(total_frames), desc='Compositing..', total=total_frames):
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
|
||||
if liveportrait_out["mismatch_method"] == "cut":
|
||||
source_frame = source_image[safe_index].unsqueeze(0).to(device)
|
||||
else:
|
||||
source_frame = _get_source_frame(source_image, i, liveportrait_out["mismatch_method"]).unsqueeze(0).to(device)
|
||||
|
||||
if not liveportrait_out["out_list"][i]:
|
||||
composited_image_list.append(source_frame)
|
||||
out_mask_list.append(torch.zeros((1, 3, H, W), device=device))
|
||||
else:
|
||||
cropped_image = torch.clamp(liveportrait_out["out_list"][i]["out"], 0, 1).permute(0, 2, 3, 1)
|
||||
|
||||
# Transform and blend
|
||||
cropped_image_to_original = _transform_img_kornia(
|
||||
cropped_image,
|
||||
crop_info["crop_info_list"][safe_index]["M_c2o"],
|
||||
dsize=(W, H),
|
||||
device=device
|
||||
)
|
||||
|
||||
mask_ori = _transform_img_kornia(
|
||||
crop_mask,
|
||||
crop_info["crop_info_list"][safe_index]["M_c2o"],
|
||||
dsize=(W, H),
|
||||
device=device
|
||||
)
|
||||
|
||||
cropped_image_to_original_blend = torch.clip(
|
||||
mask_ori * cropped_image_to_original + (1 - mask_ori) * source_frame, 0, 1
|
||||
)
|
||||
|
||||
composited_image_list.append(cropped_image_to_original_blend)
|
||||
out_mask_list.append(mask_ori)
|
||||
pbar.update(1)
|
||||
|
||||
full_tensors_out = torch.cat(composited_image_list, dim=0)
|
||||
full_tensors_out = full_tensors_out.permute(0, 2, 3, 1)
|
||||
|
||||
mask_tensors_out = torch.cat(out_mask_list, dim=0)
|
||||
mask_tensors_out = mask_tensors_out[:, 0, :, :]
|
||||
|
||||
return (
|
||||
full_tensors_out.cpu().float(),
|
||||
mask_tensors_out.cpu().float()
|
||||
)
|
||||
|
||||
def _get_source_frame(source, idx, method):
|
||||
if source.shape[0] == 1:
|
||||
return source[0]
|
||||
|
||||
if method == "constant":
|
||||
return source[min(idx, source.shape[0] - 1)]
|
||||
elif method == "cycle":
|
||||
return source[idx % source.shape[0]]
|
||||
elif method == "mirror":
|
||||
cycle_length = 2 * source.shape[0] - 2
|
||||
mirror_idx = idx % cycle_length
|
||||
if mirror_idx >= source.shape[0]:
|
||||
mirror_idx = cycle_length - mirror_idx
|
||||
return source[mirror_idx]
|
||||
|
||||
class LivePortraitLoadCropper:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"onnx_device": (
|
||||
['CPU', 'CUDA', 'ROCM', 'CoreML'], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": True})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LPCROPPER",)
|
||||
RETURN_NAMES = ("cropper",)
|
||||
FUNCTION = "crop"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def crop(self, onnx_device, keep_model_loaded):
|
||||
cropper_init_config = {
|
||||
'keep_model_loaded': keep_model_loaded,
|
||||
'onnx_device': onnx_device
|
||||
}
|
||||
|
||||
if not hasattr(self, 'cropper') or self.cropper is None or self.current_config != cropper_init_config:
|
||||
self.current_config = cropper_init_config
|
||||
self.cropper = Cropper(**cropper_init_config)
|
||||
|
||||
return (self.cropper,)
|
||||
|
||||
class LivePortraitCropper:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"cropper": ("LPCROPPER",),
|
||||
"source_image": ("IMAGE",),
|
||||
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
|
||||
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}),
|
||||
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
||||
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.001}),
|
||||
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.001}),
|
||||
"face_index": ("INT", {"default": 0, "min": 0, "max": 100}),
|
||||
"face_index_order": (
|
||||
[
|
||||
'large-small',
|
||||
'left-right',
|
||||
'right-left',
|
||||
'top-bottom',
|
||||
'bottom-top',
|
||||
'small-large',
|
||||
'distance-from-retarget-face'
|
||||
],
|
||||
),
|
||||
"rotate": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "CROPINFO",)
|
||||
RETURN_NAMES = ("cropped_image", "crop_info",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, pipeline, cropper, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, face_index_order, rotate):
|
||||
source_image_np = (source_image * 255).byte().numpy()
|
||||
|
||||
# Initialize lists
|
||||
crop_info_list = []
|
||||
cropped_images_list = []
|
||||
source_info = []
|
||||
source_rot_list = []
|
||||
f_s_list = []
|
||||
x_s_list = []
|
||||
|
||||
# Initialize a progress bar for the combined operation
|
||||
pbar = comfy.utils.ProgressBar(len(source_image_np))
|
||||
for i in tqdm(range(len(source_image_np)), desc='Detecting, cropping, and processing..', total=len(source_image_np)):
|
||||
# Cropping operation
|
||||
crop_info = cropper.crop_single_image(source_image_np[i], dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate)
|
||||
|
||||
# Processing source images
|
||||
if crop_info:
|
||||
crop_info_list.append(crop_info)
|
||||
|
||||
cropped_image = crop_info['img_crop_256x256']
|
||||
cropped_images_list.append(cropped_image)
|
||||
|
||||
I_s = pipeline.live_portrait_wrapper.prepare_source(cropped_image)
|
||||
|
||||
x_s_info = pipeline.live_portrait_wrapper.get_kp_info(I_s)
|
||||
source_info.append(x_s_info)
|
||||
|
||||
x_s = pipeline.live_portrait_wrapper.transform_keypoint(x_s_info)
|
||||
x_s_list.append(x_s)
|
||||
|
||||
R_s = get_rotation_matrix(x_s_info["pitch"], x_s_info["yaw"], x_s_info["roll"])
|
||||
source_rot_list.append(R_s)
|
||||
|
||||
f_s = pipeline.live_portrait_wrapper.extract_feature_3d(I_s)
|
||||
f_s_list.append(f_s)
|
||||
|
||||
else:
|
||||
log.warning(f"Warning: No face detected on frame {str(i)}, skipping")
|
||||
cropped_image = np.zeros((256, 256, 3), dtype=np.uint8)
|
||||
crop_info_list.append(None)
|
||||
f_s_list.append(None)
|
||||
x_s_list.append(None)
|
||||
source_info.append(None)
|
||||
source_rot_list.append(None)
|
||||
|
||||
# Update progress bar
|
||||
pbar.update(1)
|
||||
|
||||
cropped_tensors_out = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in cropped_images_list])
|
||||
/ 255
|
||||
)
|
||||
|
||||
crop_info_dict = {
|
||||
'crop_info_list': crop_info_list,
|
||||
'source_rot_list': source_rot_list,
|
||||
'f_s_list': f_s_list,
|
||||
'x_s_list': x_s_list,
|
||||
'source_info': source_info
|
||||
}
|
||||
|
||||
return (cropped_tensors_out, crop_info_dict)
|
||||
|
||||
class LivePortraitRetargeting:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"driving_crop_info": ("CROPINFO", {"default": []}),
|
||||
"eye_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"relative": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"onnx_device": (
|
||||
[
|
||||
'CPU',
|
||||
'CUDA',
|
||||
], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("cropped_images", "full_images",)
|
||||
RETURN_TYPES = ("RETARGETINGINFO",)
|
||||
RETURN_NAMES = ("retargeting_info",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, source_image, driving_images, dsize, scale, vx_ratio, vy_ratio, pipeline,
|
||||
lip_zero, eye_retargeting, lip_retargeting, stitching, relative, eyes_retargeting_multiplier, lip_retargeting_multiplier, onnx_device='CUDA'):
|
||||
source_image_np = (source_image * 255).byte().numpy()
|
||||
driving_images_np = (driving_images * 255).byte().numpy()
|
||||
def process(self, driving_crop_info, eye_retargeting, eyes_retargeting_multiplier, lip_retargeting, lip_retargeting_multiplier):
|
||||
|
||||
crop_cfg = CropConfig(
|
||||
dsize = dsize,
|
||||
scale = scale,
|
||||
vx_ratio = vx_ratio,
|
||||
vy_ratio = vy_ratio,
|
||||
)
|
||||
driving_landmarks = []
|
||||
for crop in driving_crop_info["crop_info_list"]:
|
||||
driving_landmarks.append(crop['lmk_crop'])
|
||||
|
||||
retargeting_info = {
|
||||
'eye_retargeting': eye_retargeting,
|
||||
'eyes_retargeting_multiplier': eyes_retargeting_multiplier,
|
||||
'lip_retargeting': lip_retargeting,
|
||||
'lip_retargeting_multiplier': lip_retargeting_multiplier,
|
||||
'driving_landmarks': driving_landmarks
|
||||
}
|
||||
|
||||
return (retargeting_info,)
|
||||
|
||||
|
||||
class KeypointsToImage:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": []}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("keypoints_image",)
|
||||
FUNCTION = "drawkeypoints"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def drawkeypoints(self, crop_info):
|
||||
height, width = crop_info["crop_info_list"][0]['input_image_size']
|
||||
keypoints_img_list = []
|
||||
pbar = comfy.utils.ProgressBar(len(crop_info))
|
||||
for crop in crop_info["crop_info_list"]:
|
||||
if crop:
|
||||
keypoints = crop['lmk_crop'].copy()
|
||||
# Draw each landmark as a circle
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in keypoints:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
else:
|
||||
keypoints_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
keypoints_img_list.append(keypoints_image)
|
||||
pbar.update(1)
|
||||
|
||||
keypoints_img_tensor = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in keypoints_img_list]) / 255).float()
|
||||
|
||||
|
||||
return (keypoints_img_tensor,)
|
||||
|
||||
class KeypointScaler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"offset_x": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
"offset_y": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CROPINFO", "IMAGE",)
|
||||
RETURN_NAMES = ("crop_info", "keypoints_image",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, crop_info, offset_x, offset_y, scale):
|
||||
|
||||
keypoints = crop_info['crop_info']['lmk_crop'].copy()
|
||||
|
||||
# Create an offset array
|
||||
# Calculate the centroid of the keypoints
|
||||
centroid = keypoints.mean(axis=0)
|
||||
|
||||
# Translate keypoints to origin by subtracting the centroid
|
||||
translated_keypoints = keypoints - centroid
|
||||
|
||||
# Scale the translated keypoints
|
||||
scaled_keypoints = translated_keypoints * scale
|
||||
|
||||
# Translate scaled keypoints back to original position and then apply the offset
|
||||
final_keypoints = scaled_keypoints + centroid + np.array([offset_x, offset_y])
|
||||
|
||||
crop_info['crop_info']['lmk_crop'] = final_keypoints #fix this
|
||||
|
||||
# Draw each landmark as a circle
|
||||
width, height = 512, 512
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in final_keypoints:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
keypoints_image_tensor = torch.from_numpy(keypoints_image) / 255
|
||||
keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float()
|
||||
|
||||
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
|
||||
pipeline.cropper = cropper
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = eye_retargeting
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = eyes_retargeting_multiplier
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = lip_retargeting
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = lip_retargeting_multiplier
|
||||
pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = relative
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
|
||||
|
||||
cropped_out_list = []
|
||||
full_out_list = []
|
||||
for img in source_image_np:
|
||||
cropped_frames, full_frame = pipeline.execute(img, driving_images_np)
|
||||
cropped_tensors = [torch.from_numpy(np_array) for np_array in cropped_frames]
|
||||
cropped_tensors_out = torch.stack(cropped_tensors) / 255
|
||||
cropped_tensors_out = cropped_tensors_out.cpu().float()
|
||||
|
||||
full_tensors = [torch.from_numpy(np_array) for np_array in full_frame]
|
||||
full_tensors_out = torch.stack(full_tensors) / 255
|
||||
full_tensors_out = full_tensors_out.cpu().float()
|
||||
|
||||
cropped_out_list.append(cropped_tensors_out)
|
||||
full_out_list.append(full_tensors_out)
|
||||
|
||||
cropped_tensors_out = torch.cat(cropped_out_list, dim=0)
|
||||
full_tensors_out = torch.cat(full_out_list, dim=0)
|
||||
|
||||
return (cropped_tensors_out, full_tensors_out)
|
||||
return (crop_info, keypoints_image_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
||||
"LivePortraitProcess": LivePortraitProcess,
|
||||
"LivePortraitCropper": LivePortraitCropper,
|
||||
"LivePortraitRetargeting": LivePortraitRetargeting,
|
||||
#"KeypointScaler": KeypointScaler,
|
||||
"KeypointsToImage": KeypointsToImage,
|
||||
"LivePortraitLoadCropper": LivePortraitLoadCropper,
|
||||
"LivePortraitComposite": LivePortraitComposite,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
||||
"LivePortraitProcess": "LivePortraitProcess",
|
||||
"LivePortraitCropper": "LivePortraitCropper",
|
||||
"LivePortraitRetargeting": "LivePortraitRetargeting",
|
||||
#"KeypointScaler": "KeypointScaler",
|
||||
"KeypointsToImage": "LivePortrait KeypointsToImage",
|
||||
"LivePortraitLoadCropper": "LivePortrait LoadCropper",
|
||||
"LivePortraitComposite": "LivePortrait Composite",
|
||||
}
|
||||
@@ -1,7 +1,12 @@
|
||||
# ComfyUI nodes to use [LivePortrait](https://github.com/KwaiVGI/LivePortrait)
|
||||
|
||||
## Update
|
||||
|
||||
Rework of almost the whole thing that's been in develop is now merged into main, this means old workflows will not work, but everything should be faster and there's lots of new features.
|
||||
For legacy purposes the old main branch is moved to the legacy -branch
|
||||
|
||||
|
||||
|
||||
https://github.com/kijai/ComfyUI-LivePortrait/assets/40791699/e55e10f6-af61-4d73-b162-af29eb847516
|
||||
|
||||
|
||||
I have converted all the pickle files to safetensors: https://huggingface.co/Kijai/LivePortrait_safetensors/tree/main
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
pyyaml
|
||||
numpy
|
||||
opencv-python
|
||||
onnxruntime-gpu
|
||||
pykalman
|
||||
tensorrt
|
||||
pycuda
|
||||
ctypes
|
||||
+2
-2
@@ -1,5 +1,5 @@
|
||||
pyyaml
|
||||
numpy
|
||||
opencv-python
|
||||
rich
|
||||
onnxruntime-gpu
|
||||
onnxruntime-gpu
|
||||
pykalman
|
||||
Reference in New Issue
Block a user