Squashed commit of the following:

commit b608558b9e
Merge: ad29b02 dd205ab
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 24 17:56:47 2024 +0300

    Merge branch 'develop' of https://github.com/kijai/ComfyUI-LivePortraitKJ into develop

commit ad29b02bc1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 24 17:56:46 2024 +0300

    update workflows

commit dd205ab4a4
Author: Jukka Seppänen <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 24 17:54:47 2024 +0300

    Update readme.md

commit ba0886a905
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 24 16:09:26 2024 +0300

    fix running without insightface installed

commit 068ab2c280
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 24 03:57:59 2024 +0300

    Add MediaPipe as alternative face detector

commit 6261f4e474
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 23 22:58:33 2024 +0300

    cleanup, memory fixes

commit 46675b2016
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 23 21:13:15 2024 +0300

    update workflows, cleanup

commit 806263dd25
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 20:39:43 2024 +0300

    cleanup, fixes

commit ac89dc1e2f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 19:19:46 2024 +0300

    fix no face frame skip

commit 27d745b53e
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 17:52:34 2024 +0300

    add other examples

commit 052762578c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 17:45:49 2024 +0300

    Update readme.md

commit e825c51c87
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 16:21:35 2024 +0300

    separate composition to it's own node

commit 177b324fcd
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 01:02:56 2024 +0300

    Update live_portrait_pipeline.py

commit 5c03bd8439
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 22 00:57:44 2024 +0300

    MPS fallbacks

commit ef5ff7075f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Jul 21 20:35:33 2024 +0300

    Update requirements.txt

commit 92fad03ee5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Jul 21 20:20:46 2024 +0300

    restructure a bit for more caching

commit 4cefac79b8
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Jul 21 19:42:22 2024 +0300

    Add single_frame mode for webcam

commit 5e3c92d55c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Jul 21 19:20:52 2024 +0300

    restructuring, video smoothing

commit cc0501a2db
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Jul 21 13:21:26 2024 +0300

    flag_relative_rotation_only

commit 3dc822fd2f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Jul 20 20:24:15 2024 +0300

    to use GPU for pasteback

commit 697b9a78e6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Jul 20 17:39:44 2024 +0300

    Restructure nodes, skip frames with no face detect

commit 2a7bd6116f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 10 17:14:43 2024 +0300

    Update nodes.py

commit a7d09f5d49
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 10 16:33:31 2024 +0300

    example workflow

commit 8e85d5b96d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 10 16:06:01 2024 +0300

    some optimizations

commit 30989a9d37
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Jul 10 01:17:10 2024 +0300

    Update nodes.py

commit eecf645603
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 23:02:47 2024 +0300

    rotate option for cropper

commit 1b080706df
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 22:46:51 2024 +0300

    Update nodes.py

commit 336f3f7c23
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 22:19:56 2024 +0300

    add cut method

commit f27e1cca13
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 22:11:19 2024 +0300

    remove nearest option

commit 86e91a6e9d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 22:06:56 2024 +0300

    better error for retargeting

commit 92529f7ca8
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 21:59:32 2024 +0300

    cleanup

commit c0959056ae
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 21:49:54 2024 +0300

    eye/lip retargeting fixes

commit 2e40fe3820
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 21:23:10 2024 +0300

    Update live_portrait_pipeline.py

commit 0a5e187637
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 21:15:17 2024 +0300

    keep Cropper in memory

commit 4e19dbd6d1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 20:19:53 2024 +0300

    Do video cropping on the cropped node too

commit 9c190804a7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 19:02:27 2024 +0300

    big cleanup

commit d9ca40e1d6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 15:08:17 2024 +0300

    logging

commit b68cf8788c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 14:35:51 2024 +0300

    fix warning

commit e702b26895
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 14:31:22 2024 +0300

    Update cropper.py

commit c21705edb5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 14:25:39 2024 +0300

    Don't draw keypoints for every frame by default

commit a284bb52b2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 14:04:56 2024 +0300

    Bring back mismatch_method selection

commit 9884aac18a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 14:00:43 2024 +0300

    Fix eye/lip retargeting

commit f8aada81db
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 13:00:03 2024 +0300

    face_index selection

commit 6735771664
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 11:51:38 2024 +0300

    tqdm progress bars

