accept pose videos

This commit is contained in:
Brandon Thomas
2025-01-20 02:26:37 -05:00
parent 5712d89e73
commit 713d69f390
+131 -23
View File
@@ -16,6 +16,8 @@ from animation.modules.pose_net import PoseNet
from animation.modules.unet import UNetSpatioTemporalConditionModel
from animation.pipelines.inference_pipeline_animation import InferenceAnimationPipeline
import random
import subprocess
from pathlib import Path
def seed_everything(seed):
torch.manual_seed(seed)
@@ -78,9 +80,10 @@ def export_to_gif(frames, output_gif_path, fps):
def parse_args():
parser = argparse.ArgumentParser(
description="Script to train Stable Diffusion XL for InstructPix2Pix."
description="Inference script"
)
# KEEP
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
@@ -88,8 +91,20 @@ def parse_args():
required=True
)
## KEEP
#parser.add_argument(
# "--validation_image",
# type=str,
# default=None,
# help=(
# "A set of paths to the controlnext conditioning image be evaluated every `--validation_steps`"
# " and logged to `--report_to`. Provide either a matching number of `--validation_prompt`s, a"
# " a single `--validation_prompt` to be used with all `--validation_image`s, or a single"
# " `--validation_image` that will be used with all `--validation_prompt`s."
# ),
#)
parser.add_argument(
"--validation_image",
"--start_image_path",
type=str,
default=None,
help=(
@@ -99,15 +114,41 @@ def parse_args():
" `--validation_image` that will be used with all `--validation_prompt`s."
),
)
# ===================
# Pose Input
# Supply either `pose_images_folder`, `pose_video_path`, or `pre_pose_video_path`.
# - pose_images_folder has inputs already fully processed
# - pose_video_path just needs to be converted to frames
# - pre_pose_video_path has not been converted to pose data frames
# ===================
parser.add_argument(
"--validation_control_folder",
"--pose_images_folder",
type=str,
default=None,
help=(
"the validation control image"
"the validation control images"
),
)
parser.add_argument(
"--pose_video_path",
type=str,
default=None,
help=(
"the validation control video (optional)"
),
)
parser.add_argument(
"--pre_pose_video_path",
type=str,
default=None,
help=(
"video to convert to pose"
),
)
# TODO FIX
parser.add_argument(
"--output_dir",
type=str,
@@ -115,6 +156,7 @@ def parse_args():
required=True
)
# TODO FIX
parser.add_argument(
"--height",
type=int,
@@ -122,6 +164,7 @@ def parse_args():
required=False
)
# TODO FIX
parser.add_argument(
"--width",
type=int,
@@ -129,6 +172,7 @@ def parse_args():
required=False
)
# KEEP
parser.add_argument(
"--guidance_scale",
type=float,
@@ -136,6 +180,7 @@ def parse_args():
required=False
)
# KEEP
parser.add_argument(
"--num_inference_steps",
type=int,
@@ -143,18 +188,23 @@ def parse_args():
required=False
)
# KEEP
parser.add_argument(
"--posenet_model_name_or_path",
type=str,
default=None,
help="Path to pretrained posenet model",
)
# KEEP
parser.add_argument(
"--face_encoder_model_name_or_path",
type=str,
default=None,
help="Path to pretrained face encoder model",
)
# KEEP
parser.add_argument(
"--unet_model_name_or_path",
type=str,
@@ -162,6 +212,7 @@ def parse_args():
help="Path to pretrained unet model",
)
# KEEP
parser.add_argument(
"--tile_size",
type=int,
@@ -169,6 +220,7 @@ def parse_args():
required=False
)
# KEEP
parser.add_argument(
"--overlap",
type=int,
@@ -176,23 +228,30 @@ def parse_args():
required=False
)
# KEEP
parser.add_argument(
"--noise_aug_strength",
type=float,
default=0.0, # or set to 0.02
required=False
)
# KEEP
parser.add_argument(
"--frames_overlap",
type=int,
default=4,
required=False
)
# KEEP
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
# KEEP
parser.add_argument(
"--revision",
type=str,
@@ -200,6 +259,8 @@ def parse_args():
required=False,
help="Revision of pretrained model identifier from huggingface.co/models.",
)
# KEEP
parser.add_argument(
"--decode_chunk_size",
type=int,
@@ -210,6 +271,52 @@ def parse_args():
args = parser.parse_args()
return args
def split_video_to_frames(video_path, output_path):
Path(output_path).mkdir(parents=True, exist_ok=True)
filename_format = f"{output_path}/frame_%d.png"
command = [
'ffmpeg',
'-i', video_path,
'-q:v', '1',
'-start_number', '0',
filename_format,
]
print(f"Command: {command}", flush=True)
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
#video_info = json.loads(result.stdout)
def prepare_pose_frames(args):
pose_images_dir = args.pose_images_folder
if not pose_images_dir:
raise Exception('pose_images_folder must be set, even when using other arguments')
Path(pose_images_dir).mkdir(parents=True, exist_ok=True)
if args.pre_pose_video_path:
print("Preparing pose frames from pre-pose video.")
pre_pose_frame_dir = Path(f"{pose_images_dir}/frames")
pre_pose_frame_dir.mkdir(parents=True, exist_ok=True)
split_video_to_frames(args.pre_pose_video_path, pre_pose_frame_dir)
reference_image_path = pre_pose_frame_dir / "frame_0.png"
command = [
"python", "DWPose/skeleton_extraction.py",
"--target_image_folder_path", pre_pose_frame_dir,
"--ref_image_path", reference_image_path,
"--poses_folder_path", pose_images_dir,
]
print(f"Command: {command}", flush=True)
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True)
elif args.pose_video_path:
print("Preparing pose frames from pose video.")
split_video_to_frames(args.pose_video_path, pose_images_dir)
else:
print("Pose frames assumed to already exist.")
# TODO: autodetect resolution
pose_images = load_images_from_folder(pose_images_dir, width=args.width, height=args.height)
return pose_images
if __name__ == "__main__":
args = parse_args()
@@ -333,39 +440,40 @@ if __name__ == "__main__":
os.makedirs(args.output_dir, exist_ok=True)
validation_image_path = args.validation_image
validation_image = Image.open(args.validation_image).convert('RGB')
validation_control_images = load_images_from_folder(args.validation_control_folder, width=args.width, height=args.height)
pose_images = prepare_pose_frames(args)
num_frames = len(pose_images)
start_image_path = args.start_image_path
start_image = Image.open(args.start_image_path).convert('RGB')
num_frames = len(validation_control_images)
face_model.face_helper.clean_all()
validation_face = cv2.imread(validation_image_path)
validation_image_bgr = cv2.cvtColor(validation_face, cv2.COLOR_RGB2BGR)
validation_image_face_info = face_model.app.get(validation_image_bgr)
if len(validation_image_face_info) > 0:
validation_image_face_info = sorted(validation_image_face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1]
validation_image_id_ante_embedding = validation_image_face_info['embedding']
start_image_face = cv2.imread(start_image_path)
start_image_bgr = cv2.cvtColor(start_image_face, cv2.COLOR_RGB2BGR)
start_image_face_info = face_model.app.get(start_image_bgr)
if len(start_image_face_info) > 0:
start_image_face_info = sorted(start_image_face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1]
start_image_id_ante_embedding = start_image_face_info['embedding']
else:
validation_image_id_ante_embedding = None
start_image_id_ante_embedding = None
if validation_image_id_ante_embedding is None:
face_model.face_helper.read_image(validation_image_bgr)
if start_image_id_ante_embedding is None:
face_model.face_helper.read_image(start_image_bgr)
face_model.face_helper.get_face_landmarks_5(only_center_face=True)
face_model.face_helper.align_warp_face()
if len(face_model.face_helper.cropped_faces) == 0:
validation_image_id_ante_embedding = np.zeros((512,))
start_image_id_ante_embedding = np.zeros((512,))
else:
validation_image_align_face = face_model.face_helper.cropped_faces[0]
start_image_align_face = face_model.face_helper.cropped_faces[0]
print('fail to detect face using insightface, extract embedding on align face')
validation_image_id_ante_embedding = face_model.handler_ante.get_feat(validation_image_align_face)
start_image_id_ante_embedding = face_model.handler_ante.get_feat(start_image_align_face)
# generator = torch.Generator(device=accelerator.device).manual_seed(23123134)
decode_chunk_size = args.decode_chunk_size
video_frames = pipeline(
image=validation_image,
image_pose=validation_control_images,
image=start_image,
image_pose=pose_images,
height=args.height,
width=args.width,
num_frames=num_frames,
@@ -380,7 +488,7 @@ if __name__ == "__main__":
num_inference_steps=args.num_inference_steps,
generator=generator,
output_type="pil",
validation_image_id_ante_embedding=validation_image_id_ante_embedding,
validation_image_id_ante_embedding=start_image_id_ante_embedding,
).frames[0]
out_file = os.path.join(