mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
accept pose videos
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user