commit 857ddbc6d7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 02:54:27 2024 +0300

    skip autocast if not needed for mps

commit a6edcda97d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 01:37:05 2024 +0300

    output masks

commit ca01d706d0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 00:57:13 2024 +0300

    custom mask support

commit 0dc9a8a695
Merge: ee7d5b4 ba6b3f5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 00:13:07 2024 +0300

    Merge branch 'add_video_source' into develop

commit ba6b3f5f68
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 23:09:56 2024 +0200

    bring KJ edits

commit ee7d5b4241
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Jul 9 00:08:10 2024 +0300

    revert this for compatibility

commit 03df9f35cd
Merge: ec6b5c8 8509d9a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 23:55:22 2024 +0300

    Merge branch 'add_video_source' into develop

commit ec6b5c8c85
Merge: 6f9dba7 e724da1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 23:52:58 2024 +0300

    calc_combined_eye_ratio

commit 8509d9a551
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 22:52:48 2024 +0200

    remove unused imports

commit e724da1161
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 21:38:27 2024 +0200

    fix relative mode

    use R_d_0 instead of source

commit 68d0ddf72a
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 21:22:50 2024 +0200

    remove reference frame attempt

    also use batches for driving when either retargetting is enabled

commit 6f9dba7777
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 20:50:44 2024 +0300

    fixes

commit 811ca557fb
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 20:31:00 2024 +0300

    more

commit 6d790bdcc3
Merge: ef8b426 eb5fddf
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 20:30:35 2024 +0300

    Merge branch 'add_video_source' into develop

commit ef8b4263b4
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jul 8 20:21:45 2024 +0300

    separating functions to nodes

commit eb5fddf4de
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 19:10:53 2024 +0200

    fix issues from merge

commit 9c7db3c59a
Merge: bf3410c 1f28e12
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 19:09:02 2024 +0200

    Merge branch 'main' into add_video_source

commit bf3410cd0d
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 19:04:42 2024 +0200

    trying reference frame

commit 24c65627db
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 19:03:28 2024 +0200

    local updates before merging main

commit 72bb6910e9
Author: Mel Massadian <mel@melmassadian.com>
Date:   Mon Jul 8 16:48:15 2024 +0200

    initial

    too much diff due to formatting
