add real time node
This commit is contained in:
+5
-3
@@ -1,4 +1,4 @@
|
||||
from .nodes import MuseTalk,LoadVideo,PreViewVideo,CombineAudioVideo
|
||||
from .nodes import MuseTalk,LoadVideo,PreViewVideo,CombineAudioVideo,MuseTalkRealTime
|
||||
WEB_DIRECTORY = "./web"
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
@@ -6,7 +6,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"MuseTalk": MuseTalk,
|
||||
"LoadVideo": LoadVideo,
|
||||
"PreViewVideo": PreViewVideo,
|
||||
"CombineAudioVideo": CombineAudioVideo
|
||||
"CombineAudioVideo": CombineAudioVideo,
|
||||
"MuseTalkRealTime": MuseTalkRealTime
|
||||
}
|
||||
|
||||
# A dictionary that contains the friendly/humanly readable titles for the nodes
|
||||
@@ -14,5 +15,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MuseTalk": "MuseTalk Node",
|
||||
"LoadVideo": "Video Loader",
|
||||
"PreViewVideo": "PreView Video",
|
||||
"CombineAudioVideo": "Combine Audio Video"
|
||||
"CombineAudioVideo": "Combine Audio Video",
|
||||
"MuseTalkRealTime": "MuseTalk RealTime Node"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
|
||||
import os
|
||||
import sys
|
||||
import cv2
|
||||
import json
|
||||
import torch
|
||||
import shutil
|
||||
import pickle
|
||||
import glob,time
|
||||
import queue,copy
|
||||
import threading
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
import folder_paths
|
||||
from cuda_malloc import cuda_malloc_supported
|
||||
from typing import Any
|
||||
from .musetalk.utils.utils import load_all_model,datagen
|
||||
from .musetalk.utils.preprocessing import read_imgs,get_landmark_and_bbox
|
||||
from .musetalk.utils.blending import get_image,get_image_prepare_material,get_image_blending
|
||||
|
||||
parent_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
# load model weights
|
||||
audio_processor,vae,unet,pe = load_all_model(os.path.join(parent_directory,"models"))
|
||||
device = torch.device("cuda" if cuda_malloc_supported() else "cpu")
|
||||
timesteps = torch.tensor([0], device=device)
|
||||
|
||||
output_path = folder_paths.get_output_directory()
|
||||
musetalk_out_path = os.path.join(output_path,"musetalk_realtime")
|
||||
os.makedirs(musetalk_out_path, exist_ok=True)
|
||||
|
||||
def osmakedirs(path_list):
|
||||
for path in path_list:
|
||||
os.makedirs(path) if not os.path.exists(path) else None
|
||||
|
||||
def video2imgs(vid_path, save_path, ext = '.png',cut_frame = 10000000):
|
||||
cap = cv2.VideoCapture(vid_path)
|
||||
count = 0
|
||||
while True:
|
||||
if count > cut_frame:
|
||||
break
|
||||
ret, frame = cap.read()
|
||||
if ret:
|
||||
cv2.imwrite(f"{save_path}/{count:08d}.png", frame)
|
||||
count += 1
|
||||
else:
|
||||
break
|
||||
|
||||
@torch.no_grad()
|
||||
class Avatar:
|
||||
def __init__(self, avatar_id, video_path, bbox_shift, batch_size, preparation):
|
||||
self.avatar_id = avatar_id
|
||||
self.video_path = video_path
|
||||
self.bbox_shift = bbox_shift
|
||||
self.avatar_path = os.path.join(musetalk_out_path,avatar_id)
|
||||
self.full_imgs_path = f"{self.avatar_path}/full_imgs"
|
||||
self.coords_path = f"{self.avatar_path}/coords.pkl"
|
||||
self.latents_out_path= f"{self.avatar_path}/latents.pt"
|
||||
self.video_out_path = output_path
|
||||
self.mask_out_path =f"{self.avatar_path}/mask"
|
||||
self.mask_coords_path =f"{self.avatar_path}/mask_coords.pkl"
|
||||
self.avatar_info_path = f"{self.avatar_path}/avator_info.json"
|
||||
self.avatar_info = {
|
||||
"avatar_id":avatar_id,
|
||||
"video_path":video_path,
|
||||
"bbox_shift":bbox_shift
|
||||
}
|
||||
self.preparation = preparation
|
||||
self.batch_size = batch_size
|
||||
self.idx = 0
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
if self.preparation:
|
||||
if os.path.exists(self.avatar_path):
|
||||
response = input(f"{self.avatar_id} exists, Do you want to re-create it ? (y/n)")
|
||||
if response.lower() == "y":
|
||||
shutil.rmtree(self.avatar_path)
|
||||
print("*********************************")
|
||||
print(f" creating avator: {self.avatar_id}")
|
||||
print("*********************************")
|
||||
osmakedirs([self.avatar_path,self.full_imgs_path,self.video_out_path,self.mask_out_path])
|
||||
self.prepare_material()
|
||||
else:
|
||||
self.input_latent_list_cycle = torch.load(self.latents_out_path)
|
||||
with open(self.coords_path, 'rb') as f:
|
||||
self.coord_list_cycle = pickle.load(f)
|
||||
input_img_list = glob.glob(os.path.join(self.full_imgs_path, '*.[jpJP][pnPN]*[gG]'))
|
||||
input_img_list = sorted(input_img_list, key=lambda x: int(os.path.splitext(os.path.basename(x))[0]))
|
||||
self.frame_list_cycle = read_imgs(input_img_list)
|
||||
with open(self.mask_coords_path, 'rb') as f:
|
||||
self.mask_coords_list_cycle = pickle.load(f)
|
||||
input_mask_list = glob.glob(os.path.join(self.mask_out_path, '*.[jpJP][pnPN]*[gG]'))
|
||||
input_mask_list = sorted(input_mask_list, key=lambda x: int(os.path.splitext(os.path.basename(x))[0]))
|
||||
self.mask_list_cycle = read_imgs(input_mask_list)
|
||||
else:
|
||||
print("*********************************")
|
||||
print(f" creating avator: {self.avatar_id}")
|
||||
print("*********************************")
|
||||
osmakedirs([self.avatar_path,self.full_imgs_path,self.video_out_path,self.mask_out_path])
|
||||
self.prepare_material()
|
||||
else:
|
||||
with open(self.avatar_info_path, "r") as f:
|
||||
avatar_info = json.load(f)
|
||||
|
||||
if avatar_info['bbox_shift'] != self.avatar_info['bbox_shift']:
|
||||
response = input(f" 【bbox_shift】 is changed, you need to re-create it ! (c/continue)")
|
||||
if response.lower() == "c":
|
||||
shutil.rmtree(self.avatar_path)
|
||||
print("*********************************")
|
||||
print(f" creating avator: {self.avatar_id}")
|
||||
print("*********************************")
|
||||
osmakedirs([self.avatar_path,self.full_imgs_path,self.video_out_path,self.mask_out_path])
|
||||
self.prepare_material()
|
||||
else:
|
||||
sys.exit()
|
||||
else:
|
||||
self.input_latent_list_cycle = torch.load(self.latents_out_path)
|
||||
with open(self.coords_path, 'rb') as f:
|
||||
self.coord_list_cycle = pickle.load(f)
|
||||
input_img_list = glob.glob(os.path.join(self.full_imgs_path, '*.[jpJP][pnPN]*[gG]'))
|
||||
input_img_list = sorted(input_img_list, key=lambda x: int(os.path.splitext(os.path.basename(x))[0]))
|
||||
self.frame_list_cycle = read_imgs(input_img_list)
|
||||
with open(self.mask_coords_path, 'rb') as f:
|
||||
self.mask_coords_list_cycle = pickle.load(f)
|
||||
input_mask_list = glob.glob(os.path.join(self.mask_out_path, '*.[jpJP][pnPN]*[gG]'))
|
||||
input_mask_list = sorted(input_mask_list, key=lambda x: int(os.path.splitext(os.path.basename(x))[0]))
|
||||
self.mask_list_cycle = read_imgs(input_mask_list)
|
||||
|
||||
def prepare_material(self):
|
||||
print("preparing data materials ... ...")
|
||||
with open(self.avatar_info_path, "w") as f:
|
||||
json.dump(self.avatar_info, f)
|
||||
|
||||
if os.path.isfile(self.video_path):
|
||||
video2imgs(self.video_path, self.full_imgs_path, ext = 'png')
|
||||
else:
|
||||
print(f"copy files in {self.video_path}")
|
||||
files = os.listdir(self.video_path)
|
||||
files.sort()
|
||||
files = [file for file in files if file.split(".")[-1]=="png"]
|
||||
for filename in files:
|
||||
shutil.copyfile(f"{self.video_path}/{filename}", f"{self.full_imgs_path}/{filename}")
|
||||
input_img_list = sorted(glob.glob(os.path.join(self.full_imgs_path, '*.[jpJP][pnPN]*[gG]')))
|
||||
|
||||
print("extracting landmarks...")
|
||||
coord_list, frame_list = get_landmark_and_bbox(input_img_list, self.bbox_shift)
|
||||
input_latent_list = []
|
||||
idx = -1
|
||||
# maker if the bbox is not sufficient
|
||||
coord_placeholder = (0.0,0.0,0.0,0.0)
|
||||
for bbox, frame in zip(coord_list, frame_list):
|
||||
idx = idx + 1
|
||||
if bbox == coord_placeholder:
|
||||
continue
|
||||
x1, y1, x2, y2 = bbox
|
||||
crop_frame = frame[y1:y2, x1:x2]
|
||||
resized_crop_frame = cv2.resize(crop_frame,(256,256),interpolation = cv2.INTER_LANCZOS4)
|
||||
latents = vae.get_latents_for_unet(resized_crop_frame)
|
||||
input_latent_list.append(latents)
|
||||
|
||||
self.frame_list_cycle = frame_list + frame_list[::-1]
|
||||
self.coord_list_cycle = coord_list + coord_list[::-1]
|
||||
self.input_latent_list_cycle = input_latent_list + input_latent_list[::-1]
|
||||
self.mask_coords_list_cycle = []
|
||||
self.mask_list_cycle = []
|
||||
|
||||
for i,frame in enumerate(tqdm(self.frame_list_cycle)):
|
||||
cv2.imwrite(f"{self.full_imgs_path}/{str(i).zfill(8)}.png",frame)
|
||||
|
||||
face_box = self.coord_list_cycle[i]
|
||||
mask,crop_box = get_image_prepare_material(frame,face_box)
|
||||
cv2.imwrite(f"{self.mask_out_path}/{str(i).zfill(8)}.png",mask)
|
||||
self.mask_coords_list_cycle += [crop_box]
|
||||
self.mask_list_cycle.append(mask)
|
||||
|
||||
with open(self.mask_coords_path, 'wb') as f:
|
||||
pickle.dump(self.mask_coords_list_cycle, f)
|
||||
|
||||
with open(self.coords_path, 'wb') as f:
|
||||
pickle.dump(self.coord_list_cycle, f)
|
||||
|
||||
torch.save(self.input_latent_list_cycle, os.path.join(self.latents_out_path))
|
||||
#
|
||||
|
||||
def process_frames(self, res_frame_queue,video_len):
|
||||
print(video_len)
|
||||
while True:
|
||||
if self.idx>=video_len-1:
|
||||
break
|
||||
try:
|
||||
start = time.time()
|
||||
res_frame = res_frame_queue.get(block=True, timeout=1)
|
||||
except queue.Empty:
|
||||
continue
|
||||
|
||||
bbox = self.coord_list_cycle[self.idx%(len(self.coord_list_cycle))]
|
||||
ori_frame = copy.deepcopy(self.frame_list_cycle[self.idx%(len(self.frame_list_cycle))])
|
||||
x1, y1, x2, y2 = bbox
|
||||
try:
|
||||
res_frame = cv2.resize(res_frame.astype(np.uint8),(x2-x1,y2-y1))
|
||||
except:
|
||||
continue
|
||||
mask = self.mask_list_cycle[self.idx%(len(self.mask_list_cycle))]
|
||||
mask_crop_box = self.mask_coords_list_cycle[self.idx%(len(self.mask_coords_list_cycle))]
|
||||
#combine_frame = get_image(ori_frame,res_frame,bbox)
|
||||
combine_frame = get_image_blending(ori_frame,res_frame,bbox,mask,mask_crop_box)
|
||||
|
||||
fps = 1/(time.time()-start+1e-6)
|
||||
print(f"Displaying the {self.idx}-th frame with FPS: {fps:.2f}")
|
||||
cv2.imwrite(f"{self.avatar_path}/tmp/{str(self.idx).zfill(8)}.png",combine_frame)
|
||||
self.idx = self.idx + 1
|
||||
|
||||
def inference(self, audio_path, out_vid_name, fps):
|
||||
os.makedirs(self.avatar_path+'/tmp',exist_ok =True)
|
||||
############################################## extract audio feature ##############################################
|
||||
whisper_feature = audio_processor.audio2feat(audio_path)
|
||||
whisper_chunks = audio_processor.feature2chunks(feature_array=whisper_feature,fps=fps)
|
||||
############################################## inference batch by batch ##############################################
|
||||
video_num = len(whisper_chunks)
|
||||
print("start inference")
|
||||
res_frame_queue = queue.Queue()
|
||||
self.idx = 0
|
||||
# # Create a sub-thread and start it
|
||||
process_thread = threading.Thread(target=self.process_frames, args=(res_frame_queue,video_num))
|
||||
process_thread.start()
|
||||
start_time = time.time()
|
||||
gen = datagen(whisper_chunks,self.input_latent_list_cycle, self.batch_size)
|
||||
print(f"processing audio:{audio_path} costs {(time.time() - start_time) * 1000}ms")
|
||||
start_time = time.time()
|
||||
res_frame_list = []
|
||||
|
||||
for i, (whisper_batch,latent_batch) in enumerate(tqdm(gen,total=int(np.ceil(float(video_num)/self.batch_size)))):
|
||||
start_time = time.time()
|
||||
tensor_list = [torch.FloatTensor(arr) for arr in whisper_batch]
|
||||
audio_feature_batch = torch.stack(tensor_list).to(unet.device) # torch, B, 5*N,384
|
||||
audio_feature_batch = pe(audio_feature_batch)
|
||||
|
||||
pred_latents = unet.model(latent_batch, timesteps, encoder_hidden_states=audio_feature_batch).sample
|
||||
recon = vae.decode_latents(pred_latents)
|
||||
for res_frame in recon:
|
||||
res_frame_queue.put(res_frame)
|
||||
# Close the queue and sub-thread after all tasks are completed
|
||||
process_thread.join()
|
||||
|
||||
if out_vid_name is not None:
|
||||
# optional
|
||||
cmd_img2video = f"ffmpeg -y -v warning -r {fps} -f image2 -i {self.avatar_path}/tmp/%08d.png -vcodec libx264 -vf format=rgb24,scale=out_color_matrix=bt709,format=yuv420p -crf 18 {self.avatar_path}/temp.mp4"
|
||||
print(cmd_img2video)
|
||||
os.system(cmd_img2video)
|
||||
|
||||
output_vid = os.path.join(self.video_out_path, out_vid_name+".mp4") # on
|
||||
cmd_combine_audio = f"ffmpeg -y -v warning -i {audio_path} -i {self.avatar_path}/temp.mp4 {output_vid}"
|
||||
print(cmd_combine_audio)
|
||||
os.system(cmd_combine_audio)
|
||||
|
||||
os.remove(f"{self.avatar_path}/temp.mp4")
|
||||
shutil.rmtree(f"{self.avatar_path}/tmp")
|
||||
print(f"result is save to {output_vid}")
|
||||
return output_vid
|
||||
|
||||
class Infer_Real_Time:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def __call__(self, audio_path,video_path,
|
||||
avatar_id,fps=25,batch_size=4,
|
||||
preparation=True,bbox_shift=0,
|
||||
*args: Any, **kwds: Any) -> Any:
|
||||
|
||||
avatar = Avatar(
|
||||
avatar_id = avatar_id,
|
||||
video_path = video_path,
|
||||
bbox_shift = bbox_shift,
|
||||
batch_size = batch_size,
|
||||
preparation= preparation)
|
||||
output_name = os.path.basename(audio_path)[:-4]
|
||||
return avatar.inference(audio_path,output_name,fps)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -55,3 +55,44 @@ def get_image(fp_model,image,face,face_box,upper_boundary_ratio = 0.5,expand=1.2
|
||||
body.paste(face_large, crop_box[:2], mask_image)
|
||||
body = np.array(body)
|
||||
return body[:,:,::-1]
|
||||
|
||||
def get_image_prepare_material(image,face_box,upper_boundary_ratio = 0.5,expand=1.2):
|
||||
body = Image.fromarray(image[:,:,::-1])
|
||||
|
||||
x, y, x1, y1 = face_box
|
||||
#print(x1-x,y1-y)
|
||||
crop_box, s = get_crop_box(face_box, expand)
|
||||
x_s, y_s, x_e, y_e = crop_box
|
||||
|
||||
face_large = body.crop(crop_box)
|
||||
ori_shape = face_large.size
|
||||
|
||||
mask_image = face_seg(face_large)
|
||||
mask_small = mask_image.crop((x-x_s, y-y_s, x1-x_s, y1-y_s))
|
||||
mask_image = Image.new('L', ori_shape, 0)
|
||||
mask_image.paste(mask_small, (x-x_s, y-y_s, x1-x_s, y1-y_s))
|
||||
|
||||
# keep upper_boundary_ratio of talking area
|
||||
width, height = mask_image.size
|
||||
top_boundary = int(height * upper_boundary_ratio)
|
||||
modified_mask_image = Image.new('L', ori_shape, 0)
|
||||
modified_mask_image.paste(mask_image.crop((0, top_boundary, width, height)), (0, top_boundary))
|
||||
|
||||
blur_kernel_size = int(0.1 * ori_shape[0] // 2 * 2) + 1
|
||||
mask_array = cv2.GaussianBlur(np.array(modified_mask_image), (blur_kernel_size, blur_kernel_size), 0)
|
||||
return mask_array,crop_box
|
||||
|
||||
def get_image_blending(image,face,face_box,mask_array,crop_box):
|
||||
body = Image.fromarray(image[:,:,::-1])
|
||||
face = Image.fromarray(face[:,:,::-1])
|
||||
|
||||
x, y, x1, y1 = face_box
|
||||
x_s, y_s, x_e, y_e = crop_box
|
||||
face_large = body.crop(crop_box)
|
||||
|
||||
mask_image = Image.fromarray(mask_array)
|
||||
mask_image = mask_image.convert("L")
|
||||
face_large.paste(face, (x-x_s, y-y_s, x1-x_s, y1-y_s))
|
||||
body.paste(face_large, crop_box[:2], mask_image)
|
||||
body = np.array(body)
|
||||
return body[:,:,::-1]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import folder_paths
|
||||
from .inference import MuseTalk_INFER
|
||||
from .inference_realtime import Infer_Real_Time
|
||||
from pydub import AudioSegment
|
||||
from moviepy.editor import VideoFileClip,AudioFileClip
|
||||
|
||||
@@ -8,6 +9,46 @@ parent_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
input_path = folder_paths.get_input_directory()
|
||||
out_path = folder_paths.get_output_directory()
|
||||
|
||||
class MuseTalkRealTime:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required":{
|
||||
"audio":("AUDIO",),
|
||||
"video":("VIDEO",),
|
||||
"avatar_id":("STRING",{
|
||||
"default": "talker1"
|
||||
}),
|
||||
"bbox_shift":("INT",{
|
||||
"default":0
|
||||
}),
|
||||
"fps":("INT",{
|
||||
"default":25
|
||||
}),
|
||||
"batch_size":("INT",{
|
||||
"default":4
|
||||
}),
|
||||
"preparation":("BOOLEAN",{
|
||||
"default":True
|
||||
})
|
||||
}
|
||||
}
|
||||
CATEGORY = "AIFSH_MuseTalk"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
|
||||
OUTPUT_NODE = False
|
||||
|
||||
FUNCTION = "process"
|
||||
|
||||
def process(self,audio,video,avatar_id,bbox_shift,fps,batch_size,preparation):
|
||||
muse_talk_real_time = Infer_Real_Time()
|
||||
output_vid_name = muse_talk_real_time(audio, video,avatar_id,fps=fps,batch_size=batch_size,
|
||||
preparation=preparation,bbox_shift=bbox_shift)
|
||||
return (output_vid_name,)
|
||||
|
||||
|
||||
class MuseTalk:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
Reference in New Issue
Block a user