677 lines
27 KiB
Python
Executable File
677 lines
27 KiB
Python
Executable File
import gc
|
|
import inspect
|
|
import math
|
|
import os
|
|
import shutil
|
|
import subprocess
|
|
import time
|
|
|
|
import cv2
|
|
import imageio
|
|
import numpy as np
|
|
import torch
|
|
import torchvision
|
|
from einops import rearrange
|
|
from PIL import Image
|
|
|
|
|
|
def filter_kwargs(cls, kwargs):
|
|
sig = inspect.signature(cls.__init__)
|
|
valid_params = set(sig.parameters.keys()) - {'self', 'cls'}
|
|
filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params}
|
|
return filtered_kwargs
|
|
|
|
def get_width_and_height_from_image_and_base_resolution(image, base_resolution):
|
|
target_pixels = int(base_resolution) * int(base_resolution)
|
|
original_width, original_height = Image.open(image).size
|
|
ratio = (target_pixels / (original_width * original_height)) ** 0.5
|
|
width_slider = round(original_width * ratio)
|
|
height_slider = round(original_height * ratio)
|
|
return height_slider, width_slider
|
|
|
|
def color_transfer(sc, dc):
|
|
"""
|
|
Transfer color distribution from of sc, referred to dc.
|
|
|
|
Args:
|
|
sc (numpy.ndarray): input image to be transfered.
|
|
dc (numpy.ndarray): reference image
|
|
|
|
Returns:
|
|
numpy.ndarray: Transferred color distribution on the sc.
|
|
"""
|
|
|
|
def get_mean_and_std(img):
|
|
x_mean, x_std = cv2.meanStdDev(img)
|
|
x_mean = np.hstack(np.around(x_mean, 2))
|
|
x_std = np.hstack(np.around(x_std, 2))
|
|
return x_mean, x_std
|
|
|
|
sc = cv2.cvtColor(sc, cv2.COLOR_RGB2LAB)
|
|
s_mean, s_std = get_mean_and_std(sc)
|
|
dc = cv2.cvtColor(dc, cv2.COLOR_RGB2LAB)
|
|
t_mean, t_std = get_mean_and_std(dc)
|
|
img_n = ((sc - s_mean) * (t_std / s_std)) + t_mean
|
|
np.putmask(img_n, img_n > 255, 255)
|
|
np.putmask(img_n, img_n < 0, 0)
|
|
dst = cv2.cvtColor(cv2.convertScaleAbs(img_n), cv2.COLOR_LAB2RGB)
|
|
return dst
|
|
|
|
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=12, imageio_backend=True, color_transfer_post_process=False):
|
|
videos = rearrange(videos, "b c t h w -> t b c h w")
|
|
outputs = []
|
|
for x in videos:
|
|
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
|
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
|
if rescale:
|
|
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
|
x = (x * 255).cpu().numpy().astype(np.uint8)
|
|
outputs.append(Image.fromarray(x))
|
|
|
|
if color_transfer_post_process:
|
|
for i in range(1, len(outputs)):
|
|
outputs[i] = Image.fromarray(color_transfer(np.uint8(outputs[i]), np.uint8(outputs[0])))
|
|
|
|
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
if imageio_backend:
|
|
if path.endswith("mp4"):
|
|
imageio.mimsave(path, outputs, fps=fps)
|
|
else:
|
|
imageio.mimsave(path, outputs, duration=(1000 * 1/fps))
|
|
else:
|
|
if path.endswith("mp4"):
|
|
path = path.replace('.mp4', '.gif')
|
|
outputs[0].save(path, format='GIF', append_images=outputs, save_all=True, duration=100, loop=0)
|
|
print(f"Saved video to: {path}")
|
|
|
|
class StreamVideoSaver:
|
|
"""Incrementally write video frames to an mp4 as each block is decoded.
|
|
|
|
Used as a streaming pipeline `decode_callback`: it receives one pixel chunk
|
|
per causal block ([B, C, F, H, W] float in [0, 1]) and appends its frames to
|
|
an open imageio writer, so a long video never has to be held in memory. The
|
|
file is only finalized on `close()`.
|
|
"""
|
|
def __init__(self, video_path, fps):
|
|
os.makedirs(os.path.dirname(video_path), exist_ok=True)
|
|
self.video_path = video_path
|
|
self.writer = imageio.get_writer(video_path, fps=fps)
|
|
self.num_frames = 0
|
|
|
|
def __call__(self, video_chunk, block_idx):
|
|
# video_chunk: [B, C, F, H, W] in [0, 1]; take the first sample.
|
|
chunk = video_chunk[0].permute(1, 2, 3, 0) # [F, H, W, C]
|
|
chunk = (chunk.clamp(0, 1) * 255).numpy().astype(np.uint8)
|
|
for frame in chunk:
|
|
self.writer.append_data(frame)
|
|
self.num_frames += 1
|
|
|
|
def close(self):
|
|
self.writer.close()
|
|
print(f"Saved video to: {self.video_path} ({self.num_frames} frames)")
|
|
|
|
class SegmentVideoSaver:
|
|
"""Save each decoded block as its own standalone mp4, flushed immediately,
|
|
while also assembling one complete continuous mp4.
|
|
|
|
Unlike `StreamVideoSaver` (only one continuous file finalized on close),
|
|
every block is written to a separate, fully-closed mp4 the moment it is
|
|
decoded (so partial results survive an interruption), AND its frames are
|
|
appended to a single `full.mp4` so a ready-to-play complete video is also
|
|
produced. Segments are named by block index.
|
|
"""
|
|
def __init__(self, out_dir, fps, full_name="full.mp4"):
|
|
os.makedirs(out_dir, exist_ok=True)
|
|
self.out_dir = out_dir
|
|
self.fps = fps
|
|
self.segments = []
|
|
# Continuous writer for the assembled full video.
|
|
self.full_path = os.path.join(out_dir, full_name)
|
|
self.full_writer = imageio.get_writer(self.full_path, fps=fps)
|
|
self.num_frames = 0
|
|
|
|
def __call__(self, video_chunk, block_idx):
|
|
# video_chunk: [B, C, F, H, W] in [0, 1]; take the first sample.
|
|
chunk = video_chunk[0].permute(1, 2, 3, 0) # [F, H, W, C]
|
|
chunk = (chunk.clamp(0, 1) * 255).numpy().astype(np.uint8)
|
|
# 1) Standalone, fully-flushed segment file.
|
|
seg_path = os.path.join(self.out_dir, f"seg_{block_idx:04d}.mp4")
|
|
with imageio.get_writer(seg_path, fps=self.fps) as writer:
|
|
for frame in chunk:
|
|
writer.append_data(frame)
|
|
self.segments.append(seg_path)
|
|
# 2) Append the same frames to the continuous full video.
|
|
for frame in chunk:
|
|
self.full_writer.append_data(frame)
|
|
self.num_frames += 1
|
|
print(f"Saved segment to: {seg_path} ({len(chunk)} frames)")
|
|
|
|
def close(self):
|
|
self.full_writer.close()
|
|
print(f"Saved {len(self.segments)} segments to: {self.out_dir}")
|
|
print(f"Saved full video to: {self.full_path} ({self.num_frames} frames)")
|
|
|
|
def save_videos_with_audio_grid(
|
|
videos: torch.Tensor,
|
|
audio: torch.Tensor,
|
|
path: str,
|
|
fps: int = 24,
|
|
audio_sample_rate: int = 24000,
|
|
n_rows: int = 6,
|
|
rescale: bool = False
|
|
):
|
|
"""
|
|
Save video frames with audio to a single mp4 file.
|
|
|
|
Args:
|
|
videos: Video tensor of shape (b, c, t, h, w)
|
|
audio: Audio tensor
|
|
path: Output file path
|
|
fps: Frames per second
|
|
audio_sample_rate: Audio sample rate
|
|
n_rows: Number of rows for grid layout
|
|
rescale: Whether to rescale from [-1, 1] to [0, 1]
|
|
"""
|
|
from fractions import Fraction
|
|
|
|
import av
|
|
|
|
# Convert video frames to numpy arrays
|
|
# Support both [b, c, t, h, w] and [b, t, c, h, w]
|
|
if videos.shape[1] != 3: # shape[1] is T (frames), not C (channels)
|
|
videos = rearrange(videos, "b t c h w -> b c t h w")
|
|
videos = rearrange(videos, "b c t h w -> t b c h w")
|
|
frame_list = []
|
|
for x in videos:
|
|
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
|
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
|
if rescale:
|
|
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
|
x = (x * 255).cpu().numpy().astype(np.uint8)
|
|
frame_list.append(x)
|
|
|
|
# Handle single frame case (save as image)
|
|
if len(frame_list) == 1:
|
|
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
Image.fromarray(frame_list[0]).save(path.replace('.mp4', '.png'))
|
|
print(f"Saved image to: {path.replace('.mp4', '.png')}")
|
|
return
|
|
|
|
# Prepare audio tensor
|
|
audio_tensor = audio[0].float().cpu()
|
|
if audio_tensor.ndim == 1:
|
|
audio_tensor = audio_tensor.unsqueeze(-1)
|
|
elif audio_tensor.ndim == 2 and audio_tensor.shape[0] == 1:
|
|
# [1, N] -> [N, 1]
|
|
audio_tensor = audio_tensor.squeeze(0).unsqueeze(-1)
|
|
if audio_tensor.shape[1] != 2 and audio_tensor.shape[0] == 2:
|
|
audio_tensor = audio_tensor.T
|
|
if audio_tensor.shape[1] != 2:
|
|
# mono -> duplicate to stereo
|
|
audio_tensor = audio_tensor.expand(-1, 2)
|
|
|
|
# Create output directory
|
|
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
|
|
# Create video container
|
|
height, width = frame_list[0].shape[:2]
|
|
container = av.open(path, mode="w")
|
|
v_stream = container.add_stream("libx264", rate=int(fps))
|
|
v_stream.width = width
|
|
v_stream.height = height
|
|
v_stream.pix_fmt = "yuv420p"
|
|
|
|
# Create audio stream
|
|
a_stream = container.add_stream("aac", rate=audio_sample_rate)
|
|
a_stream.codec_context.sample_rate = audio_sample_rate
|
|
a_stream.codec_context.layout = "stereo"
|
|
a_stream.codec_context.time_base = Fraction(1, audio_sample_rate)
|
|
|
|
# Write video frames
|
|
for frame_np in frame_list:
|
|
frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
|
|
for pkt in v_stream.encode(frame):
|
|
container.mux(pkt)
|
|
for pkt in v_stream.encode():
|
|
container.mux(pkt)
|
|
|
|
# Write audio
|
|
samples = audio_tensor
|
|
if samples.dtype != torch.int16:
|
|
samples = torch.clip(samples, -1.0, 1.0)
|
|
samples = (samples * 32767.0).to(torch.int16)
|
|
|
|
frame_in = av.AudioFrame.from_ndarray(
|
|
samples.contiguous().reshape(1, -1).cpu().numpy(),
|
|
format="s16",
|
|
layout="stereo",
|
|
)
|
|
frame_in.sample_rate = audio_sample_rate
|
|
|
|
cc = a_stream.codec_context
|
|
target_format = cc.format or "fltp"
|
|
target_layout = cc.layout or "stereo"
|
|
target_rate = cc.sample_rate or frame_in.sample_rate
|
|
|
|
resampler = av.audio.resampler.AudioResampler(
|
|
format=target_format,
|
|
layout=target_layout,
|
|
rate=target_rate,
|
|
)
|
|
|
|
audio_next_pts = 0
|
|
for rframe in resampler.resample(frame_in):
|
|
if rframe.pts is None:
|
|
rframe.pts = audio_next_pts
|
|
audio_next_pts += rframe.samples
|
|
rframe.sample_rate = frame_in.sample_rate
|
|
container.mux(a_stream.encode(rframe))
|
|
|
|
for packet in a_stream.encode():
|
|
container.mux(packet)
|
|
|
|
container.close()
|
|
print(f"Saved video with audio to: {path}")
|
|
|
|
def merge_video_audio(video_path: str, audio_path: str):
|
|
"""
|
|
Merge the video and audio into a new video, with the duration set to the shorter of the two,
|
|
and overwrite the original video file.
|
|
|
|
Parameters:
|
|
video_path (str): Path to the original video file
|
|
audio_path (str): Path to the audio file
|
|
"""
|
|
# check
|
|
if not os.path.exists(video_path):
|
|
raise FileNotFoundError(f"video file {video_path} does not exist")
|
|
if not os.path.exists(audio_path):
|
|
raise FileNotFoundError(f"audio file {audio_path} does not exist")
|
|
|
|
base, ext = os.path.splitext(video_path)
|
|
temp_output = f"{base}_temp{ext}"
|
|
|
|
try:
|
|
# create ffmpeg command
|
|
command = [
|
|
'ffmpeg',
|
|
'-y', # overwrite
|
|
'-i',
|
|
video_path,
|
|
'-i',
|
|
audio_path,
|
|
'-c:v',
|
|
'copy', # copy video stream
|
|
'-c:a',
|
|
'aac', # use AAC audio encoder
|
|
'-b:a',
|
|
'192k', # set audio bitrate (optional)
|
|
'-map',
|
|
'0:v:0', # select the first video stream
|
|
'-map',
|
|
'1:a:0', # select the first audio stream
|
|
'-shortest', # choose the shortest duration
|
|
temp_output
|
|
]
|
|
|
|
# execute the command
|
|
print("Start merging video and audio...")
|
|
result = subprocess.run(
|
|
command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
|
|
|
|
# check result
|
|
if result.returncode != 0:
|
|
error_msg = f"FFmpeg execute failed: {result.stderr}"
|
|
print(error_msg)
|
|
raise RuntimeError(error_msg)
|
|
|
|
shutil.move(temp_output, video_path)
|
|
print(f"Merge completed, saved to {video_path}")
|
|
|
|
except Exception as e:
|
|
if os.path.exists(temp_output):
|
|
os.remove(temp_output)
|
|
print(f"merge_video_audio failed with error: {e}")
|
|
|
|
def calculate_dimensions(target_area, ratio):
|
|
width = math.sqrt(target_area * ratio)
|
|
height = width / ratio
|
|
|
|
width = round(width / 32) * 32
|
|
height = round(height / 32) * 32
|
|
|
|
return width, height
|
|
|
|
def get_image_to_video_latent(validation_image_start, validation_image_end, video_length, sample_size):
|
|
if validation_image_start is not None and validation_image_end is not None:
|
|
if type(validation_image_start) is str and os.path.isfile(validation_image_start):
|
|
image_start = clip_image = Image.open(validation_image_start).convert("RGB")
|
|
image_start = image_start.resize([sample_size[1], sample_size[0]])
|
|
clip_image = clip_image.resize([sample_size[1], sample_size[0]])
|
|
else:
|
|
image_start = clip_image = validation_image_start
|
|
image_start = [_image_start.resize([sample_size[1], sample_size[0]]) for _image_start in image_start]
|
|
clip_image = [_clip_image.resize([sample_size[1], sample_size[0]]) for _clip_image in clip_image]
|
|
|
|
if type(validation_image_end) is str and os.path.isfile(validation_image_end):
|
|
image_end = Image.open(validation_image_end).convert("RGB")
|
|
image_end = image_end.resize([sample_size[1], sample_size[0]])
|
|
else:
|
|
image_end = validation_image_end
|
|
image_end = [_image_end.resize([sample_size[1], sample_size[0]]) for _image_end in image_end]
|
|
|
|
if type(image_start) is list:
|
|
clip_image = clip_image[0]
|
|
start_video = torch.cat(
|
|
[torch.from_numpy(np.array(_image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_start in image_start],
|
|
dim=2
|
|
)
|
|
input_video = torch.tile(start_video[:, :, :1], [1, 1, video_length, 1, 1])
|
|
input_video[:, :, :len(image_start)] = start_video
|
|
|
|
input_video_mask = torch.zeros_like(input_video[:, :1])
|
|
input_video_mask[:, :, len(image_start):] = 255
|
|
else:
|
|
input_video = torch.tile(
|
|
torch.from_numpy(np.array(image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0),
|
|
[1, 1, video_length, 1, 1]
|
|
)
|
|
input_video_mask = torch.zeros_like(input_video[:, :1])
|
|
input_video_mask[:, :, 1:] = 255
|
|
|
|
if type(image_end) is list:
|
|
image_end = [_image_end.resize(image_start[0].size if type(image_start) is list else image_start.size) for _image_end in image_end]
|
|
end_video = torch.cat(
|
|
[torch.from_numpy(np.array(_image_end)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_end in image_end],
|
|
dim=2
|
|
)
|
|
input_video[:, :, -len(end_video):] = end_video
|
|
|
|
input_video_mask[:, :, -len(image_end):] = 0
|
|
else:
|
|
image_end = image_end.resize(image_start[0].size if type(image_start) is list else image_start.size)
|
|
input_video[:, :, -1:] = torch.from_numpy(np.array(image_end)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0)
|
|
input_video_mask[:, :, -1:] = 0
|
|
|
|
input_video = input_video / 255
|
|
|
|
elif validation_image_start is not None:
|
|
if type(validation_image_start) is str and os.path.isfile(validation_image_start):
|
|
image_start = clip_image = Image.open(validation_image_start).convert("RGB")
|
|
image_start = image_start.resize([sample_size[1], sample_size[0]])
|
|
clip_image = clip_image.resize([sample_size[1], sample_size[0]])
|
|
else:
|
|
image_start = clip_image = validation_image_start
|
|
image_start = [_image_start.resize([sample_size[1], sample_size[0]]) for _image_start in image_start]
|
|
clip_image = [_clip_image.resize([sample_size[1], sample_size[0]]) for _clip_image in clip_image]
|
|
image_end = None
|
|
|
|
if type(image_start) is list:
|
|
clip_image = clip_image[0]
|
|
start_video = torch.cat(
|
|
[torch.from_numpy(np.array(_image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0) for _image_start in image_start],
|
|
dim=2
|
|
)
|
|
input_video = torch.tile(start_video[:, :, :1], [1, 1, video_length, 1, 1])
|
|
input_video[:, :, :len(image_start)] = start_video
|
|
input_video = input_video / 255
|
|
|
|
input_video_mask = torch.zeros_like(input_video[:, :1])
|
|
input_video_mask[:, :, len(image_start):] = 255
|
|
else:
|
|
input_video = torch.tile(
|
|
torch.from_numpy(np.array(image_start)).permute(2, 0, 1).unsqueeze(1).unsqueeze(0),
|
|
[1, 1, video_length, 1, 1]
|
|
) / 255
|
|
input_video_mask = torch.zeros_like(input_video[:, :1])
|
|
input_video_mask[:, :, 1:, ] = 255
|
|
else:
|
|
image_start = None
|
|
image_end = None
|
|
input_video = torch.zeros([1, 3, video_length, sample_size[0], sample_size[1]])
|
|
input_video_mask = torch.ones([1, 1, video_length, sample_size[0], sample_size[1]]) * 255
|
|
clip_image = None
|
|
|
|
del image_start
|
|
del image_end
|
|
gc.collect()
|
|
|
|
return input_video, input_video_mask, clip_image
|
|
|
|
def get_video_to_video_latent(input_video_path, video_length, sample_size, fps=None, validation_video_mask=None, ref_image=None, keep_aspect_ratio=False):
|
|
if input_video_path is not None:
|
|
if isinstance(input_video_path, str):
|
|
cap = cv2.VideoCapture(input_video_path)
|
|
input_video = []
|
|
|
|
original_fps = cap.get(cv2.CAP_PROP_FPS)
|
|
# Resample onto the `fps` timeline by timestamp instead of dropping every n-th frame: target frame
|
|
# `next_target_frame` reads the source frame closest to `next_target_frame / fps` seconds, so a 30 fps
|
|
# clip keeps its real duration on a 24 fps request instead of stretching to 1.25x its length. An
|
|
# unreadable source rate reads every frame, as before.
|
|
resample_ratio = original_fps / fps if fps is not None and original_fps > 0 else None
|
|
|
|
frame_count = 0
|
|
next_target_frame = 0
|
|
|
|
while True:
|
|
ret, frame = cap.read()
|
|
if not ret:
|
|
break
|
|
|
|
if resample_ratio is None:
|
|
emit_count = 1
|
|
else:
|
|
# Every target frame snapping onto this source frame; a source slower than `fps` repeats it,
|
|
# a faster one drops the frames in between.
|
|
emit_count = 0
|
|
while int(round(next_target_frame * resample_ratio)) == frame_count:
|
|
emit_count += 1
|
|
next_target_frame += 1
|
|
|
|
if emit_count:
|
|
if keep_aspect_ratio:
|
|
# Cover the canvas and center-crop onto it, the resize + crop geometry of the training
|
|
# collates, instead of stretching the frames onto it.
|
|
source_height, source_width = frame.shape[:2]
|
|
scale = max(sample_size[0] / source_height, sample_size[1] / source_width)
|
|
resized = cv2.resize(frame, (int(round(source_width * scale)), int(round(source_height * scale))))
|
|
top = (resized.shape[0] - sample_size[0]) // 2
|
|
left = (resized.shape[1] - sample_size[1]) // 2
|
|
frame = resized[top : top + sample_size[0], left : left + sample_size[1]]
|
|
else:
|
|
frame = cv2.resize(frame, (sample_size[1], sample_size[0]))
|
|
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
|
for _ in range(emit_count):
|
|
input_video.append(frame)
|
|
|
|
frame_count += 1
|
|
|
|
cap.release()
|
|
else:
|
|
input_video = input_video_path
|
|
|
|
if video_length is not None:
|
|
input_video = torch.from_numpy(np.array(input_video))[:video_length]
|
|
else:
|
|
input_video = torch.from_numpy(np.array(input_video))
|
|
input_video = input_video.permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
|
|
if validation_video_mask is not None:
|
|
validation_video_mask = Image.open(validation_video_mask).convert('L').resize((sample_size[1], sample_size[0]))
|
|
input_video_mask = np.where(np.array(validation_video_mask) < 240, 0, 255)
|
|
|
|
input_video_mask = torch.from_numpy(np.array(input_video_mask)).unsqueeze(0).unsqueeze(-1).permute([3, 0, 1, 2]).unsqueeze(0)
|
|
input_video_mask = torch.tile(input_video_mask, [1, 1, input_video.size()[2], 1, 1])
|
|
input_video_mask = input_video_mask.to(input_video.device, input_video.dtype)
|
|
else:
|
|
input_video_mask = torch.zeros_like(input_video[:, :1])
|
|
input_video_mask[:, :, :] = 255
|
|
else:
|
|
input_video, input_video_mask = None, None
|
|
|
|
if ref_image is not None:
|
|
if isinstance(ref_image, str):
|
|
clip_image = Image.open(ref_image).convert("RGB")
|
|
else:
|
|
clip_image = Image.fromarray(np.array(ref_image, np.uint8))
|
|
else:
|
|
clip_image = None
|
|
|
|
if ref_image is not None:
|
|
if isinstance(ref_image, str):
|
|
ref_image = Image.open(ref_image).convert("RGB")
|
|
ref_image = ref_image.resize((sample_size[1], sample_size[0]))
|
|
ref_image = torch.from_numpy(np.array(ref_image))
|
|
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
else:
|
|
ref_image = torch.from_numpy(np.array(ref_image))
|
|
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
return input_video, input_video_mask, ref_image, clip_image
|
|
|
|
def get_image_latent(ref_image=None, sample_size=None, padding=False):
|
|
if ref_image is not None:
|
|
if isinstance(ref_image, str):
|
|
ref_image = Image.open(ref_image).convert("RGB")
|
|
if padding:
|
|
ref_image = padding_image(ref_image, sample_size[1], sample_size[0])
|
|
ref_image = ref_image.resize((sample_size[1], sample_size[0]))
|
|
ref_image = torch.from_numpy(np.array(ref_image))
|
|
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
elif isinstance(ref_image, Image.Image):
|
|
ref_image = ref_image.convert("RGB")
|
|
if padding:
|
|
ref_image = padding_image(ref_image, sample_size[1], sample_size[0])
|
|
ref_image = ref_image.resize((sample_size[1], sample_size[0]))
|
|
ref_image = torch.from_numpy(np.array(ref_image))
|
|
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
else:
|
|
ref_image = torch.from_numpy(np.array(ref_image))
|
|
ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255
|
|
|
|
return ref_image
|
|
|
|
def get_image(ref_image=None):
|
|
if ref_image is not None:
|
|
if isinstance(ref_image, str):
|
|
ref_image = Image.open(ref_image).convert("RGB")
|
|
elif isinstance(ref_image, Image.Image):
|
|
ref_image = ref_image.convert("RGB")
|
|
|
|
return ref_image
|
|
|
|
def padding_image(images, new_width, new_height):
|
|
new_image = Image.new('RGB', (new_width, new_height), (255, 255, 255))
|
|
|
|
aspect_ratio = images.width / images.height
|
|
if new_width / new_height > 1:
|
|
if aspect_ratio > new_width / new_height:
|
|
new_img_width = new_width
|
|
new_img_height = int(new_img_width / aspect_ratio)
|
|
else:
|
|
new_img_height = new_height
|
|
new_img_width = int(new_img_height * aspect_ratio)
|
|
else:
|
|
if aspect_ratio > new_width / new_height:
|
|
new_img_width = new_width
|
|
new_img_height = int(new_img_width / aspect_ratio)
|
|
else:
|
|
new_img_height = new_height
|
|
new_img_width = int(new_img_height * aspect_ratio)
|
|
|
|
resized_img = images.resize((new_img_width, new_img_height))
|
|
|
|
paste_x = (new_width - new_img_width) // 2
|
|
paste_y = (new_height - new_img_height) // 2
|
|
|
|
new_image.paste(resized_img, (paste_x, paste_y))
|
|
|
|
return new_image
|
|
|
|
def timer(func):
|
|
def wrapper(*args, **kwargs):
|
|
start_time = time.time()
|
|
result = func(*args, **kwargs)
|
|
end_time = time.time()
|
|
print(f"function {func.__name__} running for {end_time - start_time} seconds")
|
|
return result
|
|
return wrapper
|
|
|
|
def timer_record(model_name=""):
|
|
def decorator(func):
|
|
def wrapper(*args, **kwargs):
|
|
torch.cuda.synchronize()
|
|
start_time = time.time()
|
|
result = func(*args, **kwargs)
|
|
torch.cuda.synchronize()
|
|
end_time = time.time()
|
|
import torch.distributed as dist
|
|
if dist.is_initialized():
|
|
if dist.get_rank() == 0:
|
|
time_sum = end_time - start_time
|
|
print('# --------------------------------------------------------- #')
|
|
print(f'# {model_name} time: {time_sum}s')
|
|
print('# --------------------------------------------------------- #')
|
|
_write_to_excel(model_name, time_sum)
|
|
else:
|
|
time_sum = end_time - start_time
|
|
print('# --------------------------------------------------------- #')
|
|
print(f'# {model_name} time: {time_sum}s')
|
|
print('# --------------------------------------------------------- #')
|
|
_write_to_excel(model_name, time_sum)
|
|
return result
|
|
return wrapper
|
|
return decorator
|
|
|
|
def _write_to_excel(model_name, time_sum):
|
|
import os
|
|
|
|
import pandas as pd
|
|
|
|
row_env = os.environ.get(f"{model_name}_EXCEL_ROW", "1") # 默认第1行
|
|
col_env = os.environ.get(f"{model_name}_EXCEL_COL", "1") # 默认第A列
|
|
file_path = os.environ.get("EXCEL_FILE", "timing_records.xlsx") # 默认文件名
|
|
|
|
try:
|
|
df = pd.read_excel(file_path, sheet_name="Sheet1", header=None)
|
|
except FileNotFoundError:
|
|
df = pd.DataFrame()
|
|
|
|
row_idx = int(row_env)
|
|
col_idx = int(col_env)
|
|
|
|
if row_idx >= len(df):
|
|
df = pd.concat([df, pd.DataFrame([ [None] * (len(df.columns) if not df.empty else 0) ] * (row_idx - len(df) + 1))], ignore_index=True)
|
|
|
|
if col_idx >= len(df.columns):
|
|
df = pd.concat([df, pd.DataFrame(columns=range(len(df.columns), col_idx + 1))], axis=1)
|
|
|
|
df.iloc[row_idx, col_idx] = time_sum
|
|
|
|
df.to_excel(file_path, index=False, header=False, sheet_name="Sheet1")
|
|
|
|
def get_autocast_dtype():
|
|
try:
|
|
if not torch.cuda.is_available():
|
|
print("CUDA not available, using float16 by default.")
|
|
return torch.float16
|
|
|
|
device = torch.cuda.current_device()
|
|
prop = torch.cuda.get_device_properties(device)
|
|
|
|
print(f"GPU: {prop.name}, Compute Capability: {prop.major}.{prop.minor}")
|
|
|
|
if prop.major >= 8:
|
|
if torch.cuda.is_bf16_supported():
|
|
print("Using bfloat16.")
|
|
return torch.bfloat16
|
|
else:
|
|
print("Compute capability >= 8.0 but bfloat16 not supported, falling back to float16.")
|
|
return torch.float16
|
|
else:
|
|
print("GPU does not support bfloat16 natively, using float16.")
|
|
return torch.float16
|
|
|
|
except Exception as e:
|
|
print(f"Error detecting GPU capability: {e}, falling back to float16.")
|
|
return torch.float16 |