Update FantasyTalking training and Dataset Loading structure (#358)
This commit is contained in:
@@ -141,14 +141,20 @@ if transformer_path is not None:
|
||||
else:
|
||||
state_dict = torch.load(transformer_path, map_location="cpu")
|
||||
|
||||
audio_processor_dict = state_dict["audio_processor"] if "audio_processor" in state_dict else state_dict
|
||||
m, u = transformer.load_state_dict(audio_processor_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
if "audio_processor" in state_dict:
|
||||
audio_processor_dict = state_dict["audio_processor"] if "audio_processor" in state_dict else state_dict
|
||||
m, u = transformer.load_state_dict(audio_processor_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
proj_model_dict = state_dict["proj_model"] if "proj_model" in state_dict else state_dict
|
||||
proj_model_dict = {"proj_model." + k : v for k, v in proj_model_dict.items()}
|
||||
m, u = transformer.load_state_dict(proj_model_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
proj_model_dict = state_dict["proj_model"] if "proj_model" in state_dict else state_dict
|
||||
proj_model_dict = {"proj_model." + k : v for k, v in proj_model_dict.items()}
|
||||
m, u = transformer.load_state_dict(proj_model_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
else:
|
||||
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
|
||||
|
||||
m, u = transformer.load_state_dict(state_dict, strict=False)
|
||||
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
|
||||
|
||||
# Get Vae
|
||||
Chosen_AutoencoderKL = {
|
||||
|
||||
Executable
+220
@@ -0,0 +1,220 @@
|
||||
## Training Code
|
||||
|
||||
The default training commands for the different versions are as follows:
|
||||
|
||||
We can choose whether to use fsdp in FantasyTalking, which can save a lot of video memory.
|
||||
|
||||
The metadata_control.json is a little different from normal json in FantasyTalking, you need to add a audio_path.
|
||||
|
||||
```json
|
||||
[
|
||||
{
|
||||
"file_path": "train/00000001.mp4",
|
||||
"audio_path": "wav/00000001.wav",
|
||||
"text": "A group of young men in suits and sunglasses are walking down a city street.",
|
||||
"type": "video"
|
||||
},
|
||||
.....
|
||||
]
|
||||
```
|
||||
|
||||
Some parameters in the sh file can be confusing, and they are explained in this document:
|
||||
|
||||
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the videos at the center, but instead, it trains the videos after grouping them into buckets based on resolution.
|
||||
- `random_frame_crop` is used for random cropping on video frames to simulate videos with different frame counts.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for videos. When `random_hw_adapt` is enabled, for training videos, the height and width will be set to `video_sample_size` as the maximum and `512` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, with `video_sample_n_frames=49`, `video_sample_size=512`, the resolution of video inputs for training is `512x512x49`.
|
||||
- `training_with_video_token_length` specifies training the model according to token length. For training videos, the height and width will be set to `video_sample_size` as the maximum and `256` as the minimum.
|
||||
- For example, when `training_with_video_token_length` is enabled, with `video_sample_n_frames=49`, `token_sample_size=512`, `video_sample_size=256`, the resolution of image inputs for training is `256x256` to `1024x1024`, and the resolution of video inputs for training is `256x256x49` to `512x512x21`.
|
||||
- The token length for a video with dimensions 512x512 and 49 frames is 13,312. We need to set the `token_sample_size = 512`.
|
||||
- At 512x512 resolution, the number of video frames is 49 (~= 512 * 512 * 49 / 512 / 512).
|
||||
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
|
||||
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
|
||||
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
|
||||
FantasyTalking without deepspeed:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/fantasytalking/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_size=512 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--transformer_path="models/FantasyTalking/fantasytalking_model.ckpt" \
|
||||
--trainable_modules "processor." "proj_model."
|
||||
```
|
||||
|
||||
FantasyTalking with deepspeed zero-2:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/fantasytalking/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_size=512 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--transformer_path="models/FantasyTalking/fantasytalking_model.ckpt" \
|
||||
--trainable_modules "processor." "proj_model."
|
||||
```
|
||||
|
||||
FantasyTalking with deepspeed zero-3:
|
||||
|
||||
```sh
|
||||
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
|
||||
```
|
||||
|
||||
Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/fantasytalking/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_size=512 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--transformer_path="models/FantasyTalking/fantasytalking_model.ckpt" \
|
||||
--trainable_modules "processor." "proj_model."
|
||||
```
|
||||
|
||||
FantasyTalking with FSDP:
|
||||
|
||||
Wan with FSDP is suitable for 14B Wan at high resolutions. Training shell command is as follows:
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap=AudioAttentionBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/fantasytalking/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_size=512 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--transformer_path="models/FantasyTalking/fantasytalking_model.ckpt" \
|
||||
--trainable_modules "processor." "proj_model."
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,40 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/fantasytalking/train.py \
|
||||
--config_path="config/wan2.1/wan_civitai.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--video_sample_size=512 \
|
||||
--token_sample_size=512 \
|
||||
--video_sample_stride=2 \
|
||||
--video_sample_n_frames=81 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--random_hw_adapt \
|
||||
--training_with_video_token_length \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--low_vram \
|
||||
--transformer_path="models/FantasyTalking/fantasytalking_model.ckpt" \
|
||||
--trainable_modules "processor." "proj_model."
|
||||
@@ -0,0 +1,9 @@
|
||||
from .dataset_image import CC15M, ImageEditDataset
|
||||
from .dataset_image_video import (ImageVideoControlDataset, ImageVideoDataset,
|
||||
ImageVideoSampler)
|
||||
from .dataset_video import VideoDataset, VideoSpeechDataset, WebVid10M
|
||||
from .utils import (VIDEO_READER_TIMEOUT, Camera, VideoReader_contextmanager,
|
||||
custom_meshgrid, get_random_mask, get_relative_pose,
|
||||
get_video_reader_batch, padding_image, process_pose_file,
|
||||
process_pose_params, ray_condition, resize_frame,
|
||||
resize_image_with_target_area)
|
||||
@@ -182,7 +182,7 @@ class ImageEditDataset(Dataset):
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = CC15M(
|
||||
csv_path="/mnt_wg/zhoumo.xjq/CCUtils/cc15m_add_index.json",
|
||||
csv_path="./cc15m_add_index.json",
|
||||
resolution=512,
|
||||
)
|
||||
|
||||
|
||||
@@ -24,238 +24,12 @@ from safetensors.torch import load_file
|
||||
from torch.utils.data import BatchSampler, Sampler
|
||||
from torch.utils.data.dataset import Dataset
|
||||
|
||||
VIDEO_READER_TIMEOUT = 20
|
||||
from .utils import (VIDEO_READER_TIMEOUT, Camera, VideoReader_contextmanager,
|
||||
custom_meshgrid, get_random_mask, get_relative_pose,
|
||||
get_video_reader_batch, padding_image, process_pose_file,
|
||||
process_pose_params, ray_condition, resize_frame,
|
||||
resize_image_with_target_area)
|
||||
|
||||
def get_random_mask(shape, image_start_only=False):
|
||||
f, c, h, w = shape
|
||||
mask = torch.zeros((f, 1, h, w), dtype=torch.uint8)
|
||||
|
||||
if not image_start_only:
|
||||
if f != 1:
|
||||
mask_index = np.random.choice([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], p=[0.05, 0.2, 0.2, 0.2, 0.05, 0.05, 0.05, 0.1, 0.05, 0.05])
|
||||
else:
|
||||
mask_index = np.random.choice([0, 1], p = [0.2, 0.8])
|
||||
if mask_index == 0:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
mask[:, :, start_y:end_y, start_x:end_x] = 1
|
||||
elif mask_index == 1:
|
||||
mask[:, :, :, :] = 1
|
||||
elif mask_index == 2:
|
||||
mask_frame_index = np.random.randint(1, 5)
|
||||
mask[mask_frame_index:, :, :, :] = 1
|
||||
elif mask_index == 3:
|
||||
mask_frame_index = np.random.randint(1, 5)
|
||||
mask[mask_frame_index:-mask_frame_index, :, :, :] = 1
|
||||
elif mask_index == 4:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
|
||||
mask_frame_before = np.random.randint(0, f // 2)
|
||||
mask_frame_after = np.random.randint(f // 2, f)
|
||||
mask[mask_frame_before:mask_frame_after, :, start_y:end_y, start_x:end_x] = 1
|
||||
elif mask_index == 5:
|
||||
mask = torch.randint(0, 2, (f, 1, h, w), dtype=torch.uint8)
|
||||
elif mask_index == 6:
|
||||
num_frames_to_mask = random.randint(1, max(f // 2, 1))
|
||||
frames_to_mask = random.sample(range(f), num_frames_to_mask)
|
||||
|
||||
for i in frames_to_mask:
|
||||
block_height = random.randint(1, h // 4)
|
||||
block_width = random.randint(1, w // 4)
|
||||
top_left_y = random.randint(0, h - block_height)
|
||||
top_left_x = random.randint(0, w - block_width)
|
||||
mask[i, 0, top_left_y:top_left_y + block_height, top_left_x:top_left_x + block_width] = 1
|
||||
elif mask_index == 7:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
a = torch.randint(min(w, h) // 8, min(w, h) // 4, (1,)).item() # 长半轴
|
||||
b = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item() # 短半轴
|
||||
|
||||
for i in range(h):
|
||||
for j in range(w):
|
||||
if ((i - center_y) ** 2) / (b ** 2) + ((j - center_x) ** 2) / (a ** 2) < 1:
|
||||
mask[:, :, i, j] = 1
|
||||
elif mask_index == 8:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
radius = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item()
|
||||
for i in range(h):
|
||||
for j in range(w):
|
||||
if (i - center_y) ** 2 + (j - center_x) ** 2 < radius ** 2:
|
||||
mask[:, :, i, j] = 1
|
||||
elif mask_index == 9:
|
||||
for idx in range(f):
|
||||
if np.random.rand() > 0.5:
|
||||
mask[idx, :, :, :] = 1
|
||||
else:
|
||||
raise ValueError(f"The mask_index {mask_index} is not define")
|
||||
else:
|
||||
if f != 1:
|
||||
mask[1:, :, :, :] = 1
|
||||
else:
|
||||
mask[:, :, :, :] = 1
|
||||
return mask
|
||||
|
||||
class Camera(object):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
def __init__(self, entry):
|
||||
fx, fy, cx, cy = entry[1:5]
|
||||
self.fx = fx
|
||||
self.fy = fy
|
||||
self.cx = cx
|
||||
self.cy = cy
|
||||
w2c_mat = np.array(entry[7:]).reshape(3, 4)
|
||||
w2c_mat_4x4 = np.eye(4)
|
||||
w2c_mat_4x4[:3, :] = w2c_mat
|
||||
self.w2c_mat = w2c_mat_4x4
|
||||
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
|
||||
|
||||
def custom_meshgrid(*args):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
|
||||
if pver.parse(torch.__version__) < pver.parse('1.10'):
|
||||
return torch.meshgrid(*args)
|
||||
else:
|
||||
return torch.meshgrid(*args, indexing='ij')
|
||||
|
||||
def get_relative_pose(cam_params):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
|
||||
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
|
||||
cam_to_origin = 0
|
||||
target_cam_c2w = np.array([
|
||||
[1, 0, 0, 0],
|
||||
[0, 1, 0, -cam_to_origin],
|
||||
[0, 0, 1, 0],
|
||||
[0, 0, 0, 1]
|
||||
])
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
|
||||
ret_poses = np.array(ret_poses, dtype=np.float32)
|
||||
return ret_poses
|
||||
|
||||
def ray_condition(K, c2w, H, W, device):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
# c2w: B, V, 4, 4
|
||||
# K: B, V, 4
|
||||
|
||||
B = K.shape[0]
|
||||
|
||||
j, i = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
|
||||
)
|
||||
i = i.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
j = j.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
|
||||
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
|
||||
|
||||
zs = torch.ones_like(i) # [B, HxW]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
zs = zs.expand_as(ys)
|
||||
|
||||
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
|
||||
|
||||
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, 3, HW
|
||||
rays_o = c2w[..., :3, 3] # B, V, 3
|
||||
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, 3, HW
|
||||
# c2w @ dirctions
|
||||
rays_dxo = torch.cross(rays_o, rays_d)
|
||||
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
|
||||
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
|
||||
# plucker = plucker.permute(0, 1, 4, 2, 3)
|
||||
return plucker
|
||||
|
||||
def process_pose_file(pose_file_path, width=672, height=384, original_pose_width=1280, original_pose_height=720, device='cpu', return_poses=False):
|
||||
"""Modified from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
with open(pose_file_path, 'r') as f:
|
||||
poses = f.readlines()
|
||||
|
||||
poses = [pose.strip().split(' ') for pose in poses[1:]]
|
||||
cam_params = [[float(x) for x in pose] for pose in poses]
|
||||
if return_poses:
|
||||
return cam_params
|
||||
else:
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
|
||||
sample_wh_ratio = width / height
|
||||
pose_wh_ratio = original_pose_width / original_pose_height # Assuming placeholder ratios, change as needed
|
||||
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / width
|
||||
else:
|
||||
resized_ori_h = width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / height
|
||||
|
||||
intrinsic = np.asarray([[cam_param.fx * width,
|
||||
cam_param.fy * height,
|
||||
cam_param.cx * width,
|
||||
cam_param.cy * height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params) # Assuming this function is defined elsewhere
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, height, width, device=device)[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None]
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b f h w c")[0]
|
||||
return plucker_embedding
|
||||
|
||||
def process_pose_params(cam_params, width=672, height=384, original_pose_width=1280, original_pose_height=720, device='cpu'):
|
||||
"""Modified from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
|
||||
sample_wh_ratio = width / height
|
||||
pose_wh_ratio = original_pose_width / original_pose_height # Assuming placeholder ratios, change as needed
|
||||
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / width
|
||||
else:
|
||||
resized_ori_h = width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / height
|
||||
|
||||
intrinsic = np.asarray([[cam_param.fx * width,
|
||||
cam_param.fy * height,
|
||||
cam_param.cx * width,
|
||||
cam_param.cy * height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params) # Assuming this function is defined elsewhere
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, height, width, device=device)[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None]
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b f h w c")[0]
|
||||
return plucker_embedding
|
||||
|
||||
class ImageVideoSampler(BatchSampler):
|
||||
"""A sampler wrapper for grouping images with similar aspect ratio into a same batch.
|
||||
@@ -304,34 +78,6 @@ class ImageVideoSampler(BatchSampler):
|
||||
yield bucket[:]
|
||||
del bucket[:]
|
||||
|
||||
@contextmanager
|
||||
def VideoReader_contextmanager(*args, **kwargs):
|
||||
vr = VideoReader(*args, **kwargs)
|
||||
try:
|
||||
yield vr
|
||||
finally:
|
||||
del vr
|
||||
gc.collect()
|
||||
|
||||
def get_video_reader_batch(video_reader, batch_index):
|
||||
frames = video_reader.get_batch(batch_index).asnumpy()
|
||||
return frames
|
||||
|
||||
def resize_frame(frame, target_short_side):
|
||||
h, w, _ = frame.shape
|
||||
if h < w:
|
||||
if target_short_side > h:
|
||||
return frame
|
||||
new_h = target_short_side
|
||||
new_w = int(target_short_side * w / h)
|
||||
else:
|
||||
if target_short_side > w:
|
||||
return frame
|
||||
new_w = target_short_side
|
||||
new_h = int(target_short_side * h / w)
|
||||
|
||||
resized_frame = cv2.resize(frame, (new_w, new_h))
|
||||
return resized_frame
|
||||
|
||||
class ImageVideoDataset(Dataset):
|
||||
def __init__(
|
||||
@@ -513,33 +259,6 @@ class ImageVideoDataset(Dataset):
|
||||
|
||||
return sample
|
||||
|
||||
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
|
||||
|
||||
class ImageVideoControlDataset(Dataset):
|
||||
def __init__(
|
||||
@@ -556,6 +275,7 @@ class ImageVideoControlDataset(Dataset):
|
||||
enable_camera_info=False,
|
||||
return_file_name=False,
|
||||
enable_subject_info=False,
|
||||
padding_subject_info=True,
|
||||
):
|
||||
# Loading annotations from files
|
||||
print(f"loading annotations from {ann_path} ...")
|
||||
@@ -590,6 +310,7 @@ class ImageVideoControlDataset(Dataset):
|
||||
self.enable_inpaint = enable_inpaint
|
||||
self.enable_camera_info = enable_camera_info
|
||||
self.enable_subject_info = enable_subject_info
|
||||
self.padding_subject_info = padding_subject_info
|
||||
|
||||
self.video_length_drop_start = video_length_drop_start
|
||||
self.video_length_drop_end = video_length_drop_end
|
||||
@@ -757,11 +478,18 @@ class ImageVideoControlDataset(Dataset):
|
||||
width, height = subject_image.size
|
||||
total_pixels = width * height
|
||||
|
||||
img = padding_image(subject_image, visual_width, visual_height)
|
||||
if self.padding_subject_info:
|
||||
img = padding_image(subject_image, visual_width, visual_height)
|
||||
else:
|
||||
img = resize_image_with_target_area(subject_image, 1024 * 1024)
|
||||
|
||||
if random.random() < 0.5:
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
subject_images.append(img)
|
||||
subject_image = np.array(subject_images)
|
||||
subject_images.append(np.array(img))
|
||||
if self.padding_subject_info:
|
||||
subject_image = np.array(subject_images)
|
||||
else:
|
||||
subject_image = subject_images
|
||||
else:
|
||||
subject_image = None
|
||||
|
||||
@@ -806,15 +534,23 @@ class ImageVideoControlDataset(Dataset):
|
||||
width, height = subject_image.size
|
||||
total_pixels = width * height
|
||||
|
||||
img = padding_image(subject_image, visual_width, visual_height)
|
||||
if self.padding_subject_info:
|
||||
img = padding_image(subject_image, visual_width, visual_height)
|
||||
else:
|
||||
img = resize_image_with_target_area(subject_image, 1024 * 1024)
|
||||
|
||||
if random.random() < 0.5:
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
subject_images.append(img)
|
||||
subject_image = np.array(subject_images)
|
||||
subject_images.append(np.array(img))
|
||||
if self.padding_subject_info:
|
||||
subject_image = np.array(subject_images)
|
||||
else:
|
||||
subject_image = subject_images
|
||||
else:
|
||||
subject_image = None
|
||||
|
||||
return image, control_image, subject_image, None, text, 'image'
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from threading import Thread
|
||||
|
||||
import albumentations
|
||||
import cv2
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms as transforms
|
||||
@@ -20,61 +21,11 @@ from PIL import Image
|
||||
from torch.utils.data import BatchSampler, Sampler
|
||||
from torch.utils.data.dataset import Dataset
|
||||
|
||||
VIDEO_READER_TIMEOUT = 20
|
||||
|
||||
def get_random_mask(shape):
|
||||
f, c, h, w = shape
|
||||
|
||||
mask_index = np.random.randint(0, 4)
|
||||
mask = torch.zeros((f, 1, h, w), dtype=torch.uint8)
|
||||
if mask_index == 0:
|
||||
mask[1:, :, :, :] = 1
|
||||
elif mask_index == 1:
|
||||
mask_frame_index = 1
|
||||
mask[mask_frame_index:-mask_frame_index, :, :, :] = 1
|
||||
elif mask_index == 2:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
mask[:, :, start_y:end_y, start_x:end_x] = 1
|
||||
elif mask_index == 3:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
|
||||
mask_frame_before = np.random.randint(0, f // 2)
|
||||
mask_frame_after = np.random.randint(f // 2, f)
|
||||
mask[mask_frame_before:mask_frame_after, :, start_y:end_y, start_x:end_x] = 1
|
||||
else:
|
||||
raise ValueError(f"The mask_index {mask_index} is not define")
|
||||
return mask
|
||||
|
||||
|
||||
@contextmanager
|
||||
def VideoReader_contextmanager(*args, **kwargs):
|
||||
vr = VideoReader(*args, **kwargs)
|
||||
try:
|
||||
yield vr
|
||||
finally:
|
||||
del vr
|
||||
gc.collect()
|
||||
|
||||
|
||||
def get_video_reader_batch(video_reader, batch_index):
|
||||
frames = video_reader.get_batch(batch_index).asnumpy()
|
||||
return frames
|
||||
from .utils import (VIDEO_READER_TIMEOUT, Camera, VideoReader_contextmanager,
|
||||
custom_meshgrid, get_random_mask, get_relative_pose,
|
||||
get_video_reader_batch, padding_image, process_pose_file,
|
||||
process_pose_params, ray_condition, resize_frame,
|
||||
resize_image_with_target_area)
|
||||
|
||||
|
||||
class WebVid10M(Dataset):
|
||||
@@ -157,16 +108,16 @@ class WebVid10M(Dataset):
|
||||
class VideoDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
json_path, video_folder=None,
|
||||
ann_path, data_root=None,
|
||||
sample_size=256, sample_stride=4, sample_n_frames=16,
|
||||
enable_bucket=False, enable_inpaint=False
|
||||
):
|
||||
print(f"loading annotations from {json_path} ...")
|
||||
self.dataset = json.load(open(json_path, 'r'))
|
||||
print(f"loading annotations from {ann_path} ...")
|
||||
self.dataset = json.load(open(ann_path, 'r'))
|
||||
self.length = len(self.dataset)
|
||||
print(f"data scale: {self.length}")
|
||||
|
||||
self.video_folder = video_folder
|
||||
self.data_root = data_root
|
||||
self.sample_stride = sample_stride
|
||||
self.sample_n_frames = sample_n_frames
|
||||
self.enable_bucket = enable_bucket
|
||||
@@ -183,19 +134,25 @@ class VideoDataset(Dataset):
|
||||
|
||||
def get_batch(self, idx):
|
||||
video_dict = self.dataset[idx]
|
||||
video_id, name = video_dict['file_path'], video_dict['text']
|
||||
video_id, text = video_dict['file_path'], video_dict['text']
|
||||
|
||||
if self.video_folder is None:
|
||||
if self.data_root is None:
|
||||
video_dir = video_id
|
||||
else:
|
||||
video_dir = os.path.join(self.video_folder, video_id)
|
||||
video_dir = os.path.join(self.data_root, video_id)
|
||||
|
||||
with VideoReader_contextmanager(video_dir, num_threads=2) as video_reader:
|
||||
video_length = len(video_reader)
|
||||
|
||||
clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1)
|
||||
start_idx = random.randint(0, video_length - clip_length)
|
||||
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int)
|
||||
min_sample_n_frames = min(
|
||||
self.video_sample_n_frames,
|
||||
int(len(video_reader) * (self.video_length_drop_end - self.video_length_drop_start) // self.video_sample_stride)
|
||||
)
|
||||
if min_sample_n_frames == 0:
|
||||
raise ValueError(f"No Frames in video.")
|
||||
|
||||
video_length = int(self.video_length_drop_end * len(video_reader))
|
||||
clip_length = min(video_length, (min_sample_n_frames - 1) * self.video_sample_stride + 1)
|
||||
start_idx = random.randint(int(self.video_length_drop_start * video_length), video_length - clip_length) if video_length != clip_length else 0
|
||||
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, min_sample_n_frames, dtype=int)
|
||||
|
||||
try:
|
||||
sample_args = (video_reader, batch_index)
|
||||
@@ -214,29 +171,184 @@ class VideoDataset(Dataset):
|
||||
else:
|
||||
pixel_values = pixel_values
|
||||
|
||||
return pixel_values, name
|
||||
if not self.enable_bucket:
|
||||
pixel_values = self.video_transforms(pixel_values)
|
||||
|
||||
# Random use no text generation
|
||||
if random.random() < self.text_drop_ratio:
|
||||
text = ''
|
||||
return pixel_values, text
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
while True:
|
||||
sample = {}
|
||||
try:
|
||||
pixel_values, name = self.get_batch(idx)
|
||||
break
|
||||
sample["pixel_values"] = pixel_values
|
||||
sample["text"] = name
|
||||
sample["idx"] = idx
|
||||
if len(sample) > 0:
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
print("Error info:", e)
|
||||
print(e, self.dataset[idx % len(self.dataset)])
|
||||
idx = random.randint(0, self.length-1)
|
||||
|
||||
if not self.enable_bucket:
|
||||
pixel_values = self.pixel_transforms(pixel_values)
|
||||
if self.enable_inpaint:
|
||||
if self.enable_inpaint and not self.enable_bucket:
|
||||
mask = get_random_mask(pixel_values.size())
|
||||
mask_pixel_values = pixel_values * (1 - mask) + torch.ones_like(pixel_values) * -1 * mask
|
||||
sample = dict(pixel_values=pixel_values, mask_pixel_values=mask_pixel_values, mask=mask, text=name)
|
||||
mask_pixel_values = pixel_values * (1 - mask) + torch.zeros_like(pixel_values) * mask
|
||||
sample["mask_pixel_values"] = mask_pixel_values
|
||||
sample["mask"] = mask
|
||||
|
||||
clip_pixel_values = sample["pixel_values"][0].permute(1, 2, 0).contiguous()
|
||||
clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
|
||||
sample["clip_pixel_values"] = clip_pixel_values
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
class VideoSpeechDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
ann_path, data_root=None,
|
||||
video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16,
|
||||
enable_bucket=False, enable_inpaint=False,
|
||||
audio_sr=16000, # 新增:目标音频采样率
|
||||
text_drop_ratio=0.1 # 新增:文本丢弃概率
|
||||
):
|
||||
print(f"loading annotations from {ann_path} ...")
|
||||
self.dataset = json.load(open(ann_path, 'r'))
|
||||
self.length = len(self.dataset)
|
||||
print(f"data scale: {self.length}")
|
||||
|
||||
self.data_root = data_root
|
||||
self.video_sample_stride = video_sample_stride
|
||||
self.video_sample_n_frames = video_sample_n_frames
|
||||
self.enable_bucket = enable_bucket
|
||||
self.enable_inpaint = enable_inpaint
|
||||
self.audio_sr = audio_sr
|
||||
self.text_drop_ratio = text_drop_ratio
|
||||
|
||||
video_sample_size = tuple(video_sample_size) if not isinstance(video_sample_size, int) else (video_sample_size, video_sample_size)
|
||||
self.pixel_transforms = transforms.Compose(
|
||||
[
|
||||
transforms.Resize(video_sample_size[0]),
|
||||
transforms.CenterCrop(video_sample_size),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
]
|
||||
)
|
||||
|
||||
def get_batch(self, idx):
|
||||
video_dict = self.dataset[idx]
|
||||
video_id, text = video_dict['file_path'], video_dict['text']
|
||||
audio_id = video_dict['audio_path']
|
||||
|
||||
if self.data_root is None:
|
||||
video_path = video_id
|
||||
else:
|
||||
sample = dict(pixel_values=pixel_values, text=name)
|
||||
video_path = os.path.join(self.data_root, video_id)
|
||||
|
||||
if self.data_root is None:
|
||||
audio_path = audio_id
|
||||
else:
|
||||
audio_path = os.path.join(self.data_root, audio_id)
|
||||
|
||||
if not os.path.exists(audio_path):
|
||||
raise FileNotFoundError(f"Audio file not found for {video_path}")
|
||||
|
||||
with VideoReader_contextmanager(video_path, num_threads=2) as video_reader:
|
||||
total_frames = len(video_reader)
|
||||
fps = video_reader.get_avg_fps() # 获取原始视频帧率
|
||||
|
||||
# 计算实际采样的视频帧数(考虑边界)
|
||||
max_possible_frames = (total_frames - 1) // self.video_sample_stride + 1
|
||||
actual_n_frames = min(self.video_sample_n_frames, max_possible_frames)
|
||||
if actual_n_frames <= 0:
|
||||
raise ValueError(f"Video too short: {video_path}")
|
||||
|
||||
# 随机选择起始帧
|
||||
max_start = total_frames - (actual_n_frames - 1) * self.video_sample_stride - 1
|
||||
start_frame = random.randint(0, max_start) if max_start > 0 else 0
|
||||
frame_indices = [start_frame + i * self.video_sample_stride for i in range(actual_n_frames)]
|
||||
|
||||
# 读取视频帧
|
||||
try:
|
||||
sample_args = (video_reader, frame_indices)
|
||||
pixel_values = func_timeout(
|
||||
VIDEO_READER_TIMEOUT, get_video_reader_batch, args=sample_args
|
||||
)
|
||||
except FunctionTimedOut:
|
||||
raise ValueError(f"Read {idx} timeout.")
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to extract frames from video. Error is {e}.")
|
||||
|
||||
# 视频后处理
|
||||
if not self.enable_bucket:
|
||||
pixel_values = torch.from_numpy(pixel_values).permute(0, 3, 1, 2).contiguous()
|
||||
pixel_values = pixel_values / 255.
|
||||
pixel_values = self.pixel_transforms(pixel_values)
|
||||
|
||||
# === 新增:加载并截取对应音频 ===
|
||||
# 视频片段的起止时间(秒)
|
||||
start_time = start_frame / fps
|
||||
end_time = (start_frame + (actual_n_frames - 1) * self.video_sample_stride) / fps
|
||||
duration = end_time - start_time
|
||||
|
||||
# 使用 librosa 加载整个音频(或仅加载所需部分,但 librosa.load 不支持精确 seek,所以先加载再切)
|
||||
audio_input, sample_rate = librosa.load(audio_path, sr=self.audio_sr) # 重采样到目标 sr
|
||||
|
||||
# 转换为样本索引
|
||||
start_sample = int(start_time * self.audio_sr)
|
||||
end_sample = int(end_time * self.audio_sr)
|
||||
|
||||
# 安全截取
|
||||
if start_sample >= len(audio_input):
|
||||
# 音频太短,用零填充或截断
|
||||
audio_segment = np.zeros(int(duration * self.audio_sr), dtype=np.float32)
|
||||
else:
|
||||
audio_segment = audio_input[start_sample:end_sample]
|
||||
# 如果太短,补零
|
||||
target_len = int(duration * self.audio_sr)
|
||||
if len(audio_segment) < target_len:
|
||||
audio_segment = np.pad(audio_segment, (0, target_len - len(audio_segment)), mode='constant')
|
||||
|
||||
# === 文本随机丢弃 ===
|
||||
if random.random() < self.text_drop_ratio:
|
||||
text = ''
|
||||
|
||||
return pixel_values, text, audio_segment, sample_rate
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, idx):
|
||||
while True:
|
||||
sample = {}
|
||||
try:
|
||||
pixel_values, text, audio, sample_rate = self.get_batch(idx)
|
||||
sample["pixel_values"] = pixel_values
|
||||
sample["text"] = text
|
||||
sample["audio"] = torch.from_numpy(audio).float() # 转为 tensor
|
||||
sample["sample_rate"] = sample_rate
|
||||
sample["idx"] = idx
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"Error processing {idx}: {e}, retrying with random idx...")
|
||||
idx = random.randint(0, self.length - 1)
|
||||
|
||||
if self.enable_inpaint and not self.enable_bucket:
|
||||
mask = get_random_mask(pixel_values.size(), image_start_only=True)
|
||||
mask_pixel_values = pixel_values * (1 - mask) + torch.zeros_like(pixel_values) * mask
|
||||
sample["mask_pixel_values"] = mask_pixel_values
|
||||
sample["mask"] = mask
|
||||
|
||||
clip_pixel_values = sample["pixel_values"][0].permute(1, 2, 0).contiguous()
|
||||
clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255
|
||||
sample["clip_pixel_values"] = clip_pixel_values
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,347 @@
|
||||
import csv
|
||||
import gc
|
||||
import io
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from contextlib import contextmanager
|
||||
from random import shuffle
|
||||
from threading import Thread
|
||||
|
||||
import albumentations
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as transforms
|
||||
from decord import VideoReader
|
||||
from einops import rearrange
|
||||
from func_timeout import FunctionTimedOut, func_timeout
|
||||
from packaging import version as pver
|
||||
from PIL import Image
|
||||
from safetensors.torch import load_file
|
||||
from torch.utils.data import BatchSampler, Sampler
|
||||
from torch.utils.data.dataset import Dataset
|
||||
|
||||
VIDEO_READER_TIMEOUT = 20
|
||||
|
||||
def get_random_mask(shape, image_start_only=False):
|
||||
f, c, h, w = shape
|
||||
mask = torch.zeros((f, 1, h, w), dtype=torch.uint8)
|
||||
|
||||
if not image_start_only:
|
||||
if f != 1:
|
||||
mask_index = np.random.choice([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], p=[0.05, 0.2, 0.2, 0.2, 0.05, 0.05, 0.05, 0.1, 0.05, 0.05])
|
||||
else:
|
||||
mask_index = np.random.choice([0, 1, 7, 8], p = [0.2, 0.7, 0.05, 0.05])
|
||||
if mask_index == 0:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
mask[:, :, start_y:end_y, start_x:end_x] = 1
|
||||
elif mask_index == 1:
|
||||
mask[:, :, :, :] = 1
|
||||
elif mask_index == 2:
|
||||
mask_frame_index = np.random.randint(1, 5)
|
||||
mask[mask_frame_index:, :, :, :] = 1
|
||||
elif mask_index == 3:
|
||||
mask_frame_index = np.random.randint(1, 5)
|
||||
mask[mask_frame_index:-mask_frame_index, :, :, :] = 1
|
||||
elif mask_index == 4:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
block_size_x = torch.randint(w // 4, w // 4 * 3, (1,)).item() # 方块的宽度范围
|
||||
block_size_y = torch.randint(h // 4, h // 4 * 3, (1,)).item() # 方块的高度范围
|
||||
|
||||
start_x = max(center_x - block_size_x // 2, 0)
|
||||
end_x = min(center_x + block_size_x // 2, w)
|
||||
start_y = max(center_y - block_size_y // 2, 0)
|
||||
end_y = min(center_y + block_size_y // 2, h)
|
||||
|
||||
mask_frame_before = np.random.randint(0, f // 2)
|
||||
mask_frame_after = np.random.randint(f // 2, f)
|
||||
mask[mask_frame_before:mask_frame_after, :, start_y:end_y, start_x:end_x] = 1
|
||||
elif mask_index == 5:
|
||||
mask = torch.randint(0, 2, (f, 1, h, w), dtype=torch.uint8)
|
||||
elif mask_index == 6:
|
||||
num_frames_to_mask = random.randint(1, max(f // 2, 1))
|
||||
frames_to_mask = random.sample(range(f), num_frames_to_mask)
|
||||
|
||||
for i in frames_to_mask:
|
||||
block_height = random.randint(1, h // 4)
|
||||
block_width = random.randint(1, w // 4)
|
||||
top_left_y = random.randint(0, h - block_height)
|
||||
top_left_x = random.randint(0, w - block_width)
|
||||
mask[i, 0, top_left_y:top_left_y + block_height, top_left_x:top_left_x + block_width] = 1
|
||||
elif mask_index == 7:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
a = torch.randint(min(w, h) // 8, min(w, h) // 4, (1,)).item() # 长半轴
|
||||
b = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item() # 短半轴
|
||||
|
||||
for i in range(h):
|
||||
for j in range(w):
|
||||
if ((i - center_y) ** 2) / (b ** 2) + ((j - center_x) ** 2) / (a ** 2) < 1:
|
||||
mask[:, :, i, j] = 1
|
||||
elif mask_index == 8:
|
||||
center_x = torch.randint(0, w, (1,)).item()
|
||||
center_y = torch.randint(0, h, (1,)).item()
|
||||
radius = torch.randint(min(h, w) // 8, min(h, w) // 4, (1,)).item()
|
||||
for i in range(h):
|
||||
for j in range(w):
|
||||
if (i - center_y) ** 2 + (j - center_x) ** 2 < radius ** 2:
|
||||
mask[:, :, i, j] = 1
|
||||
elif mask_index == 9:
|
||||
for idx in range(f):
|
||||
if np.random.rand() > 0.5:
|
||||
mask[idx, :, :, :] = 1
|
||||
else:
|
||||
raise ValueError(f"The mask_index {mask_index} is not define")
|
||||
else:
|
||||
if f != 1:
|
||||
mask[1:, :, :, :] = 1
|
||||
else:
|
||||
mask[:, :, :, :] = 1
|
||||
return mask
|
||||
|
||||
@contextmanager
|
||||
def VideoReader_contextmanager(*args, **kwargs):
|
||||
vr = VideoReader(*args, **kwargs)
|
||||
try:
|
||||
yield vr
|
||||
finally:
|
||||
del vr
|
||||
gc.collect()
|
||||
|
||||
def get_video_reader_batch(video_reader, batch_index):
|
||||
frames = video_reader.get_batch(batch_index).asnumpy()
|
||||
return frames
|
||||
|
||||
def resize_frame(frame, target_short_side):
|
||||
h, w, _ = frame.shape
|
||||
if h < w:
|
||||
if target_short_side > h:
|
||||
return frame
|
||||
new_h = target_short_side
|
||||
new_w = int(target_short_side * w / h)
|
||||
else:
|
||||
if target_short_side > w:
|
||||
return frame
|
||||
new_w = target_short_side
|
||||
new_h = int(target_short_side * h / w)
|
||||
|
||||
resized_frame = cv2.resize(frame, (new_w, new_h))
|
||||
return resized_frame
|
||||
|
||||
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 resize_image_with_target_area(img: Image.Image, target_area: int = 1024 * 1024) -> Image.Image:
|
||||
"""
|
||||
将 PIL 图像缩放到接近指定像素面积(target_area),保持原始宽高比,
|
||||
并确保新宽度和高度均为 32 的整数倍。
|
||||
|
||||
参数:
|
||||
img (PIL.Image.Image): 输入图像
|
||||
target_area (int): 目标像素总面积,例如 1024*1024 = 1048576
|
||||
|
||||
返回:
|
||||
PIL.Image.Image: Resize 后的图像
|
||||
"""
|
||||
orig_w, orig_h = img.size
|
||||
if orig_w == 0 or orig_h == 0:
|
||||
raise ValueError("Input image has zero width or height.")
|
||||
|
||||
ratio = orig_w / orig_h
|
||||
ideal_width = math.sqrt(target_area * ratio)
|
||||
ideal_height = ideal_width / ratio
|
||||
|
||||
new_width = round(ideal_width / 32) * 32
|
||||
new_height = round(ideal_height / 32) * 32
|
||||
|
||||
new_width = max(32, new_width)
|
||||
new_height = max(32, new_height)
|
||||
|
||||
new_width = int(new_width)
|
||||
new_height = int(new_height)
|
||||
|
||||
resized_img = img.resize((new_width, new_height), Image.LANCZOS)
|
||||
return resized_img
|
||||
|
||||
class Camera(object):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
def __init__(self, entry):
|
||||
fx, fy, cx, cy = entry[1:5]
|
||||
self.fx = fx
|
||||
self.fy = fy
|
||||
self.cx = cx
|
||||
self.cy = cy
|
||||
w2c_mat = np.array(entry[7:]).reshape(3, 4)
|
||||
w2c_mat_4x4 = np.eye(4)
|
||||
w2c_mat_4x4[:3, :] = w2c_mat
|
||||
self.w2c_mat = w2c_mat_4x4
|
||||
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
|
||||
|
||||
def custom_meshgrid(*args):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
|
||||
if pver.parse(torch.__version__) < pver.parse('1.10'):
|
||||
return torch.meshgrid(*args)
|
||||
else:
|
||||
return torch.meshgrid(*args, indexing='ij')
|
||||
|
||||
def get_relative_pose(cam_params):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
|
||||
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
|
||||
cam_to_origin = 0
|
||||
target_cam_c2w = np.array([
|
||||
[1, 0, 0, 0],
|
||||
[0, 1, 0, -cam_to_origin],
|
||||
[0, 0, 1, 0],
|
||||
[0, 0, 0, 1]
|
||||
])
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
|
||||
ret_poses = np.array(ret_poses, dtype=np.float32)
|
||||
return ret_poses
|
||||
|
||||
def ray_condition(K, c2w, H, W, device):
|
||||
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
# c2w: B, V, 4, 4
|
||||
# K: B, V, 4
|
||||
|
||||
B = K.shape[0]
|
||||
|
||||
j, i = custom_meshgrid(
|
||||
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
|
||||
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
|
||||
)
|
||||
i = i.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
j = j.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
|
||||
|
||||
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
|
||||
|
||||
zs = torch.ones_like(i) # [B, HxW]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
zs = zs.expand_as(ys)
|
||||
|
||||
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
|
||||
|
||||
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, 3, HW
|
||||
rays_o = c2w[..., :3, 3] # B, V, 3
|
||||
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, 3, HW
|
||||
# c2w @ dirctions
|
||||
rays_dxo = torch.cross(rays_o, rays_d)
|
||||
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
|
||||
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
|
||||
# plucker = plucker.permute(0, 1, 4, 2, 3)
|
||||
return plucker
|
||||
|
||||
def process_pose_file(pose_file_path, width=672, height=384, original_pose_width=1280, original_pose_height=720, device='cpu', return_poses=False):
|
||||
"""Modified from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
with open(pose_file_path, 'r') as f:
|
||||
poses = f.readlines()
|
||||
|
||||
poses = [pose.strip().split(' ') for pose in poses[1:]]
|
||||
cam_params = [[float(x) for x in pose] for pose in poses]
|
||||
if return_poses:
|
||||
return cam_params
|
||||
else:
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
|
||||
sample_wh_ratio = width / height
|
||||
pose_wh_ratio = original_pose_width / original_pose_height # Assuming placeholder ratios, change as needed
|
||||
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / width
|
||||
else:
|
||||
resized_ori_h = width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / height
|
||||
|
||||
intrinsic = np.asarray([[cam_param.fx * width,
|
||||
cam_param.fy * height,
|
||||
cam_param.cx * width,
|
||||
cam_param.cy * height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params) # Assuming this function is defined elsewhere
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, height, width, device=device)[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None]
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b f h w c")[0]
|
||||
return plucker_embedding
|
||||
|
||||
def process_pose_params(cam_params, width=672, height=384, original_pose_width=1280, original_pose_height=720, device='cpu'):
|
||||
"""Modified from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
|
||||
"""
|
||||
cam_params = [Camera(cam_param) for cam_param in cam_params]
|
||||
|
||||
sample_wh_ratio = width / height
|
||||
pose_wh_ratio = original_pose_width / original_pose_height # Assuming placeholder ratios, change as needed
|
||||
|
||||
if pose_wh_ratio > sample_wh_ratio:
|
||||
resized_ori_w = height * pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fx = resized_ori_w * cam_param.fx / width
|
||||
else:
|
||||
resized_ori_h = width / pose_wh_ratio
|
||||
for cam_param in cam_params:
|
||||
cam_param.fy = resized_ori_h * cam_param.fy / height
|
||||
|
||||
intrinsic = np.asarray([[cam_param.fx * width,
|
||||
cam_param.fy * height,
|
||||
cam_param.cx * width,
|
||||
cam_param.cy * height]
|
||||
for cam_param in cam_params], dtype=np.float32)
|
||||
|
||||
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
|
||||
c2ws = get_relative_pose(cam_params) # Assuming this function is defined elsewhere
|
||||
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
|
||||
plucker_embedding = ray_condition(K, c2ws, height, width, device=device)[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
|
||||
plucker_embedding = plucker_embedding[None]
|
||||
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b f h w c")[0]
|
||||
return plucker_embedding
|
||||
@@ -38,6 +38,15 @@ class FantasyTalkingAudioEncoder(ModelMixin, ConfigMixin, FromOriginalModelMixin
|
||||
audio_segment, sampling_rate=sample_rate, return_tensors="pt"
|
||||
).input_values.to(self.model.device, self.model.dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
fea = self.model(input_values).last_hidden_state
|
||||
return fea
|
||||
|
||||
def extract_audio_feat_without_file_load(self, audio_segment, sample_rate):
|
||||
input_values = self.processor(
|
||||
audio_segment, sampling_rate=sample_rate, return_tensors="pt"
|
||||
).input_values.to(self.model.device, self.model.dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
fea = self.model(input_values).last_hidden_state
|
||||
return fea
|
||||
@@ -695,7 +695,7 @@ class FantasyTalkingPipeline(DiffusionPipeline):
|
||||
)
|
||||
|
||||
audio_scale = torch.tensor(
|
||||
[0, 1]
|
||||
[0.75, 1]
|
||||
).to(latent_model_input.device, latent_model_input.dtype)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
|
||||
Reference in New Issue
Block a user