This commit is contained in:
kijai
2024-07-24 18:01:19 +03:00
parent cae686d921
commit 3d195208db
27 changed files with 7995 additions and 1270 deletions
File diff suppressed because it is too large Load Diff
@@ -1,70 +1,96 @@
{
"last_node_id": 31,
"last_link_id": 68,
"last_node_id": 208,
"last_link_id": 480,
"nodes": [
{
"id": 4,
"type": "LoadImage",
"id": 204,
"type": "LivePortraitLoadMediaPipeCropper",
"pos": [
138,
323
-1059,
-767
],
"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
478
],
"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": "LivePortraitLoadMediaPipeCropper"
},
"widgets_values": [
"oldman.jpg",
"image"
"CPU",
true
]
},
{
"id": 19,
"id": 1,
"type": "DownloadAndLoadLivePortraitModels",
"pos": [
-1046,
-904
],
"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
-715,
-617
],
"size": {
"0": 315,
"1": 242
},
"flags": {},
"order": 4,
"order": 6,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 30
"link": 466
},
{
"name": "get_image_size",
"type": "IMAGE",
"link": 68
"link": null
},
{
"name": "width_input",
@@ -88,7 +114,7 @@
"name": "IMAGE",
"type": "IMAGE",
"links": [
32
434
],
"shape": 3,
"slot_index": 0
@@ -112,290 +138,282 @@
"widgets_values": [
512,
512,
"nearest-exact",
false,
"lanczos",
true,
2,
0,
0
]
},
{
"id": 18,
"type": "ImageConcatMulti",
"id": 78,
"type": "GetImageSizeAndCount",
"pos": [
860,
679
-364,
-619
],
"size": {
"0": 210,
"1": 150
"1": 86
},
"flags": {},
"order": 5,
"order": 7,
"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": [
"auto"
]
},
{
"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": "384 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,
"type": "LivePortraitProcess",
"id": 207,
"type": "Note",
"pos": [
500,
249
-850,
-202
],
"size": [
230.3331222759989,
105.34622315856666
],
"flags": {},
"order": 2,
"mode": 0,
"properties": {
"text": ""
},
"widgets_values": [
"Example live inputs, direct webcam capture using cv2 or screencapture using mss. Both are about same speed."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 206,
"type": "Screencap_mss",
"pos": [
-568,
-34
],
"size": {
"0": 367.79998779296875,
"1": 362
"0": 315,
"1": 178
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "Screencap_mss"
},
"widgets_values": [
0,
0,
512,
512,
1,
0.1
]
},
{
"id": 198,
"type": "PreviewImage",
"pos": [
413,
-826
],
"size": {
"0": 521.2196044921875,
"1": 566.1187133789062
},
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 470
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 196,
"type": "LoadImage",
"pos": [
-1058,
-623
],
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 4,
"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": [
"Mona-Lisa-oil-wood-panel-Leonardo-da.webp",
"image"
]
},
{
"id": 205,
"type": "WebcamCaptureCV2",
"pos": [
-577,
-283
],
"size": {
"0": 315,
"1": 178
},
"flags": {},
"order": 5,
"mode": 0,
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": [
479
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WebcamCaptureCV2"
},
"widgets_values": [
0,
0,
512,
512,
0,
false
]
},
{
"id": 190,
"type": "LivePortraitProcess",
"pos": [
-79,
-552
],
"size": {
"0": 430.8000183105469,
"1": 282
},
"flags": {},
"order": 9,
"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": 479,
"slot_index": 3
},
{
"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 +421,160 @@
"properties": {
"Node name for S&R": "LivePortraitProcess"
},
"widgets_values": [
false,
0.03,
true,
1,
"constant",
"single_frame",
0.000003
]
},
{
"id": 189,
"type": "LivePortraitCropper",
"pos": [
-73,
-876
],
"size": {
"0": 330,
"1": 242
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "pipeline",
"type": "LIVEPORTRAITPIPE",
"link": 446,
"slot_index": 0
},
{
"name": "cropper",
"type": "LPCROPPER",
"link": 478,
"slot_index": 1
},
{
"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,
2.34,
0.099,
0.148,
0,
-0.11,
true,
false,
1,
false,
1,
true,
true,
"CPU"
"large-small",
false
]
}
],
"links": [
[
30,
8,
434,
165,
0,
19,
78,
0,
"IMAGE"
],
[
32,
19,
445,
78,
0,
18,
0,
"IMAGE"
],
[
58,
1,
0,
30,
0,
"LIVEPORTRAITPIPE"
],
[
59,
4,
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,
475,
78,
0,
190,
2,
"IMAGE"
],
[
478,
204,
0,
189,
1,
"LPCROPPER"
],
[
479,
205,
0,
190,
3,
"IMAGE"
]
],
@@ -489,10 +582,10 @@
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"scale": 0.7513148009015781,
"offset": {
"0": 173.40487670898438,
"1": -0.9636010527610779
"0": 1468.4081568988054,
"1": 1224.8414164288351
}
}
},
File diff suppressed because it is too large Load Diff
-44
View File
@@ -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"
-18
View File
@@ -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
-8
View File
@@ -29,21 +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
flag_do_rot: bool = True # whether to conduct the rotation when flag_do_crop is True
+234 -142
View File
@@ -4,184 +4,276 @@
Pipeline of LivePortrait
"""
import cv2
import numpy as np
import os.path as osp
from tqdm import tqdm
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
import comfy.model_management as mm
import gc
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 execute(
self, 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
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)
############################################
######## 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)
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 tqdm(range(n_frames), desc='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']
source_images_num = len(crop_info["crop_info_list"])
if mismatch_method == "cut" or relative_motion_mode == "source_video_smoothed":
total_frames = source_images_num
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, source_images_num - 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_images_num), desc='Smoothing...', total=source_images_num):
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
delta_new = delta_new * delta_multiplier
t_new[..., 2].fill_(0) # zero tz
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)
#with eye/lip retargeting
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)
if relative_motion_mode != "off": # 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
)
)
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)
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)
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)
out_dict = {
"out_list": out_list,
"crop_info": crop_info,
"mismatch_method": mismatch_method,
}
return I_p_lst, I_p_paste_lst
return out_dict
+13 -72
View File
@@ -32,11 +32,6 @@ 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)
def prepare_source(self, img: np.ndarray) -> torch.Tensor:
""" construct the input as standard
img: HxWx3, uint8, 256x256
@@ -58,24 +53,6 @@ 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
@@ -193,35 +170,6 @@ class LivePortraitWrapper(object):
return delta
def retarget_keypoints(self, frame_idx, num_keypoints, input_eye_ratios, input_lip_ratios, source_landmarks, portrait_wrapper, kp_source, driving_transformed_kp):
# TODO: GPT style, refactor it...
if self.cfg.flag_eye_retargeting:
print("Retargeting eye...")
# ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
eye_delta = compute_eye_delta(frame_idx, input_eye_ratios, source_landmarks, portrait_wrapper, kp_source)
else:
# α_eyes = 0
eye_delta = None
if self.cfg.flag_lip_retargeting:
print("Retargeting lip...")
# ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
lip_delta = compute_lip_delta(frame_idx, input_lip_ratios, source_landmarks, portrait_wrapper, kp_source)
else:
# α_lip = 0
lip_delta = None
if self.cfg.flag_relative: # use x_s
new_driving_kp = kp_source + \
(eye_delta.reshape(-1, num_keypoints, 3) if eye_delta is not None else 0) + \
(lip_delta.reshape(-1, num_keypoints, 3) if lip_delta is not None else 0)
else: # use x_d,i
new_driving_kp = driving_transformed_kp + \
(eye_delta.reshape(-1, num_keypoints, 3) if eye_delta is not None else 0) + \
(lip_delta.reshape(-1, num_keypoints, 3) if lip_delta is not None else 0)
return new_driving_kp
def stitch(self, kp_source: torch.Tensor, kp_driving: torch.Tensor) -> torch.Tensor:
"""
kp_source: BxNx3
@@ -264,7 +212,7 @@ class LivePortraitWrapper(object):
kp_source: BxNx3
kp_driving: BxNx3
"""
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i))
# 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)
@@ -272,23 +220,14 @@ class LivePortraitWrapper(object):
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()
for k, v in ret_dct.items():
if isinstance(v, torch.Tensor):
ret_dct[k] = v.cpu()
if self.cfg.flag_use_half_precision:
ret_dct[k] = ret_dct[k].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
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
def calc_retargeting_ratio(self, source_lmk, driving_lmk_lst):
input_eye_ratio_lst = []
input_lip_ratio_lst = []
@@ -301,17 +240,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)
+8 -2
View File
@@ -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)
+5 -1
View File
@@ -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
+5 -1
View File
@@ -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:
-65
View File
@@ -1,65 +0,0 @@
# coding: utf-8
"""
Make video template
"""
import os
import cv2
import numpy as np
import pickle
from tqdm import tqdm
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 tqdm(range(n_frames), desc='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}")
+111 -24
View File
@@ -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
@@ -39,6 +74,23 @@ def _transform_pts(pts, M):
return pts @ M[:2, :2].T + M[:2, 2]
def parse_pt2_from_pt478(pt478, use_lip=True):
"""
parsing the 2 points according to the 101 points, which cancels the roll
"""
# the former version use the eye center, but it is not robust, now use interpolation
pt_left_eye = pt478[468] # left eye center
pt_right_eye = pt478[473] # right eye center
if use_lip:
# use lip
pt_center_eye = (pt_left_eye + pt_right_eye) / 2
pt_center_lip = pt478[14]
pt2 = np.stack([pt_center_eye, pt_center_lip], axis=0)
else:
pt2 = np.stack([pt_left_eye, pt_right_eye], axis=0)
return pt2
def parse_pt2_from_pt101(pt101, use_lip=True):
"""
parsing the 2 points according to the 101 points, which cancels the roll
@@ -89,29 +141,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
@@ -145,9 +228,13 @@ def parse_pt2_from_pt_x(pts, use_lip=True):
pt2 = parse_pt2_from_pt5(pts, use_lip=use_lip)
elif pts.shape[0] == 203:
pt2 = parse_pt2_from_pt203(pts, use_lip=use_lip)
elif pts.shape[0] == 478:
pt2 = parse_pt2_from_pt478(pts, use_lip=use_lip)
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 +437,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:
@@ -379,15 +468,13 @@ def crop_image(img, pts: np.ndarray, **kwargs):
ret_dct = {
'M_o2c': M_o2c, # from the original image to the cropped image 3x3
'M_c2o': M_c2o, # from the cropped image to the original image 3x3
'img_crop': img_crop, # the cropped image
'pt_crop': pt_crop, # the landmarks of the cropped image
}
return ret_dct
return ret_dct, img_crop
def average_bbox_lst(bbox_lst):
if len(bbox_lst) == 0:
return None
bbox_arr = np.array(bbox_lst)
return np.mean(bbox_arr, axis=0).tolist()
+76 -87
View File
@@ -1,49 +1,39 @@
# 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
frame_rgb_crop_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame crop list
class Cropper(object):
def __init__(self, provider, **kwargs) -> None:
class CropperInsightFace(object):
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
)
self.landmark_runner.warmup()
from .face_analysis_diy import FaceAnalysisDIY
self.face_analysis_wrapper = FaceAnalysisDIY(
name='buffalo_l',
root=os.path.join(folder_paths.models_dir, 'insightface'),
@@ -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,85 @@ 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
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(
ret_dct, image_crop = 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)
cropped_image_256 = cv2.resize(image_crop, (256, 256), interpolation=cv2.INTER_AREA)
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
return ret_dct, cropped_image_256
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
class CropperMediaPipe(object):
def __init__(self, **kwargs) -> None:
device_id = kwargs.get('device_id', 0)
provider = kwargs.get('onnx_device', 'CPU')
self.landmark_runner = LandmarkRunner(
ckpt_path=os.path.join(folder_paths.models_dir, 'liveportrait', 'landmark.onnx'),
onnx_provider=provider,
device_id=device_id
)
self.landmark_runner.warmup()
from ...media_pipe.mp_utils import LMKExtractor
self.lmk_extractor = LMKExtractor()
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
def crop_single_image(self, img_rgb, dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate):
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)
face_result = self.lmk_extractor(img_rgb)
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']
if face_result is None:
ret_dct = {}
cropped_image_256 = None
return ret_dct, cropped_image_256
face_landmarks = face_result[face_index]
lmks = []
for index in range(len(face_landmarks)):
x = face_landmarks[index].x * img_rgb.shape[1]
y = face_landmarks[index].y * img_rgb.shape[0]
lmks.append([x, y])
pts = np.array(lmks)
# crop the face
ret_dct, image_crop = crop_image(
img_rgb, # ndarray
pts, # 106x2 or Nx2
dsize=dsize,
scale=scale,
vy_ratio=vy_ratio,
vx_ratio=vx_ratio,
rotate=rotate
)
# update a 256x256 version for network input or else
cropped_image_256 = cv2.resize(image_crop, (256, 256), interpolation=cv2.INTER_AREA)
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, cropped_image_256
+18 -4
View File
@@ -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')
+30
View File
@@ -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
-48
View File
@@ -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]
-97
View File
@@ -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}")
+6 -16
View File
@@ -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
@@ -52,8 +44,7 @@ class LandmarkRunner(object):
def run(self, img_rgb: np.ndarray, lmk=None):
if lmk is not None:
crop_dct = crop_image(img_rgb, lmk, dsize=self.dsize, scale=1.5, vy_ratio=-0.1)
img_crop_rgb = crop_dct['img_crop']
crop_dct, img_crop_rgb = crop_image(img_rgb, lmk, dsize=self.dsize, scale=1.5, vy_ratio=-0.1)
else:
img_crop_rgb = cv2.resize(img_rgb, (self.dsize, self.dsize))
scale = max(img_rgb.shape[:2]) / self.dsize
@@ -72,13 +63,12 @@ class LandmarkRunner(object):
pts = to_ndarray(out_pts[0]).reshape(-1, 2) * self.dsize # scale to 0-224
pts = _transform_pts(pts, M=crop_dct['M_c2o'])
del crop_dct, img_crop_rgb
return {
'pts': pts, # 2d landmarks 203 points
}
def warmup(self):
# 构造dummy image进行warmup
self.timer.tic()
dummy_image = np.zeros((1, 3, self.dsize, self.dsize), dtype=np.float32)
@@ -86,4 +76,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')
-16
View File
@@ -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
-139
View File
@@ -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 tqdm import tqdm
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 tqdm(range(n), desc='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 tqdm(enumerate(I_p_lst), total=len(I_p_lst), desc='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
+1
View File
@@ -0,0 +1 @@
from .mp_utils import LMKExtractor
File diff suppressed because it is too large Load Diff
+38
View File
@@ -0,0 +1,38 @@
import os
import mediapipe as mp
from mediapipe.tasks import python
from mediapipe.tasks.python import vision
from . import face_landmark
CUR_DIR = os.path.dirname(__file__)
class LMKExtractor():
def __init__(self):
# Create an FaceLandmarker object.
self.mode = mp.tasks.vision.FaceDetectorOptions.running_mode.IMAGE
base_options = python.BaseOptions(model_asset_path=os.path.join(CUR_DIR, 'mp_models','face_landmarker_v2_with_blendshapes.task'))
base_options.delegate = mp.tasks.BaseOptions.Delegate.CPU
options = vision.FaceLandmarkerOptions(base_options=base_options,
running_mode=self.mode,
output_face_blendshapes=False,
output_facial_transformation_matrixes=True,
num_faces=1,
min_face_detection_confidence=0.5,
min_face_presence_confidence=0.5,
min_tracking_confidence=0.5)
self.detector = face_landmark.FaceLandmarker.create_from_options(options)
det_base_options = python.BaseOptions(model_asset_path=os.path.join(CUR_DIR, 'mp_models','blaze_face_short_range.tflite'))
det_options = vision.FaceDetectorOptions(base_options=det_base_options)
self.det_detector = vision.FaceDetector.create_from_options(det_options)
def __call__(self, img):
image = mp.Image(image_format=mp.ImageFormat.SRGB, data=img)
try:
detection_result, _ = self.detector.detect(image)
except:
return None
return detection_result.face_landmarks
+626 -178
View File
@@ -4,106 +4,77 @@ import yaml
import folder_paths
import comfy.model_management as mm
import comfy.utils
import numpy as np
import cv2
from tqdm import tqdm
import gc
script_directory = os.path.dirname(os.path.abspath(__file__))
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
from .liveportrait.utils.cropper import Cropper
try:
from .liveportrait.utils.cropper import CropperMediaPipe
except:
raise ModuleNotFoundError("Can't load MediaPipe, MediaPipeCropper not available")
try:
from .liveportrait.utils.cropper import CropperInsightFace
except:
raise ModuleNotFoundError("Can't load InsightFace, InsightFaceCropper not available")
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,
input_shape=(256, 256),
device_id=0,
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
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.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": (
[
'auto',
'fp16',
'fp32',
], {
"default": 'auto'
}),
}
"fp16",
"fp32",
"auto",
],
{"default": "auto"},
),
},
}
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
@@ -111,26 +82,26 @@ class DownloadAndLoadLivePortraitModels:
FUNCTION = "loadmodel"
CATEGORY = "LivePortrait"
def loadmodel(self, precision='auto'):
def loadmodel(self, precision="fp16"):
device = mm.get_torch_device()
mm.soft_empty_cache()
if precision == 'auto':
try:
if mm.is_device_mps(device):
print("LivePortrait using fp32 for MPS")
log.info("LivePortrait using fp32 for MPS")
dtype = 'fp32'
elif mm.should_use_fp16():
print("LivePortrait using fp16")
log.info("LivePortrait using fp16")
dtype = 'fp16'
else:
print("LivePortrait using fp32")
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
print(f"LivePortrait using {dtype}")
log.info(f"LivePortrait using {dtype}")
pbar = comfy.utils.ProgressBar(3)
@@ -138,86 +109,108 @@ class DownloadAndLoadLivePortraitModels:
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()
stitcher = stitcher.to(device).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()
retargetor_lip = retargetor_lip.to(device).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.')
retargetor_eye = retargetor_eye.to(device).eval()
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(
@@ -228,97 +221,552 @@ class DownloadAndLoadLivePortraitModels:
self.stich_retargeting_module,
InferenceConfig(
device_id=device,
flag_use_half_precision = True if dtype == 'fp16' else False
)
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.")
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 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(
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()
gc.collect()
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.cpu())
out_mask_list.append(torch.zeros((1, 3, H, W), device="cpu"))
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.cpu())
out_mask_list.append(mask_ori.cpu())
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.float(),
mask_tensors_out.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 = CropperInsightFace(**cropper_init_config)
return (self.cropper,)
class LivePortraitLoadMediaPipeCropper:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"landmarkrunner_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, landmarkrunner_onnx_device, keep_model_loaded):
cropper_init_config = {
'keep_model_loaded': keep_model_loaded,
'onnx_device': landmarkrunner_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 = CropperMediaPipe(**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.contiguous() * 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, cropped_image_256 = 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_images_list.append(cropped_image_256)
I_s = pipeline.live_portrait_wrapper.prepare_source(cropped_image_256)
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)
del I_s
else:
log.warning(f"Warning: No face detected on frame {str(i)}, skipping")
cropped_images_list.append(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'])
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
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
}
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()
return (retargeting_info,)
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)
class KeypointsToImage:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"crop_info": ("CROPINFO", {"default": []}),
},
}
cropped_tensors_out = torch.cat(cropped_out_list, dim=0)
full_tensors_out = torch.cat(full_out_list, dim=0)
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("keypoints_image",)
FUNCTION = "drawkeypoints"
CATEGORY = "LivePortrait"
return (cropped_tensors_out, full_tensors_out)
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()
return (crop_info, keypoints_image_tensor,)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
"LivePortraitProcess": LivePortraitProcess,
"LivePortraitCropper": LivePortraitCropper,
"LivePortraitRetargeting": LivePortraitRetargeting,
#"KeypointScaler": KeypointScaler,
"KeypointsToImage": KeypointsToImage,
"LivePortraitLoadCropper": LivePortraitLoadCropper,
"LivePortraitLoadMediaPipeCropper": LivePortraitLoadMediaPipeCropper,
"LivePortraitComposite": LivePortraitComposite,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
"LivePortraitProcess": "LivePortraitProcess",
"LivePortraitProcess": "LivePortrait Process",
"LivePortraitCropper": "LivePortrait Cropper",
"LivePortraitRetargeting": "LivePortrait Retargeting",
#"KeypointScaler": "KeypointScaler",
"KeypointsToImage": "LivePortrait KeypointsToImage",
"LivePortraitLoadCropper": "LivePortrait Load InsightFaceCropper",
"LivePortraitLoadMediaPipeCropper": "LivePortrait Load MediaPipeCropper",
"LivePortraitComposite": "LivePortrait Composite",
}
+33 -2
View File
@@ -1,15 +1,46 @@
# ComfyUI nodes to use [LivePortrait](https://github.com/KwaiVGI/LivePortrait)
## Update
https://github.com/kijai/ComfyUI-LivePortrait/assets/40791699/e55e10f6-af61-4d73-b162-af29eb847516
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
Changes
- Added MediaPipe as alternative to Insightface, everything should now be covered under MIT and Apache-2.0 licenses when using it.
- Proper Vid2vid including smoothing algorhitm (thanks @melMass)
- Improved speed and efficiency, allows for near realtime view even in Comfy (~80-100ms delay)
- Restructured nodes for more options
- Auto skipping frames with no face detected
- Numerous other things I have forgotten about at this point, it's been a lot
- Better Mac support on MPS (thanks @Grant-CP
# Examples:
Realtime with webcam feed:
https://github.com/user-attachments/assets/31f77c10-b757-44ae-bb26-39e45ec0b2d9
Image2vid:
https://github.com/user-attachments/assets/cfec0419-d1eb-4e67-8913-890eeb155eef
Vid2Vid:
https://github.com/user-attachments/assets/28438fcb-fbb0-4e4e-baf4-00fe06c455de
I have converted all the pickle files to safetensors: https://huggingface.co/Kijai/LivePortrait_safetensors/tree/main
They go here (and are automatically downloaded if the folder is not present) `ComfyUI/models/liveportrait`
# Face detectors
Insightface is also required.
You can either use the original default Insightface, or Google's MediaPipe.
Biggest difference is the license: Insightface is strictly for NON-COMMERCIAL use.
MediaPipe is a bit worse at detection, and can't run on GPU in Windows, though it's much faster on CPU compared to Insightface
Insightface is not automatically installed, if you wish to use it follow these instructions:
If you have a working compile environment, installing it can be as easy as:
`pip install insightface`
+2 -1
View File
@@ -1,5 +1,6 @@
pyyaml
numpy
opencv-python
rich
onnxruntime-gpu
pykalman
mediapipe