Merge pull request #2 from smthemex/Pr1

init
This commit is contained in:
smthemex
2026-02-06 09:36:39 +08:00
committed by GitHub
82 changed files with 11859 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
*.mp4 filter=lfs diff=lfs merge=lfs -text
+8
View File
@@ -0,0 +1,8 @@
__pyc*
*/__pyc*
*/*/__pyc*
readme2*
results
temp
hostfile*
env.txt
Binary file not shown.
@@ -0,0 +1,71 @@
[
{
"mp4_path": null,
"mp4_img": "./InterDemo/ActionAndSong/img/001.png",
"dwpose_img": "./InterDemo/ActionAndSong/dwpose/001.png",
"dwpose_mp4": null,
"prompt": [
"Make an OK sign with one hand",
"Clap hands slowly with both hands"
],
"detailed_prompt": [
null,
null
],
"structured_prompt": [
"(Raise one hand) (Form an OK sign with fingers)",
"(Bring hands apart) (Clap hands slowly with both hands)"
],
"prompt_zh": "一只手比OK的手势 双手鼓掌",
"obj_img": null,
"obj_name": "None",
"video_id": "001_ok_clap",
"audio": "./InterDemo/ActionAndSong/audio/001.WAV"
},
{
"mp4_path": null,
"mp4_img": "./InterDemo/ActionAndSong/img/002.png",
"dwpose_img": "./InterDemo/ActionAndSong/dwpose/002.png",
"dwpose_mp4": null,
"prompt": [
"Wave slowly with both hands for greeting",
"Give a thumbs-up with one hand"
],
"detailed_prompt": [
null,
null
],
"structured_prompt": [
"(Raise both hands slowly) (Wave both hands side to side for greeting)",
"(Make a fist) (Extend the thumb upwards)"
],
"prompt_zh": "两只手打招呼 伸出大拇指点赞",
"obj_img": null,
"obj_name": "None",
"video_id": "002_greet_thumb",
"audio": "./InterDemo/ActionAndSong/audio/002.WAV"
},
{
"mp4_path": null,
"mp4_img": "./InterDemo/ActionAndSong/img/003.png",
"dwpose_img": "./InterDemo/ActionAndSong/dwpose/003.png",
"dwpose_mp4": null,
"prompt": [
"Hand clenches in a fist to cheer",
"Rest chin on one hand in thought"
],
"detailed_prompt": [
null,
null
],
"structured_prompt": [
"(Hand clenches in a fist) (Cheer with clenched fist)",
"(Raise one hand slowly to the chin) (Rest chin on the hand)"
],
"prompt_zh": "一只手握拳加油 一只手托下巴思考",
"obj_img": null,
"obj_name": "None",
"video_id": "003_cheer_rest",
"audio": "./InterDemo/ActionAndSong/audio/003.WAV"
}
]
+44
View File
@@ -0,0 +1,44 @@
[
{
"mp4_path": null,
"mp4_img": "./InterDemo/TIA2MV/img/001.jpeg",
"dwpose_img": "./InterDemo/TIA2MV/dwpose/001_obj_handbag.png",
"dwpose_mp4": null,
"prompt": "The man picks up his handbag.",
"detailed_prompt": "The man holds the handle of the handbag with his left hand while his right hand hangs naturally.",
"structured_prompt": "(Left hand holds the handle of the handbag)(Right hand hangs naturally)",
"prompt_zh": "男人提起手提包",
"obj_img": "./InterDemo/TIA2MV/obj_img/001/handbag_0_0.jpg",
"obj_name": "handbag",
"speech_path": "./InterDemo/TIA2MV/audio/001_0_spoken_text_en_medium/01-seedtts-01_promptvn.wav",
"video_id": "001_inter_handbag"
},
{
"mp4_path": null,
"mp4_img": "./InterDemo/TIA2MV/img/002.jpeg",
"dwpose_img": "./InterDemo/TIA2MV/dwpose/002_obj_camera.png",
"dwpose_mp4": null,
"prompt": "The man picks up the camera next to him to take a picture.",
"detailed_prompt": "The man first uses his right hand to pick up the camera on the table, then aims the camera at the lens, preparing to take a photo.",
"structured_prompt": "(Right hand picks up the camera on the table)(Aims the camera at the lens)",
"prompt_zh": "男人拿起手边的相机拍照",
"obj_img": "./InterDemo/TIA2MV/obj_img/002/camera_0_0.jpg",
"obj_name": "camera",
"speech_path": "./InterDemo/TIA2MV/audio/002_0_spoken_text_en_medium/01-seedtts-01_promptvn.wav",
"video_id": "002_inter_camera"
},
{
"mp4_path": null,
"mp4_img": "./InterDemo/TIA2MV/img/003.jpeg",
"dwpose_img": "./InterDemo/TIA2MV/dwpose/003_obj_coffee.png",
"dwpose_mp4": null,
"prompt": "The woman took the coffee from the table and drank it.",
"detailed_prompt": "The woman first moved her right hand away from the coffee cup, then picked up the coffee cup with her left hand, brought the cup close to her mouth, and finally took a sip.",
"structured_prompt": "(Moved her right hand away from the coffee cup)(Picked up the coffee cup with her left hand)(Brought the cup close to her mouth)(Took a sip)",
"prompt_zh": "女人拿起桌上的咖啡喝",
"obj_img": "./InterDemo/TIA2MV/obj_img/003/coffee_0_0.jpg",
"obj_name": "coffee",
"speech_path": "./InterDemo/TIA2MV/audio/003_1_spoken_text_en_medium/02-seedtts-02_promptvn.wav",
"video_id": "003_inter_coffee"
}
]
Binary file not shown.

After

Width:  |  Height:  |  Size: 12 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 13 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 11 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.0 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 864 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 143 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 13 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 37 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 37 KiB

+47
View File
@@ -0,0 +1,47 @@
[
{
"mp4_path": "./InterDemo/TIAP2V/mp4/self_collect_20250508_trans_slpit_self_collect_20250508_24_000_crop_0125.mp4",
"mp4_img": null,
"dwpose_img": null,
"dwpose_mp4": "./InterDemo/TIAP2V/dwpose/self_collect_20250508_trans_slpit_self_collect_20250508_24_000_crop_0125.mp4",
"prompt": "A woman with long dark hair, wearing large gold hoop earrings and a navy blue hoodie, sits holding a thin, flat, black rectangular object with vertical ridges. She is indoors in a room with white walls, a wooden cabinet visible behind her, and several cardboard boxes on the floor around her.",
"detailed_prompt": null,
"structured_prompt": null,
"prompt_zh": "A woman with long dark hair, wearing large gold hoop earrings and a navy blue hoodie, sits holding a thin, flat, black rectangular object with vertical ridges. She is indoors in a room with white walls, a wooden cabinet visible behind her, and several cardboard boxes on the floor around her.",
"obj_img": "./InterDemo/TIAP2V/obj_img/self_collect_20250508_trans_slpit_self_collect_20250508_24_000_crop_0125_##frame##_28_##object##_laptop.png",
"obj_name": "laptop",
"speech_path": "./InterDemo/TIAP2V/speech/bzya5p228n0xlz96_clip_0000_118.wav",
"video_id": "self_collect_20250508_trans_slpit_self_collect_20250508_24_000_crop_0125",
"homa_img": "/apdcephfs_cq8/share_1367250/ziyaohuang/projects/data_pipeline/testdata_processor/testdata-selfcollect_20250508/pasted_human_images_whiteback/self_collect_20250508_trans_slpit_self_collect_20250508_24_000_crop_0125_##frame##_28_##human##_laptop.png"
},
{
"mp4_path": "./InterDemo/TIAP2V/mp4/self_collect_20250508_trans_slpit_self_collect_20250508_25_000_crop_0050.mp4",
"mp4_img": null,
"dwpose_img": null,
"dwpose_mp4": "./InterDemo/TIAP2V/dwpose/self_collect_20250508_trans_slpit_self_collect_20250508_25_000_crop_0050.mp4",
"prompt": "An Asian male wearing a blue sweater, purple shirt, and black baseball cap holds up a black Nike Dunk High sneaker. He is wearing glasses and has short black hair. Behind him are several stuffed animals, including Mickey Mouse, and framed pictures on a light blue wall.",
"detailed_prompt": null,
"structured_prompt": null,
"prompt_zh": "An Asian male wearing a blue sweater, purple shirt, and black baseball cap holds up a black Nike Dunk High sneaker. He is wearing glasses and has short black hair. Behind him are several stuffed animals, including Mickey Mouse, and framed pictures on a light blue wall.",
"obj_img": "./InterDemo/TIAP2V/obj_img/self_collect_20250508_trans_slpit_self_collect_20250508_25_000_crop_0050_##frame##_44_##object##_shoe.png",
"obj_name": "shoe",
"speech_path": "./InterDemo/TIAP2V/speech/zlhjs3erm2y06y5v_clip_0000_20.wav",
"video_id": "self_collect_20250508_trans_slpit_self_collect_20250508_25_000_crop_0050",
"homa_img": "/apdcephfs_cq8/share_1367250/ziyaohuang/projects/data_pipeline/testdata_processor/testdata-selfcollect_20250508/pasted_human_images_whiteback/self_collect_20250508_trans_slpit_self_collect_20250508_25_000_crop_0050_##frame##_44_##human##_shoe.png"
},
{
"mp4_path": "./InterDemo/TIAP2V/mp4/self_collect_20250508_trans_slpit_self_collect_20250508_29_000_crop_0025.mp4",
"mp4_img": null,
"dwpose_img": null,
"dwpose_mp4": "./InterDemo/TIAP2V/dwpose/self_collect_20250508_trans_slpit_self_collect_20250508_29_000_crop_0025.mp4",
"prompt": "An Asian woman with fair skin, short black hair, and wearing a beige sweater sits indoors. She is holding a white Nintendo Switch Lite console in her hands. She has a silver necklace with a round pendant and bracelets on both wrists. Behind her are two windows, a white dresser, and a white shelf unit.",
"detailed_prompt": null,
"structured_prompt": null,
"prompt_zh": "An Asian woman with fair skin, short black hair, and wearing a beige sweater sits indoors. She is holding a white Nintendo Switch Lite console in her hands. She has a silver necklace with a round pendant and bracelets on both wrists. Behind her are two windows, a white dresser, and a white shelf unit.",
"obj_img": "./InterDemo/TIAP2V/obj_img/self_collect_20250508_trans_slpit_self_collect_20250508_29_000_crop_0025_##frame##_108_##object##_unknown_3.png",
"obj_name": "3",
"speech_path": "./InterDemo/TIAP2V/speech/fexqre0ojfx8x7s5_clip_0000_19.wav",
"video_id": "self_collect_20250508_trans_slpit_self_collect_20250508_29_000_crop_0025",
"homa_img": "/apdcephfs_cq8/share_1367250/ziyaohuang/projects/data_pipeline/testdata_processor/testdata-selfcollect_20250508/pasted_human_images_whiteback/self_collect_20250508_trans_slpit_self_collect_20250508_29_000_crop_0025_##frame##_108_##human##_unknown_3.png"
}
]
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:63b246bbcaf71295d2fe2a9a23102f679f00b9d2172adc3aaf6ce0e308324f5f
size 550358
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:aa10c01b5e62d397b2c8923f2b4b7a882a95baa89139174cdf6696e24b10c913
size 395180
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:70916b7c880774999c98ec1962b32d12be6a3b70fddb35ba21d002d066529729
size 549991
+189
View File
@@ -0,0 +1,189 @@
# !/usr/bin/env python
# -*- coding: UTF-8 -*-
import numpy as np
import torch
import os
import folder_paths
from typing_extensions import override
from comfy_api.latest import ComfyExtension, io
from .test_wanx_tia2mv_obj_back import generate_video,perdata,load_model,get_z_scale
from .model_loader_utils import trans2path,tensor2pillist,clear_comfyui_cache,get_wav2vec_repo
import nodes
MAX_SEED = np.iinfo(np.int32).max
device = torch.device(
"cuda:0") if torch.cuda.is_available() else torch.device(
"mps") if torch.backends.mps.is_available() else torch.device("cpu")
node_cr_path = os.path.dirname(os.path.abspath(__file__))
weigths_gguf_current_path = os.path.join(folder_paths.models_dir, "gguf")
if not os.path.exists(weigths_gguf_current_path):
os.makedirs(weigths_gguf_current_path)
folder_paths.add_model_folder_path("gguf", weigths_gguf_current_path) # gguf dir
class InteractAvatar_SM_Model(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Model",
display_name="InteractAvatar_SM_Model",
category="InteractAvatar",
inputs=[
io.Combo.Input("dit",options= ["none"] +folder_paths.get_filename_list("diffusion_models") ),
io.Combo.Input("gguf",options= ["none"] +folder_paths.get_filename_list("gguf") ),
io.Combo.Input("lora",options= ["none"] + folder_paths.get_filename_list("loras") ),
io.Combo.Input("back_append_frame",options= [1,2]),
io.Boolean.Input("offload",default=True),
],
outputs=[
io.Model.Output(display_name="model"),
],
)
@classmethod
def execute(cls, dit,gguf,lora,back_append_frame,offload) -> io.NodeOutput:
dit_path=folder_paths.get_full_path("diffusion_models", dit) if dit != "none" else None
gguf_path=folder_paths.get_full_path("gguf", gguf) if gguf != "none" else None
lora_path=folder_paths.get_full_path("loras", lora) if lora != "none" else None
assert dit_path != None , "Please select a model"
clear_comfyui_cache()
model = load_model(
checkpoint_path=dit_path if dit_path != None else gguf_path if gguf_path != None else None,
dit_path=dit_path if dit_path != None else gguf_path if gguf_path != None else None,
lora_path=lora_path,
back_append_frame=int(back_append_frame),
offload=offload
)
return io.NodeOutput(model)
class InteractAvatar_SM_Predata(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Predata",
display_name="InteractAvatar_SM_Predata",
category="InteractAvatar",
inputs=[
io.Clip.Input("clip"),
io.Vae.Input("vae"),
io.Image.Input("images"), # image or video
io.Image.Input("pose_images"), # image or video
io.Int.Input("short_side", default=512, min=256, max=nodes.MAX_RESOLUTION,step=32,display_mode=io.NumberDisplay.number),
io.String.Input("prompt",multiline=True, default="两只手打招呼 伸出大拇指点赞"),
io.String.Input("negative_prompt",multiline=True, default="bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"),
io.String.Input("structured_prompt",multiline=True, default=" (Raise both hands slowly) (Wave both hands side to side for greeting),\n (Make a fist) (Extend the thumb upwards)"),
io.Int.Input("num_frames", default=81, min=8, max=10000,step=1,display_mode=io.NumberDisplay.number),
io.Combo.Input("mode",options=['a2mv','ap2v','mv','a2v','p2v'], default="a2mv"),
io.Combo.Input("back_append_frame",options= [1,2]),
io.Boolean.Input("all_text",default=True),
io.String.Input("wav2vec_repo",multiline=False, default=""),
io.String.Input("object_name",multiline=False, default=""),
io.Audio.Input("audio",optional=True),
io.Image.Input("object_images",optional=True),
io.Mask.Input("object_mask",optional=True),
],
outputs=[
io.Conditioning.Output(display_name="data_dict"),
],
)
@classmethod
def execute(cls, clip,vae,images,pose_images,short_side,prompt,negative_prompt,structured_prompt,num_frames,mode,back_append_frame,all_text,wav2vec_repo,object_name,audio=None,object_images=None,object_mask=None) -> io.NodeOutput:
if back_append_frame==2:
prompt=[ii for ii in prompt.splitlines() if ii]
structured_prompt=[ii for ii in structured_prompt.splitlines() if ii]
assert len(prompt)==len(structured_prompt), "Please make sure the number of prompts and structured prompts are the same,人物提示词和动作提示词的行数要一致"
img=tensor2pillist(images)
dw_img=tensor2pillist(pose_images)
audio_path=trans2path(audio)
data_dict = perdata(
clip,
vae,
img,
dw_img,
object_images,
object_mask,
audio_path,
mode,
prompt=prompt,
negative_prompt=negative_prompt,
structured_prompt=structured_prompt,
frame_num=num_frames,
short_side=short_side,
back_append_frame=back_append_frame,
wav2vec_dir=get_wav2vec_repo(wav2vec_repo),
device=device,
all_text=all_text,
)
return io.NodeOutput(data_dict)
class InteractAvatar_SM_Sampler(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Sampler",
display_name="InteractAvatar_SM_Sampler",
category="InteractAvatar",
inputs=[
io.Model.Input("model"),
io.Conditioning.Input("data_dict"),
io.Int.Input("steps", default=20, min=1, max=10000,display_mode=io.NumberDisplay.number),
io.Int.Input("seed", default=0, min=0, max=MAX_SEED),
io.Float.Input("sample_shift", default=5.0, min=0.0, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Float.Input("text_guide_scale", default=5.0, min=0.1, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Float.Input("audio_guide_scale", default=7.5, min=0.1, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Int.Input("bad_thres", default=800, min=0, max=100000,step=1,display_mode=io.NumberDisplay.number),
io.Int.Input("motion_frame",default=25, min=8, max=10000,step=1,display_mode=io.NumberDisplay.number),
io.Boolean.Input("bad_cfg",default=True),
io.Boolean.Input("three_cfg",default=False),
],
outputs=[
io.Latent.Output(display_name="Latent"),
],
)
@classmethod
def execute(cls, model,data_dict,steps,seed,sample_shift,text_guide_scale,audio_guide_scale,bad_thres,motion_frame,bad_cfg,three_cfg) -> io.NodeOutput:
clear_comfyui_cache()
param_dict = dict(
shift=sample_shift,
sampling_steps=steps,
text_guide_scale=text_guide_scale,
audio_guide_scale=audio_guide_scale,
seed=seed,
bad_cfg=bad_cfg,
three_cfg=three_cfg,
bad_thres=bad_thres,
offload_model=True,
motion_frame=motion_frame,
)
data_dict.update(param_dict)
video = generate_video(model,data_dict)
clear_comfyui_cache() # clear cache
if data_dict["back_append_frame"]==2: # don't need decoder image # B H W C
# import time
# pre_fix=time.strftime("%Y%m%d-%H%M%S")
# save_files=os.path.join(folder_paths.get_output_directory(), f"output_{pre_fix}.pt")
# torch.save(video, save_files) #本地解码
video=model.vae.encode(video) # 统一输出接口
else:
video=get_z_scale(video) #torch.Size([1, 16, 21, 90, 136])
output={"samples":video}
return io.NodeOutput(output)
class InteractAvatar_SM_Extension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
InteractAvatar_SM_Model,
InteractAvatar_SM_Predata,
InteractAvatar_SM_Sampler,
]
async def comfy_entrypoint() -> InteractAvatar_SM_Extension: # ComfyUI calls this to load your extension and its nodes.
return InteractAvatar_SM_Extension()
+3
View File
@@ -0,0 +1,3 @@
from .InteractAvatar_node import *
+29
View File
@@ -0,0 +1,29 @@
compute_environment: LOCAL_MACHINE
deepspeed_config:
deepspeed_hostfile: hostfile
deepspeed_multinode_launcher: pdsh
gradient_accumulation_steps: 1
gradient_clipping: 1.5
offload_optimizer_device: none
offload_param_device: none
zero3_init_flag: true
zero_stage: 2
reduce_scatter: false
overlap_comm: true
distributed_type: DEEPSPEED
downcast_bf16: 'no'
dynamo_backend: 'NO'
fsdp_config: {}
machine_rank: 0
main_process_ip: xxxx
main_process_port: 8181
main_training_function: main
megatron_lm_config: {}
mixed_precision: fp16
num_machines: 8
num_processes: 64
rdzv_backend: static
same_network: true
use_cpu: false
# standard
+495
View File
@@ -0,0 +1,495 @@
# !/usr/bin/env python
# -*- coding: UTF-8 -*-
import os
import torch
import gc
from PIL import Image
import numpy as np
import math
import comfy.utils
import cv2
import random
import torchaudio
import folder_paths
from comfy.utils import common_upscale,ProgressBar
from safetensors.torch import load_file
import comfy.model_management as mm
from pathlib import PureWindowsPath
cur_path = os.path.dirname(os.path.abspath(__file__))
def covert_obj_img(images, masks,target_size):
background_color=(255, 255, 255)
if images is not None and masks is not None:
output_list = []
# Convert to numpy for easier processing
images_np = images.cpu().numpy()
# Handle mask shapes
if masks.dim() == 2:
# Single mask HW, expand to BHW
masks_np = masks.unsqueeze(0).cpu().numpy()
elif masks.dim() == 3:
# Batch masks BHW
masks_np = masks.cpu().numpy()
else:
raise ValueError("Mask must be of shape HW or BHW")
assert masks_np.shape[0] == images_np.shape[0] , "Masks and images must had same batch size"
batch_size = images_np.shape[0]
for i in range(batch_size):
# Get image and corresponding mask
img = images_np[i] # HWC
mask = masks_np[i]
# Ensure image is in range [0, 255]
if img.max() <= 1.0:
img = img * 255
img = img.astype(np.uint8)
# Create RGBA image using mask as alpha channel
rgba_img = np.zeros((img.shape[0], img.shape[1], 4), dtype=np.uint8)
rgba_img[:, :, :3] = img
rgba_img[:, :, 3] = ((1 - mask) * 255).astype(np.uint8) #
# Convert to PIL Image
pil_img = Image.fromarray(rgba_img, 'RGBA')
pil_img.save(f"{i}temp.png")
output_list.append(pil_img)
img=output_list[0]
original_width, original_height = img.size
target_width, target_height = target_size
# 1. 计算缩放比例,确保图片能完整放入目标框内
ratio = min(target_width / original_width, target_height / original_height)
# 2. 计算缩放后的新尺寸
new_width = int(original_width * ratio)
new_height = int(original_height * ratio)
# 3. 高质量缩放图片
resized_img = img.resize((new_width, new_height), Image.Resampling.LANCZOS)
# 4. 创建一个新的纯色背景画布
# 注意:颜色需要是RGBA格式,所以白色是 (255, 255, 255, 255)
# 最后一个值255代表完全不透明
final_background_color = background_color + (255,)
background = Image.new('RGBA', target_size, final_background_color)
# 5. 计算粘贴位置,使其居中
paste_x = (target_width - new_width) // 2
paste_y = (target_height - new_height) // 2
# 6. 将缩放后的图片粘贴到背景画布上
# 第三个参数 `resized_img` 作为蒙版,可以正确处理PNG的透明通道
background.paste(resized_img, (paste_x, paste_y), resized_img)
else:
final_background_color = background_color + (255,)
background = Image.new('RGBA', target_size, final_background_color)
background.save("temp.png")
return background.convert('RGB')
def get_wav2vec_repo(repo):
required_files = ["chinese-wav2vec2-base-fairseq-ckpt.pt", "config.json", "model.safetensors", "preprocessor_config.json"]
if not repo:
wav2vec_repo=os.path.join(folder_paths.models_dir, "wav2vec2-base")
if not os.path.exists(wav2vec_repo):
os.makedirs(wav2vec_repo)
download_file(folder_paths.models_dir, required_files)
else:
if not check_files_exist(wav2vec_repo,required_files):
download_file(folder_paths.models_dir, required_files)
else:
wav2vec_repo=PureWindowsPath(repo).as_posix()
return wav2vec_repo
def download_file(local_dir, required_files):
from huggingface_hub import hf_hub_download
for i in required_files:
hf_hub_download(
repo_id="youliang1233214/InteractAvatar",
subfolder="wav2vec2-base",
filename=i,
local_dir = local_dir,
)
def check_files_exist(folder_path, required_files):
for file_name in required_files:
file_path = os.path.join(folder_path, file_name)
if not os.path.exists(file_path):
return False
return True
def clear_comfyui_cache():
cf_models=mm.loaded_models()
try:
for pipe in cf_models:
pipe.unpatch_model(device_to=torch.device("cpu"))
print(f"Unpatching models.{pipe}")
except: pass
mm.soft_empty_cache()
torch.cuda.empty_cache()
max_gpu_memory = torch.cuda.max_memory_allocated()
print(f"After Max GPU memory allocated: {max_gpu_memory / 1000 ** 3:.2f} GB")
def trans2path(audio):
if audio is None:
return None
import io as io_base
audio_file_prefix = ''.join(random.choice("0123456789") for _ in range(6))
audio_file = os.path.join(folder_paths.get_input_directory(), f"audio_{audio_file_prefix}_temp.wav")
buff = io_base.BytesIO()
torchaudio.save(buff, audio["waveform"].squeeze(0), audio["sample_rate"], format="FLAC")
with open(audio_file, 'wb') as f:
f.write(buff.getbuffer())
return audio_file
def encode_image( image, vae):
if image is None:
return None
ref_latents=None
samples = image.movedim(-1, 1)
total = int(1024 * 1024)
scale_by = math.sqrt(total / (samples.shape[3] * samples.shape[2]))
width = round(samples.shape[3] * scale_by)
height = round(samples.shape[2] * scale_by)
s = comfy.utils.common_upscale(samples, width, height, "area", "disabled")
image = s.movedim(1, -1)
if vae is not None:
ref_latents = vae.encode(image[:, :, :, :3])
return ref_latents
def add_mean(latents):
vae_config={"latents_mean": [
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921
],
"latents_std": [
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.916
],}
latents_mean = (torch.tensor(vae_config["latents_mean"]).view(1, 16, 1, 1, 1).to(latents.device, latents.dtype))
latents_std = 1.0 / torch.tensor(vae_config["latents_std"]).view(1, 16, 1, 1, 1).to(latents.device, latents.dtype)
latents = latents / latents_std + latents_mean
image_latent_height, image_latent_width = latents.shape[3:]
image_latents = pack_latents_(
latents, 1, 16, image_latent_height, image_latent_width)
return image_latents
def pack_latents_(latents, batch_size, num_channels_latents, height, width):
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
latents = latents.permute(0, 2, 4, 1, 3, 5)
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
return latents
def load_lora(model, lora_1, lora_2, lora_scale1, lora_scale2):
lora_path_1=folder_paths.get_full_path("loras", lora_1) if lora_1 != "none" else None
lora_path_2=folder_paths.get_full_path("loras", lora_2) if lora_2 != "none" else None
# lora_list=[i for i in [lora_path_1,lora_path_2] if i is not None]
# lora_scales=[lora_scale1,lora_scale2]
all_adapters = model.get_list_adapters()
dit_list=[]
if all_adapters:
dit_list= all_adapters.get('transformer',[])+all_adapters.get('transformer_2',[])
if lora_path_1 is not None:
adapter_name=os.path.splitext(os.path.basename(lora_path_1))[0].replace(".", "_")
dit_list2=all_adapters.get('transformer_2',[])
if dit_list2:
if adapter_name in dit_list: #dit_list
pass
else:
for i in dit_list2:
model.delete_adapters(i)
print(f"去除dit中未加载的lora: {i}")
try:
model.load_lora_weights(lora_path_1, adapter_name=adapter_name,**{"load_into_transformer_2": True})
model.set_adapters([adapter_name], adapter_weights=lora_scale1)
except KeyError as e:
try:
print(f"检测到特殊的 LoRA 格式,尝试手动处理: {lora_path_1}")
state_dict = torch.load(lora_path_1, map_location="cpu",weights_only=False) if not lora_path_1.endswith(".safetensors") else load_file(lora_path_1,)
processed_state_dict = preprocess_lora_state_dict(state_dict)
model.load_lora_weights(processed_state_dict, adapter_name=adapter_name,**{"load_into_transformer_2": True})
model.set_adapters([adapter_name], adapter_weights=lora_scale1)
except:
print(f"加载LoRA权重失败: {e}")
pass
else:
try:
model.load_lora_weights(lora_path_1, adapter_name=adapter_name,**{"load_into_transformer_2": True})
model.set_adapters([adapter_name], adapter_weights=lora_scale1)
except KeyError as e:
try:
print(f"检测到特殊的 LoRA 格式,尝试手动处理: {lora_path_1}")
state_dict = torch.load(lora_path_1, map_location="cpu",weights_only=False) if not lora_path_1.endswith(".safetensors") else load_file(lora_path_1,)
processed_state_dict = preprocess_lora_state_dict(state_dict)
model.load_lora_weights(processed_state_dict, adapter_name=adapter_name,**{"load_into_transformer_2": True})
model.set_adapters([adapter_name], adapter_weights=lora_scale1)
del processed_state_dict
except:
print(f"加载LoRA权重失败: {e}")
pass
if lora_path_2 is not None:
adapter_name=os.path.splitext(os.path.basename(lora_path_2))[0].replace(".", "_")
dit_list=all_adapters.get('transformer',[])
if dit_list:
if adapter_name in dit_list: #dit_list
pass
else:
for i in dit_list:
model.delete_adapters(i)
print(f"去除dit中未加载的lora: {i}")
try:
model.load_lora_weights(lora_path_2, adapter_name=adapter_name,**{"load_into_transformer_2": False})
model.set_adapters([adapter_name], adapter_weights=lora_scale2)
except KeyError as e:
try:
print(f"检测到特殊的 LoRA 格式,尝试手动处理: {lora_path_2}")
state_dict = torch.load(lora_path_2, map_location="cpu",weights_only=False) if not lora_path_2.endswith(".safetensors") else load_file(lora_path_2,)
processed_state_dict = preprocess_lora_state_dict(state_dict)
model.load_lora_weights(processed_state_dict, adapter_name=adapter_name,**{"load_into_transformer_2": False})
model.set_adapters([adapter_name], adapter_weights=lora_scale2)
del processed_state_dict
except:
print(f"加载LoRA权重失败: {e}")
pass
else:
try:
model.load_lora_weights(lora_path_2, adapter_name=adapter_name,**{"load_into_transformer_2": False})
model.set_adapters([adapter_name], adapter_weights=lora_scale2)
except KeyError as e:
try:
print(f"检测到特殊的 LoRA 格式,尝试手动处理: {lora_path_2}")
state_dict = torch.load(lora_path_2, map_location="cpu",weights_only=False) if not lora_path_2.endswith(".safetensors") else load_file(lora_path_2,)
processed_state_dict = preprocess_lora_state_dict(state_dict)
model.load_lora_weights(processed_state_dict, adapter_name=adapter_name,**{"load_into_transformer_2": False})
model.set_adapters([adapter_name], adapter_weights=lora_scale2)
del processed_state_dict
except:
print(f"加载LoRA权重失败: {e}")
pass
return model
def preprocess_lora_state_dict(state_dict):
processed_dict = state_dict.copy()
keys_to_remove = [
'head.head.diff_b',
'head.head.diff_m',
'head.head.diff',
'patch_embedding.diff',
'patch_embedding.diff_b',
'blocks.*.diff_m', # 匹配所有blocks的diff_m
'head.head.lora_down'
'diffusion_model.head.head.diff'
'diffusion_model.head.head.diff_b'
'diffusion_model.head.lora_down'
]
keys_to_delete = []
for key in processed_dict.keys():
if key.endswith('.diff_m'):
keys_to_delete.append(key)
for key in keys_to_delete:
processed_dict.pop(key, None)
print(f"移除键: {key}")
for key in keys_to_remove:
if key in processed_dict:
processed_dict.pop(key, None)
print(f"移除键: {key}")
return processed_dict
def gc_cleanup():
gc.collect()
torch.cuda.empty_cache()
def tensor2cv(tensor_image):
if len(tensor_image.shape)==4:# b hwc to hwc
tensor_image=tensor_image.squeeze(0)
if tensor_image.is_cuda:
tensor_image = tensor_image.cpu()
tensor_image=tensor_image.numpy()
#反归一化
maxValue=tensor_image.max()
tensor_image=tensor_image*255/maxValue
img_cv2=np.uint8(tensor_image)#32 to uint8
img_cv2=cv2.cvtColor(img_cv2,cv2.COLOR_RGB2BGR)
return img_cv2
def phi2narry(img):
img = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
return img
def tensor2image(tensor):
tensor = tensor.cpu()
image_np = tensor.squeeze().mul(255).clamp(0, 255).byte().numpy()
image = Image.fromarray(image_np, mode='RGB')
return image
def tensor2pillist(tensor_in):
d1, _, _, _ = tensor_in.size()
if d1 == 1:
img_list = [tensor2image(tensor_in)]
else:
tensor_list = torch.chunk(tensor_in, chunks=d1)
img_list=[tensor2image(i) for i in tensor_list]
return img_list
def tensor2pillist_upscale(tensor_in,width,height):
d1, _, _, _ = tensor_in.size()
if d1 == 1:
img_list = [nomarl_upscale(tensor_in,width,height)]
else:
tensor_list = torch.chunk(tensor_in, chunks=d1)
img_list=[nomarl_upscale(i,width,height) for i in tensor_list]
return img_list
def tensor2list(tensor_in,width,height):
if tensor_in is None:
return None
d1, _, _, _ = tensor_in.size()
if d1 == 1:
tensor_list = [tensor_upscale(tensor_in,width,height)]
else:
tensor_list_ = torch.chunk(tensor_in, chunks=d1)
tensor_list=[tensor_upscale(i,width,height) for i in tensor_list_]
return tensor_list
def tensor_upscale(tensor, width, height):
samples = tensor.movedim(-1, 1)
samples = common_upscale(samples, width, height, "bilinear", "center")
samples = samples.movedim(1, -1)
return samples
def nomarl_upscale(img, width, height):
samples = img.movedim(-1, 1)
img = common_upscale(samples, width, height, "bilinear", "center")
samples = img.movedim(1, -1)
img = tensor2image(samples)
return img
def cv2tensor(img,bgr2rgb=True):
assert type(img) == np.ndarray, 'the img type is {}, but ndarry expected'.format(type(img))
if bgr2rgb:
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
img = torch.from_numpy(img.transpose((2, 0, 1)))
return img.float().div(255).permute(1, 2, 0).unsqueeze(0)
def images_generator(img_list: list, ):
# get img size
sizes = {}
for image_ in img_list:
if isinstance(image_, Image.Image):
count = sizes.get(image_.size, 0)
sizes[image_.size] = count + 1
elif isinstance(image_, np.ndarray):
count = sizes.get(image_.shape[:2][::-1], 0)
sizes[image_.shape[:2][::-1]] = count + 1
else:
raise "unsupport image list,must be pil or cv2!!!"
size = max(sizes.items(), key=lambda x: x[1])[0]
yield size[0], size[1]
# any to tensor
def load_image(img_in):
if isinstance(img_in, Image.Image):
img_in = img_in.convert("RGB")
i = np.array(img_in, dtype=np.float32)
i = torch.from_numpy(i).div_(255)
if i.shape[0] != size[1] or i.shape[1] != size[0]:
i = torch.from_numpy(i).movedim(-1, 0).unsqueeze(0)
i = common_upscale(i, size[0], size[1], "lanczos", "center")
i = i.squeeze(0).movedim(0, -1).numpy()
return i
elif isinstance(img_in, np.ndarray):
i = cv2.cvtColor(img_in, cv2.COLOR_BGR2RGB).astype(np.float32)
i = torch.from_numpy(i).div_(255)
print(i.shape)
return i
else:
raise "unsupport image list,must be pil,cv2 or tensor!!!"
total_images = len(img_list)
processed_images = 0
pbar = ProgressBar(total_images)
images = map(load_image, img_list)
try:
prev_image = next(images)
while True:
next_image = next(images)
yield prev_image
processed_images += 1
pbar.update_absolute(processed_images, total_images)
prev_image = next_image
except StopIteration:
pass
if prev_image is not None:
yield prev_image
def load_images_list(img_list: list, ):
gen = images_generator(img_list)
(width, height) = next(gen)
images = torch.from_numpy(np.fromiter(gen, np.dtype((np.float32, (height, width, 3)))))
if len(images) == 0:
raise FileNotFoundError(f"No images could be loaded .")
return images
def get_video_files(directory, extensions=None):
if extensions is None:
extensions = ['webm', 'mp4', 'mkv', 'gif', 'mov']
extensions = [ext.lower() for ext in extensions]
video_files = []
for root, dirs, files in os.walk(directory):
for file in files:
_, ext = os.path.splitext(file)
ext = ext.lower()[1:]
if ext in extensions:
full_path = os.path.join(root, file)
video_files.append(full_path)
return video_files
+128
View File
@@ -0,0 +1,128 @@
import torch
import bitsandbytes
import bitsandbytes.functional as F
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# AdamW8bitKahan: AdamW8bit with Kahan summation
class AdamW8bitKahan(bitsandbytes.optim.AdamW8bit):
def __init__(self, *args, stabilize=True, **kwargs):
super().__init__(*args, **kwargs)
self.stabilize = stabilize
@torch.no_grad()
def init_state(self, group, p, gindex, pindex):
super().init_state(group, p, gindex, pindex)
self.state[p]['shift'] = self.get_state_buffer(p, dtype=p.dtype)
@torch.no_grad()
def update_step(self, group, p, gindex, pindex):
# avoid update error from non-contiguous memory layout
p.data = p.data.contiguous()
p.grad = p.grad.contiguous()
state = self.state[p]
grad = p.grad
# Kahan summation
config = self.get_config(gindex, pindex, group)
state["step"] += 1
step = state["step"]
# percentile clipping
if config["percentile_clipping"] < 100:
current_gnorm, clip_value, gnorm_scale = F.percentile_clipping(
grad,
state["gnorm_vec"],
step,
config["percentile_clipping"],
)
else:
gnorm_scale = 1.0
shift = state['shift']
# StableAdamW
if self.stabilize:
exp_avg_sq = state['state2']
eps_sq = torch.tensor(config['eps']**2, dtype=exp_avg_sq.dtype, device=exp_avg_sq.device)
rms = grad.pow(2).div_(exp_avg_sq.maximum(eps_sq)).mean().sqrt()
lr = config['lr'] / max(1, rms.item())
else:
lr = config['lr']
# Kahan summation
if state["state1"].dtype == torch.float:
F.optimizer_update_32bit(
self.optimizer_name,
grad,
shift,
state["state1"],
config["betas"][0],
config["eps"],
step,
lr,
state["state2"],
config["betas"][1],
config["betas"][2] if len(config["betas"]) >= 3 else 0.0,
config["alpha"],
config["weight_decay"],
gnorm_scale,
state["unorm_vec"] if config["max_unorm"] > 0.0 else None,
max_unorm=config["max_unorm"],
skip_zeros=config["skip_zeros"],
)
# 8-bit update
elif state["state1"].dtype == torch.uint8 and not config["block_wise"]:
F.optimizer_update_8bit(
self.optimizer_name,
grad,
shift,
state["state1"],
state["state2"],
config["betas"][0],
config["betas"][1],
config["eps"],
step,
lr,
state["qmap1"],
state["qmap2"],
state["max1"],
state["max2"],
state["new_max1"],
state["new_max2"],
config["weight_decay"],
gnorm_scale=gnorm_scale,
unorm_vec=state["unorm_vec"] if config["max_unorm"] > 0.0 else None,
max_unorm=config["max_unorm"],
)
# swap maxes
state["max1"], state["new_max1"] = state["new_max1"], state["max1"]
state["max2"], state["new_max2"] = state["new_max2"], state["max2"]
elif state["state1"].dtype == torch.uint8 and config["block_wise"]:
F.optimizer_update_8bit_blockwise(
self.optimizer_name,
grad,
shift,
state["state1"],
state["state2"],
config["betas"][0],
config["betas"][1],
config["betas"][2] if len(config["betas"]) >= 3 else 0.0,
config["alpha"],
config["eps"],
step,
lr,
state["qmap1"],
state["qmap2"],
state["absmax1"],
state["absmax2"],
config["weight_decay"],
gnorm_scale=gnorm_scale,
skip_zeros=config["skip_zeros"],
)
# update p
buffer = p.clone()
p.add_(shift)
# Kahan summation
shift.add_(buffer.sub_(p))
+28
View File
@@ -0,0 +1,28 @@
import torch
# Simple wrapper for use with gradient release. Grad hooks do the optimizer steps, so this no-ops
# the step() and zero_grad() methods. It also handles state_dict.
class GradientReleaseOptimizerWrapper(torch.optim.Optimizer):
def __init__(self, optimizers):
self.optimizers = optimizers
@property
def param_groups(self):
ret = []
for opt in self.optimizers:
ret.extend(opt.param_groups)
return ret
def state_dict(self):
return {i: opt.state_dict() for i, opt in enumerate(self.optimizers)}
# load_state_dict: load state dict
def load_state_dict(self, state_dict):
for i, sd in state_dict.items():
self.optimizers[i].load_state_dict(sd)
def step(self):
pass
def zero_grad(self, set_to_none=True):
pass
+16
View File
@@ -0,0 +1,16 @@
[project]
name = "interactavatar"
description = "Making Avatars InteractTowards Text-Driven Human-Object Interaction for Controllable Talking Avatars"
version = "1.0.0"
license = {file = "LICENSE"}
dependencies = ["toml", "transformers", "diffusers", "datasets", "pillow", "sentencepiece", "protobuf", "peft", "torch-optimi", "tensorboard", "tqdm", "safetensors", "bitsandbytes", "imageio[ffmpeg]", "av", "einops", "accelerate", "loguru", "easydict", "ftfy", "decord", "pyloudnorm", "#deepspeed"]
[project.urls]
Repository = "https://github.com/smthemex/ComfyUI_InteractAvatar"
# Used by Comfy Registry https://registry.comfy.org
[tool.comfy]
PublisherId = "smthemex"
DisplayName = "ComfyUI_InteractAvatar"
Icon = ""
includes = []
+24
View File
@@ -0,0 +1,24 @@
toml
transformers
diffusers
datasets
pillow
sentencepiece
protobuf
peft
torch-optimi
tensorboard
tqdm
safetensors
bitsandbytes
imageio[ffmpeg]
av
einops
accelerate
loguru
easydict
ftfy
decord
pyloudnorm
#deepspeed
+63
View File
@@ -0,0 +1,63 @@
#!/bin/bash
# ======== 基本配置 (保持不变) ========
GPUS_NUM=1
inference_steps=40
model_path="./model_pretrain/Wan2.2-TI2V-5B"
exp_name="interact-demo"
caption="struct"
GPU_ID_LIST=(0 1 2)
GPU_ID=0
TOTAL_GPUS=8
wav2vec_path="./ckpt/wav2vec2-base/"
# for a2mv
mode="a2mv"
transformer_path="./ckpt/interact-avatar/"
test_data_path="./InterDemo/TIA2MV/demo_tia2mv.json"
back_append_frame=1
frame_num=101
# for pose dirven with audio
# mode="ap2v"
# transformer_path="./ckpt/interact-avatar/"
# test_data_path="./InterDemo/TIAP2V/demo_ap2v.json"
# back_append_frame=1
# frame_num=133
# for song with action
# mode="a2mv"
# transformer_path="./ckpt/interact-avatar-long/"
# test_data_path="./InterDemo/ActionAndSong/demo_song_action.json"
# back_append_frame=2
# frame_num=1000
base_save_path="./output/"$mode"_"$back_append_frame
for GPU_ID in "${GPU_ID_LIST[@]}"; do
echo "======> Checkpoint on GPU $GPU_ID with mode $mode cfg <======"
save_path=${base_save_path}
CUDA_VISIBLE_DEVICES=$GPU_ID python test_wanx_tia2mv_obj_back.py \
--task "ti2v-5B" \
--ckpt_dir $model_path \
--ulysses_size $GPUS_NUM \
--frame_num $frame_num \
--sample_shift 5.0 \
--text_guide_scale 5.0 \
--audio_guide_scale 7.5 \
--transformer_dir $transformer_path \
--sample_steps $inference_steps \
--save_path $save_path \
--test_data_path $test_data_path\
--base_seed 2025 \
--mode $mode \
--start $(($GPU_ID)) \
--end $(($GPU_ID + 1)) \
--caption $caption \
--short_side 704 \
--bad_cfg \
--bad_thres 800 \
--back_append_frame $back_append_frame \
--wav2vec_dir $wav2vec_path &
done
echo "all task started..."
wait
echo "all task done"
+539
View File
@@ -0,0 +1,539 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import logging
import sys
import warnings
import pyloudnorm as pyln
import librosa
warnings.filterwarnings('ignore')
import numpy as np
import torch
from PIL import Image
import librosa
import re
import math
#import wan
from .wan.tia2mv_obj_back_id_prefix import WanTIA2MVRefBackIDPrefix,distribute_prompts_with_gaps
from einops import rearrange
from .utils.audio_analysis.wav2vec2 import Wav2Vec2Model
from transformers import Wav2Vec2FeatureExtractor
from .wan.configs import WAN_CONFIGS
from .utils.img_utils import process_images_final,resize_images,resize_short_side
from .model_loader_utils import covert_obj_img,clear_comfyui_cache
mean = torch.tensor(
[
-0.2289,
-0.0052,
-0.1323,
-0.2339,
-0.2799,
0.0174,
0.1838,
0.1557,
-0.1382,
0.0542,
0.2813,
0.0891,
0.1570,
-0.0098,
0.0375,
-0.1825,
-0.2246,
-0.1207,
-0.0698,
0.5109,
0.2665,
-0.2108,
-0.2158,
0.2502,
-0.2055,
-0.0322,
0.1109,
0.1567,
-0.0729,
0.0899,
-0.2799,
-0.1230,
-0.0313,
-0.1649,
0.0117,
0.0723,
-0.2839,
-0.2083,
-0.0520,
0.3748,
0.0152,
0.1957,
0.1433,
-0.2944,
0.3573,
-0.0548,
-0.1681,
-0.0667,
],
device=torch.device('cuda'),
)
std = torch.tensor(
[
0.4765,
1.0364,
0.4514,
1.1677,
0.5313,
0.4990,
0.4818,
0.5013,
0.8158,
1.0344,
0.5894,
1.0901,
0.6885,
0.6165,
0.8454,
0.4978,
0.5759,
0.3523,
0.7135,
0.6804,
0.5833,
1.4146,
0.8986,
0.5659,
0.7069,
0.5338,
0.4889,
0.4917,
0.4069,
0.4999,
0.6866,
0.4093,
0.5709,
0.6065,
0.6415,
0.4944,
0.5726,
1.2042,
0.5458,
1.6887,
0.3971,
1.0600,
0.3943,
0.5537,
0.5444,
0.4089,
0.7468,
0.7744,
],
device=torch.device('cuda'),
)
wanvae_scale = [mean, 1.0 / std]
def get_mu_scale(mu):
mu = (mu - wanvae_scale[0].view(1, 48, 1, 1, 1).to(mu.device,mu.dtype)) * wanvae_scale[1].view(
1, 48, 1, 1, 1).to(mu.device,mu.dtype)
return mu
def get_z_scale(z):
z = z / wanvae_scale[1].view(1, 48, 1, 1, 1).to(z.device,z.dtype) + wanvae_scale[0].view(
1, 48, 1, 1, 1).to(z.device,z.dtype)
return z
def phi2narry(img):
img = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
return img
def loudness_norm(audio_array, sr=16000, lufs=-23):
meter = pyln.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > 100:
return audio_array
normalized_audio = pyln.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def align_floor_to(value, alignment):
return int(math.floor(value / alignment) * alignment)
def align_ceil_to(value, alignment):
return int(math.ceil(value / alignment) * alignment)
def custom_init(device, wav2vec):
audio_encoder = Wav2Vec2Model.from_pretrained(wav2vec, local_files_only=True).to(device)
audio_encoder.feature_extractor._freeze_parameters()
wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(wav2vec, local_files_only=True)
return wav2vec_feature_extractor, audio_encoder
def get_embedding(speech_array, wav2vec_feature_extractor, audio_encoder, sr=16000, device='cpu'):
audio_duration = len(speech_array) / sr
video_length = audio_duration * 25 # Assume the video fps is 25
# wav2vec_feature_extractor
audio_feature = np.squeeze(
wav2vec_feature_extractor(speech_array, sampling_rate=sr).input_values
)
audio_feature = torch.from_numpy(audio_feature).float().to(device=device)
audio_feature = audio_feature.unsqueeze(0)
# audio encoder
with torch.no_grad():
embeddings = audio_encoder(audio_feature, seq_len=int(video_length), output_hidden_states=True)
if len(embeddings) == 0:
print("Fail to extract audio embedding")
return None
audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0)
audio_emb = rearrange(audio_emb, "b s d -> s b d")
return audio_emb
def _init_logging(rank):
# logging
if rank == 0:
# set format
logging.basicConfig(
level=logging.INFO,
format="[%(asctime)s] %(levelname)s: %(message)s",
handlers=[logging.StreamHandler(stream=sys.stdout)])
else:
logging.basicConfig(level=logging.ERROR)
def load_model(checkpoint_path,dit_path,lora_path,short_side=512,back_append_frame=1,offload=False):
logging.info("Creating WanI2V pipeline.")
cfg = WAN_CONFIGS['ti2v-5B']
wan_a2v = WanTIA2MVRefBackIDPrefix(
config=cfg,
checkpoint_dir=checkpoint_path,
transformer_dir=dit_path,
lora_path=lora_path,
device_id=0,
rank=0,
t5_fsdp=False,
dit_fsdp=False,
use_usp=False,#(args.ulysses_size > 1 or args.ring_size > 1),
t5_cpu=True,
load_from_merged_model=None,
short_side=short_side,
back_append_frame=back_append_frame,
offload=offload,
)
return wan_a2v
def generate_video(wan_a2v,gen_kwargs):
logging.info("Generating video ...")
if gen_kwargs["back_append_frame"] == 1:
video, _ = wan_a2v.generate(
None, None, None, None, None, **gen_kwargs
)
else:
wan_a2v.vae=gen_kwargs["vae"]
video, _ = wan_a2v.generate_long(
None, None, None, None, None, **gen_kwargs
)
return video
def perdata( clip,vae,images,dw_iamges,object_images,object_mask,audio_path,mode,prompt,negative_prompt,structured_prompt,frame_num,short_side,back_append_frame,
wav2vec_dir,device,all_text=False,clips=None,max_frames_num=1000,
):
img=images[0]
dw_img=dw_iamges[0]
if len(dw_iamges) >1:
dw_seqs = dw_iamges
if clips is not None:
#clips = batch['clips']
dw_seqs = dw_iamges[clips[0]:clips[1]]
else:
dwpose_frame_num = frame_num
dw_seqs = None
# pre size
w, h = img.size
img = resize_short_side(img, short_side)
small_img = resize_short_side(img, 256)
dw_img = resize_short_side(dw_img, 256)
if short_side == 512:
_, small_img = resize_images(img, small_img)
img, dw_img = resize_images(img, dw_img)
else:
_, small_img = process_images_final(img, small_img)
img, dw_img = process_images_final(img, dw_img)
dw_w, dw_h = dw_img.size
if small_img.size != dw_img.size:
small_img = img.resize(dw_img.size, Image.LANCZOS)
if dw_seqs is not None:
dw_seqs = [dw.resize(dw_img.size, Image.LANCZOS) for dw in dw_seqs]
dwpose_len = len(dw_seqs)
dwpose_frame_num = (dwpose_len - 1) // 4 * 4 + 1
# pre audio
if audio_path is not None:
audio_input, sampling_rate = librosa.load(audio_path, sr=16000)
audio_frames_clip = loudness_norm(audio_input, sampling_rate)
audio_frame_len = int(len(audio_frames_clip) / sampling_rate * 25)
audio_frame_num = (audio_frame_len - 1) // 4 * 4 + 1
else:
audio_frames_clip = np.zeros((int(dwpose_frame_num / 25 * 16000)))
audio_frame_num = dwpose_frame_num
sampling_rate = 16000
if dw_seqs is not None:
dw_seqs = dw_seqs[:frame_num]
if audio_frame_num > dwpose_frame_num:
padding_dwpose = dw_seqs[-1]
dw_seqs = dw_seqs + [padding_dwpose] * (audio_frame_num - dwpose_frame_num)
dwpose_frame_num = audio_frame_num
if audio_frame_num < dwpose_frame_num:
padding_audio = np.zeros((int(dwpose_frame_num / 25 * 16000)))
audio_frames_clip = np.concatenate([audio_frames_clip, padding_audio], axis=0)
audio_frame_num = dwpose_frame_num
if dw_seqs is None and mode == 'ap2v':
dw_seqs = [dw_img] * audio_frame_num
frame_num = min(frame_num,min(audio_frame_num,dwpose_frame_num))
audio_frames_clip = audio_frames_clip[:int(frame_num * sampling_rate / 25)]
wav2vec_feature_extractor, audio_encoder= custom_init(device, wav2vec_dir)
zero_audio_embedding = get_embedding(np.zeros_like(audio_frames_clip), wav2vec_feature_extractor, audio_encoder, device=device)
audio_embedding = get_embedding(audio_frames_clip, wav2vec_feature_extractor, audio_encoder, device=device)
# obj_img_path = obj_img
obj_img = covert_obj_img(object_images,object_mask,target_size=img.size) #RGBA
caption_prompt_list = []
clear_state = " clear hands and face, objects are clear and stable. human movements are slow and steady, with a strong sense of reality."
if back_append_frame == 1:
segments = re.findall(r'\(.*?\)', structured_prompt)
if all_text:
first_three_segments = segments
else:
first_three_segments = segments[:3]
print(first_three_segments)
result_string = "".join(first_three_segments)
caption_prompt = prompt + clear_state + result_string.lower().replace(')(',') (')
else:
for i in range(len(structured_prompt)):
segments = re.findall(r'\(.*?\)', structured_prompt[i])
if all_text:
first_three_segments = segments
else:
first_three_segments = segments[:3]
print(first_three_segments)
result_string = "".join(first_three_segments)
caption_prompt = prompt[i] + clear_state + result_string.lower().replace(')(',') (')
caption_prompt_list.append(caption_prompt)
if mode in ['a2v','a2mv','mv','i2v']:
dw_seqs = None
vae_stride=WAN_CONFIGS['ti2v-5B'].vae_stride
patch_size=WAN_CONFIGS['ti2v-5B'].patch_size
sp_size=1
if back_append_frame==1:
cond_image=phi2narry(img) #BHWC
obj_image=phi2narry(obj_img)
pose_ref_img=phi2narry(dw_img)
small_img=phi2narry(small_img)
#print(cond_image.shape, obj_image.shape, pose_ref_img.shape, small_img.shape) #torch.Size([1, 768, 512, 3]) torch.Size([1, 768, 512, 3]) torch.Size([1, 384, 256, 3]) torch.Size([1, 384, 256, 3])
# 2
if dw_seqs is not None:
if isinstance(dw_seqs, list): # If input is a list of PIL Images
processed_poses = [phi2narry(p) for p in dw_seqs]
cond_pose_sequence = torch.stack(processed_poses).to(device)
else: # If input is already a tensor
cond_pose_sequence = dw_seqs.to(device)
dwpose_len = cond_pose_sequence.shape[0]
dwpose_len = (dwpose_len - 1) // vae_stride[0] * vae_stride[0] + 1
frame_num = min(frame_num, dwpose_len)
cond_pose_sequence = cond_pose_sequence[:frame_num]
else:
cond_pose_sequence = torch.ones((frame_num, pose_ref_img.shape[1], pose_ref_img.shape[2], pose_ref_img.shape[3]), device=pose_ref_img.device, dtype=pose_ref_img.dtype) * 0.5
if mode in ['a2v','a2mv','mv','i2v']:
if dw_seqs is not None:
cond_pose_sequence = torch.ones((frame_num, cond_pose_sequence.shape[1], cond_pose_sequence.shape[2], cond_pose_sequence.shape[-1]), device=cond_pose_sequence.device, dtype=cond_pose_sequence.dtype) * 0.5
# 3
if back_append_frame==1:
tokens_p = clip.tokenize(caption_prompt)
context = clip.encode_from_tokens_scheduled(tokens_p)[0][0].to(device,torch.bfloat16) #use bf16 #torch.Size([1, 512, 4096])
tokens_n = clip.tokenize(negative_prompt)
context_null = clip.encode_from_tokens_scheduled(tokens_n)[0][0].to(device,torch.bfloat16)
else:
pass
clear_comfyui_cache()
# 4
h, w = cond_image.shape[1], cond_image.shape[2] #cf BHWC
lat_h, lat_w = h // vae_stride[1], w // vae_stride[2]
max_seq_len = ((frame_num - 1) // vae_stride[0] + 1) * lat_h * lat_w // (
patch_size[1] * patch_size[2])
max_seq_len = int(math.ceil(max_seq_len / sp_size)) * sp_size
with torch.no_grad():
cond_image =vae.encode(cond_image[:, :, :, :3])
cond_image = get_mu_scale(cond_image)[0].unsqueeze(0) #torch.Size([1, 48, 1, 48, 32])
obj_image=vae.encode(obj_image.repeat( 4, 1,1,1)[:, :, :, :3])
obj_image = get_mu_scale(obj_image)[0].unsqueeze(0)
small_img=vae.encode(small_img.repeat( 4, 1, 1,1)[:, :, :, :3])
small_img = get_mu_scale(small_img)[0].unsqueeze(0)
# 5
cond_dw_img = pose_ref_img # Shape: B, C, 1, H, W. (1, C, N, H, W)
motion_h, motion_w = cond_dw_img.shape[1], cond_dw_img.shape[2] #cf BHWC
motion_lat_h, motion_lat_w = motion_h // vae_stride[1], motion_w // vae_stride[2]
motion_max_seq_len = ((frame_num - 1) // vae_stride[0] + 1) * motion_lat_h * motion_lat_w // (
patch_size[1] * patch_size[2])
motion_max_seq_len = int(math.ceil(motion_max_seq_len / sp_size)) * sp_size
with torch.no_grad():
cond_dw_img=vae.encode(cond_dw_img[:, :, :, :3])
cond_dw_img = get_mu_scale(cond_dw_img)[0].unsqueeze(0)
cond_pose_sequence=vae.encode(cond_pose_sequence[:, :, :, :3])
cond_pose_sequence = get_mu_scale(cond_pose_sequence)[0].unsqueeze(0)
data_dict=dict(
context_null=context_null,
context=context,
lat_h=lat_h,
lat_w=lat_w,
max_seq_len=max_seq_len,
motion_lat_h=motion_lat_h,
motion_lat_w=motion_lat_w,
motion_max_seq_len=motion_max_seq_len,
cond_image=cond_image,
obj_image=obj_image,
small_image=small_img,
cond_dw_img=cond_dw_img,
cond_pose_sequence=cond_pose_sequence,
audio_embedding=audio_embedding,
zero_audio_embedding=zero_audio_embedding,
back_append_frame=back_append_frame,
mode=mode,
frame_num=frame_num,
max_frames_num=frame_num,
)
else:
curr_cond_image=phi2narry(img) #BHWC
curr_pose_ref_img=phi2narry(dw_img) #BCHW
# 固定的尾帧条件 (Fixed End Condition)
obj_image=phi2narry(obj_img)
small_img_tensor=phi2narry(small_img)
id_img=phi2narry(img)#BHWC use img as id_img
id_small_img=phi2narry(small_img) #BHWC
# Pose Sequence 处理
if dw_seqs is not None:
if isinstance(dw_seqs, list):
processed_poses = [phi2narry(p) for p in dw_seqs]
cond_pose_sequence_full = torch.stack(processed_poses).to(device)
else:
cond_pose_sequence_full = dw_seqs.to(device)
else:
cond_pose_sequence_full = torch.ones((frame_num+ 2000, curr_pose_ref_img.shape[1], curr_pose_ref_img.shape[2], curr_pose_ref_img.shape[3]), device=curr_pose_ref_img.device, dtype=curr_pose_ref_img.dtype) * 0.5
if mode in ['a2v', 'a2mv', 'mv', 'i2v']:
if dw_seqs is not None:
cond_pose_sequence_full = torch.ones((cond_pose_sequence_full.shape[0], cond_pose_sequence_full.shape[1], cond_pose_sequence_full.shape[2], cond_pose_sequence_full.shape[3]), device=cond_pose_sequence_full.device, dtype=cond_pose_sequence_full.dtype) * 0.5
h, w = curr_cond_image.shape[1], curr_cond_image.shape[2] #cf BHWC
lat_h, lat_w = h // vae_stride[1], w // vae_stride[2]
motion_h, motion_w = curr_pose_ref_img.shape[1], curr_pose_ref_img.shape[2]
motion_lat_h, motion_lat_w = motion_h // vae_stride[1], motion_w // vae_stride[2]
with torch.no_grad():
encoded_obj_image=vae.encode(obj_image.repeat(4, 1, 1,1)[:, :, :, :3])
encoded_obj_image = get_mu_scale(encoded_obj_image)[0].unsqueeze(0)
encoded_small_img=vae.encode(small_img_tensor.repeat(4, 1, 1,1)[:, :, :, :3])
encoded_small_img = get_mu_scale(encoded_small_img)[0].unsqueeze(0) # torch.Size([1, 48, 1, 24, 16]) encoded_obj_image.shape: torch.Size([1, 48, 1, 48, 32])
encoded_id_img=vae.encode(id_img[:, :, :, :3])
encoded_id_img = get_mu_scale(encoded_id_img)[0].unsqueeze(0)
encoded_id_small_img=vae.encode(id_small_img[:, :, :, :3])
encoded_id_small_img = get_mu_scale(encoded_id_small_img)[0].unsqueeze(0) #torch.Size([1, 48, 1, 48, 32]) encoded_id_small_img.shape: torch.Size([1, 48, 1, 24, 16])
# 循环参数设定
seg_frames = 101
target_latent_len = (seg_frames - 1) // vae_stride[0] + 1 # 26 latents
total_model_seq_len = target_latent_len + 1 + 1 # 26 + 1(ID) + 1(Obj) = 28
# --- 设置 Loop Logic 参数 ---
prefix_latents_num = 5 # 后续段复用的 Latent 数量
vae_temp_stride = 4 # VAE 时间维度的下采样率
effective_new_latents = target_latent_len - prefix_latents_num
stride_frames = effective_new_latents * vae_temp_stride # 84帧
total_video_len = min((audio_embedding.shape[0]-1)//4*4+1, max_frames_num//4*4+1)
if total_video_len <= seg_frames:
num_loops = 1
else:
num_loops = 1 + int(math.ceil((total_video_len - seg_frames) / stride_frames))
print(f"Total Video Length: {total_video_len}, Num Loops: {num_loops}, Stride: {stride_frames}")
if num_loops == 0: num_loops = 1
prompt_dict=caption_prompt_list
num_action = len(prompt_dict) #TODO:
#print('before prompt_dict:', prompt_dict)
prompt_dict = distribute_prompts_with_gaps(prompt_dict, num_loops)
print('after prompt_dict:', prompt_dict)
context_loop = []
context_null_loop = []
for loop_idx in range(num_loops):
if loop_idx < num_loops:
input_prompt = prompt_dict[loop_idx]
else:
input_prompt = prompt_dict[-1]
# input_prompt = short_prompt + " ".join(current_chunk) + clear_state
# 文本编码
tokens_p = clip.tokenize(input_prompt)
context = clip.encode_from_tokens_scheduled(tokens_p)[0][0].to(device,torch.bfloat16)
tokens_n = clip.tokenize(negative_prompt)
context_null = clip.encode_from_tokens_scheduled(tokens_n)[0][0].to(device,torch.bfloat16)
context_loop.append(context)
context_null_loop.append(context_null)
data_dict=dict(
context_loop=context_loop,
context_null_loop=context_null_loop,
lat_h=lat_h,
lat_w=lat_w,
motion_lat_h=motion_lat_h,
motion_lat_w=motion_lat_w,
encoded_small_img=encoded_small_img,
encoded_id_img=encoded_id_img,
encoded_id_small_img=encoded_id_small_img,
encoded_obj_image=encoded_obj_image,
curr_cond_image=curr_cond_image,
curr_pose_ref_img=curr_pose_ref_img,
cond_pose_sequence_full=cond_pose_sequence_full,
audio_embedding=audio_embedding,
zero_audio_embedding=zero_audio_embedding,
back_append_frame=back_append_frame,
mode=mode,
frame_num=frame_num,
vae=vae,
max_frames_num=frame_num,
)
return data_dict
View File
+79
View File
@@ -0,0 +1,79 @@
import json
import os
from pathlib import Path
meta_file = ['/apdcephfs_jn2/share_302243908/terohu/data/human-avideo/openhumanvid0812.list',
'/apdcephfs_jn2/share_302243908/terohu/data/human-avideo/talking_head0808.list',
'/apdcephfs_jn2/share_302243908/terohu/data/human-avideo/openhumanvid_140w_caped_0827.list']
def process_list_files():
"""
读取每个list文件中的JSON数据,根据lang字段分类写入新文件
"""
for file_path in meta_file:
if not os.path.exists(file_path):
print(f"文件不存在: {file_path}")
continue
# 获取文件名(不含扩展名)
file_name = Path(file_path).stem
file_dir = Path(file_path).parent
# 创建输出文件路径
en_output = file_dir / f"{file_name}_en.list"
zh_output = file_dir / f"{file_name}_zh.list"
# 初始化计数器
en_count = 0
zh_count = 0
total_count = 0
print(f"正在处理文件: {file_path}")
try:
with open(file_path, 'r', encoding='utf-8') as f:
with open(en_output, 'w', encoding='utf-8') as en_file, \
open(zh_output, 'w', encoding='utf-8') as zh_file:
for line_num, line in enumerate(f, 1):
if line_num % 10000 == 0:
print(f"已处理 {line_num} 行...")
line = line.strip()
if not line:
continue
try:
# 解析JSON
data=json.load(open(line, "r"))
total_count += 1
# 检查lang字段
if 'language' in data:
lang = data['language'].lower()
if lang == 'en':
en_file.write(line + '\n')
en_count += 1
elif lang == 'zh':
zh_file.write(line + '\n')
zh_count += 1
else:
print(f"未知语言类型 '{lang}' 在第 {line_num} 行: {file_path}")
else:
print(f"缺少 'language' 字段在第 {line_num} 行: {file_path}")
except json.JSONDecodeError as e:
print(f"JSON解析错误在第 {line_num} 行: {file_path}, 错误: {e}")
continue
except Exception as e:
print(f"处理文件时出错 {file_path}: {e}")
continue
print(f"完成处理 {file_path}:")
print(f" 总计: {total_count} 条记录")
print(f" 英文: {en_count} 条记录 -> {en_output}")
print(f" 中文: {zh_count} 条记录 -> {zh_output}")
print("-" * 50)
if __name__ == "__main__":
process_list_files()
+18
View File
@@ -0,0 +1,18 @@
import torch
import torch.nn.functional as F
# get_mask_from_lengths
def get_mask_from_lengths(lengths, max_len=None):
lengths = lengths.to(torch.long)
if max_len is None:
max_len = torch.max(lengths).item()
ids = torch.arange(0, max_len).unsqueeze(0).expand(lengths.shape[0], -1).to(lengths.device)
mask = ids < lengths.unsqueeze(1).expand(-1, max_len)
return mask
# linear_interpolation
def linear_interpolation(features, seq_len):
features = features.transpose(1, 2)
output_features = F.interpolate(features, size=seq_len, align_corners=True, mode='linear')
return output_features.transpose(1, 2)
+80
View File
@@ -0,0 +1,80 @@
# the implementation of Wav2Vec2Model is borrowed from
# https://github.com/huggingface/transformers/blob/HEAD/src/transformers/models/wav2vec2/modeling_wav2vec2.py
# initialize our encoder with the pre-trained wav2vec 2.0 weights.
from transformers import Wav2Vec2Config, Wav2Vec2Model
from transformers.modeling_outputs import BaseModelOutput
from .torch_utils import linear_interpolation
# the implementation of Wav2Vec2Model is borrowed from
# https://github.com/huggingface/transformers/blob/HEAD/src/transformers/models/wav2vec2/modeling_wav2vec2.py
# initialize our encoder with the pre-trained wav2vec 2.0 weights.
class Wav2Vec2Model(Wav2Vec2Model):
def __init__(self, config: Wav2Vec2Config):
super().__init__(config)
def forward(
self,
input_values,
seq_len,
attention_mask=None,
mask_time_indices=None,
output_attentions=None,
output_hidden_states=None,
return_dict=None,
):
self.config.output_attentions = False# use sdpa need set it to False
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
extract_features = self.feature_extractor(input_values)
extract_features = extract_features.transpose(1, 2)
extract_features = linear_interpolation(extract_features, seq_len=seq_len)
if attention_mask is not None:
# compute reduced attention_mask corresponding to feature vectors
attention_mask = self._get_feature_vector_attention_mask(
extract_features.shape[1], attention_mask, add_adapter=False
)
hidden_states, extract_features = self.feature_projection(extract_features)
hidden_states = self._mask_hidden_states(
hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask
)
encoder_outputs = self.encoder(
hidden_states,
attention_mask=attention_mask,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
hidden_states = encoder_outputs[0]
if self.adapter is not None:
hidden_states = self.adapter(hidden_states)
if not return_dict:
return (hidden_states, ) + encoder_outputs[1:]
return BaseModelOutput(
last_hidden_state=hidden_states,
hidden_states=encoder_outputs.hidden_states,
attentions=encoder_outputs.attentions,
)
# extract features from wav2vec 2.0 encoder
def feature_extract(
self,
input_values,
seq_len,
):
extract_features = self.feature_extractor(input_values)
extract_features = extract_features.transpose(1, 2)
extract_features = linear_interpolation(extract_features, seq_len=seq_len)
return extract_features
+56
View File
@@ -0,0 +1,56 @@
from contextlib import contextmanager
import gc
import time
import torch
import deepspeed.comm.comm as dist
# import imageio
from safetensors import safe_open
DTYPE_MAP = {'float32': torch.float32,
'float16': torch.float16,
'bfloat16': torch.bfloat16,
'float8': torch.float8_e4m3fn
}
# VIDEO_EXTENSIONS = set(x.extension for x in imageio.config.video_extensions)
AUTOCAST_DTYPE = None
def get_rank():
return dist.get_rank()
# is_main_process: check if current process is the main process
def is_main_process():
return get_rank() == 0
# zero_first: zero first in distributed training
@contextmanager
def zero_first():
if not is_main_process():
dist.barrier()
yield
if is_main_process():
dist.barrier()
# empty_cuda_cache: empty cuda cache
def empty_cuda_cache():
gc.collect()
torch.cuda.empty_cache()
@contextmanager
def log_duration(name):
start = time.time()
try:
yield
finally:
print(f'{name}: {time.time()-start:.3f}')
# load_safetensors: load safetensors file
def load_safetensors(path):
tensors = {}
with safe_open(path, framework="pt", device="cpu") as f:
for key in f.keys():
tensors[key] = f.get_tensor(key)
return tensors
+101
View File
@@ -0,0 +1,101 @@
import os
import importlib.util
from pathlib import Path
import json
class Config(dict):
"""A dictionary that allows dot notation access for configuration settings."""
def __init__(self, data=None):
super().__init__()
if data:
for key, value in data.items():
# Set the key-value pair using the key and the converted value
self[key] = self._convert(value)
def _convert(self, value):
"""Recursively convert nested dictionaries to Config."""
if isinstance(value, dict):
return Config(value) # Convert all nested dicts to Config
elif isinstance(value, list):
return [self._convert(item) for item in value] # Convert items in lists
elif isinstance(value, Path):
return str(value) # Convert Path objects to string
return value
def __getattr__(self, item):
"""Allow access to dictionary keys via dot notation."""
if item in self:
return self[item]
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{item}'")
def __setattr__(self, key, value):
"""Allow setting dictionary keys via dot notation."""
self[key] = self._convert(value)
@classmethod
def load_config(cls, file_path):
"""Load a Python config file and return it as a Config instance."""
spec = importlib.util.spec_from_file_location("config_module", file_path)
config_module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(config_module)
if hasattr(config_module, "config") and isinstance(config_module.config, dict):
return cls(config_module.config) # Ensure full conversion of nested dicts
else:
raise ValueError("The config file does not define a 'config' dictionary.")
def update_config(self, new_data):
"""Recursively update Config with new dictionary values."""
for key, value in new_data.items():
if isinstance(value, dict) and isinstance(self.get(key), Config):
self[key].update_config(value) # Recursive update for nested dicts
else:
self[key] = self._convert(value) # Convert and assign
def dump_config(self, file_path=None):
"""Dump the Config object into a .py file with a Pythonic and readable format."""
# Ensure the provided path is valid
if file_path is None:
file_path = self.log.output_dir + '/config.py'
else:
dir_name = os.path.dirname(file_path)
if dir_name and not os.path.exists(dir_name):
os.makedirs(dir_name)
# Convert the Config instance into a regular dictionary
def config_to_dict(config):
"""Recursively convert a Config instance into a regular dictionary."""
if isinstance(config, Config):
return {key: config_to_dict(value) if isinstance(value, Config) else value
for key, value in config.items()}
return config
config_dict = config_to_dict(self)
# Write the config dictionary to a .py file in a formatted, readable way
with open(file_path, 'w') as f:
# Use json.dumps to pretty-print the dictionary with indentation and spaces
f.write("config = ")
f.write(json.dumps(config_dict, indent=4)) # Pretty print with indentation
f.write("\n")
print(f"Config has been saved to {file_path}")
def prepare_log(self):
def make_folder(folder):
if not os.path.exists(folder):
os.makedirs(folder)
if self.log.output_dir is not None:
make_folder(self.log.output_dir)
if self.log.model_dir is not None:
make_folder(self.log.model_dir)
if self.log.log_dir is not None:
make_folder(self.log.log_dir)
if self.log.result_dir is not None:
make_folder(self.log.result_dir)
+807
View File
@@ -0,0 +1,807 @@
from pathlib import Path
import os.path
import random
from collections import defaultdict
import math
import os
import numpy as np
import torch
from deepspeed.utils.logging import logger
from deepspeed import comm as dist
import datasets
from datasets.fingerprint import Hasher
from PIL import Image
import imageio
import multiprocess as mp
from utils.common import is_main_process, VIDEO_EXTENSIONS, log_duration
DEBUG = False
IMAGE_SIZE_ROUND_TO_MULTIPLE = 32
NUM_PROC = min(8, os.cpu_count())
def shuffle_with_seed(l, seed=None):
rng_state = random.getstate()
random.seed(seed)
random.shuffle(l)
random.setstate(rng_state)
def process_caption_fn(shuffle_tags=False, caption_prefix=''):
def fn(example):
with open(example['caption_file']) as f:
caption = f.read().strip()
if shuffle_tags:
tags = [tag.strip() for tag in caption.split(',')]
random.shuffle(tags)
caption = ', '.join(tags)
caption = caption_prefix + caption
example['caption'] = caption
return example
return fn
def round_to_multiple(x, multiple):
return int(round(x / multiple) * multiple)
def _map_and_cache(dataset, map_fn, cache_dir, cache_file_prefix='', new_fingerprint_args=None, regenerate_cache=False, caching_batch_size=1, with_indices=False):
# Do the fingerprinting ourselves, because otherwise map() does it by serializing the map function.
# That goes poorly when the function is capturing huge models (slow, OOMs, etc).
new_fingerprint_args = [] if new_fingerprint_args is None else new_fingerprint_args
new_fingerprint_args.append(dataset._fingerprint)
new_fingerprint = Hasher.hash(new_fingerprint_args)
cache_file = cache_dir / f'{cache_file_prefix}{new_fingerprint}.arrow'
cache_file = str(cache_file)
dataset = dataset.map(
map_fn,
cache_file_name=cache_file,
load_from_cache_file=(not regenerate_cache),
writer_batch_size=100,
new_fingerprint=new_fingerprint,
remove_columns=dataset.column_names,
batched=True,
batch_size=caching_batch_size,
with_indices=with_indices,
num_proc=NUM_PROC,
)
dataset.set_format('torch')
return dataset
# The smallest unit of a dataset. Represents a single size bucket from a single folder of images
# and captions on disk. Not batched; returns individual items.
class SizeBucketDataset:
def __init__(self, metadata_dataset, directory_config, size_bucket, model_name):
self.metadata_dataset = metadata_dataset
self.directory_config = directory_config
self.size_bucket = size_bucket
self.model_name = model_name
self.path = Path(self.directory_config['path'])
self.cache_dir = self.path / 'cache' / self.model_name / f'cache_{size_bucket[0]}x{size_bucket[1]}x{size_bucket[2]}'
os.makedirs(self.cache_dir, exist_ok=True)
self.text_embedding_datasets = []
self.num_repeats = self.directory_config.get('num_repeats', 1)
def cache_latents(self, map_fn, regenerate_cache=False, caching_batch_size=1):
print(f'caching latents: {self.size_bucket}')
self.latent_dataset = _map_and_cache(
self.metadata_dataset,
map_fn,
self.cache_dir,
cache_file_prefix='latents_',
regenerate_cache=regenerate_cache,
caching_batch_size=caching_batch_size,
with_indices=True,
)
# Shuffle again, since one media file can produce multiple training examples. E.g. video, or maybe
# in the future data augmentation. Don't need to shuffle text embeddings since those are looked
# up by index.
self.latent_dataset = self.latent_dataset.shuffle(seed=123)
# TODO: should we do dataset.flatten_indices() to make it contiguous on disk again?
# self.latent_dataset = self.latent_dataset.flatten_indices(
# cache_file_name=str(self.cache_dir / 'latents_flattened.arrow')
# )
def add_text_embedding_dataset(self, te_dataset):
self.text_embedding_datasets.append(te_dataset)
def __getitem__(self, idx):
idx = idx % len(self.latent_dataset)
ret = self.latent_dataset[idx]
te_idx = ret['te_idx'].item()
if DEBUG:
print(Path(self.metadata_dataset[te_idx]['image_file']).stem)
for ds in self.text_embedding_datasets:
ret.update(ds[te_idx])
return ret
def __len__(self):
return len(self.latent_dataset) * self.num_repeats
# Logical concatenation of multiple SizeBucketDataset, for the same size bucket. It returns items
# as batches.
class ConcatenatedBatchedDataset:
def __init__(self, datasets):
self.datasets = datasets
self.post_init_called = False
def post_init(self, batch_size):
iteration_order = []
for i, ds in enumerate(self.datasets):
print(i, len(ds))
iteration_order.extend([i]*len(ds))
shuffle_with_seed(iteration_order, 0)
cumulative_sums = [0] * len(self.datasets)
for k, dataset_idx in enumerate(iteration_order):
iteration_order[k] = (dataset_idx, cumulative_sums[dataset_idx])
cumulative_sums[dataset_idx] += 1
self.iteration_order = iteration_order
assert len(self.iteration_order) > 0, 'ConcatenatedBatchedDataset is empty. Are your file paths correct?'
self.batch_size = batch_size
self._make_divisible_by(self.batch_size)
self.post_init_called = True
def __len__(self):
assert self.post_init_called
return len(self.iteration_order) // self.batch_size
def __getitem__(self, idx):
assert self.post_init_called
start = idx * self.batch_size
end = start + self.batch_size
return [self.datasets[i][j] for i, j in self.iteration_order[start:end]]
def _make_divisible_by(self, n):
new_length = (len(self.iteration_order) // n) * n
self.iteration_order = self.iteration_order[:new_length]
if new_length == 0 and is_main_process():
logger.warning(f"size bucket {self.datasets[0].size_bucket} is being completely dropped because it doesn't have enough images")
class ARBucketDataset:
def __init__(self, ar_frames, resolutions, metadata_dataset, directory_config, model_name):
self.ar_frames = ar_frames
self.resolutions = resolutions
self.metadata_dataset = metadata_dataset
self.directory_config = directory_config
self.model_name = model_name
self.size_buckets = []
self.path = Path(directory_config['path'])
self.cache_dir = self.path / 'cache' / self.model_name / f'ar_frames_{self.ar_frames[0]:.3f}_{self.ar_frames[1]}'
os.makedirs(self.cache_dir, exist_ok=True)
for res in resolutions:
area = res**2
w = math.sqrt(area * self.ar_frames[0])
h = area / w
w = round_to_multiple(w, IMAGE_SIZE_ROUND_TO_MULTIPLE)
h = round_to_multiple(h, IMAGE_SIZE_ROUND_TO_MULTIPLE)
size_bucket = (w, h, self.ar_frames[1])
metadata_with_size_bucket = self.metadata_dataset.map(lambda example: {'size_bucket': size_bucket}, keep_in_memory=True)
sub_dataset = SizeBucketDataset(metadata_with_size_bucket, directory_config, size_bucket, model_name)
print(size_bucket)
self.size_buckets.append(
sub_dataset
)
def get_size_bucket_datasets(self):
return self.size_buckets
def cache_latents(self, map_fn, regenerate_cache=False, caching_batch_size=1):
print(f'caching latents: {self.ar_frames}')
for ds in self.size_buckets:
ds.cache_latents(map_fn, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
def cache_text_embeddings(self, map_fn, i, regenerate_cache=False, caching_batch_size=1):
print(f'caching text embeddings: {self.ar_frames}')
te_dataset = _map_and_cache(
self.metadata_dataset,
map_fn,
self.cache_dir,
cache_file_prefix=f'text_embeddings_{i}_',
new_fingerprint_args=[i],
regenerate_cache=regenerate_cache,
caching_batch_size=caching_batch_size,
)
for size_bucket_dataset in self.size_buckets:
size_bucket_dataset.add_text_embedding_dataset(te_dataset)
class DirectoryDataset:
def __init__(self, directory_config, dataset_config, model_name):
self._set_defaults(directory_config, dataset_config)
self.directory_config = directory_config
self.dataset_config = dataset_config
self.model_name = model_name
self.enable_ar_bucket = directory_config.get('enable_ar_bucket', dataset_config.get('enable_ar_bucket', False))
self.resolutions = self._process_user_provided_resolutions(
directory_config.get('resolutions', dataset_config['resolutions'])
)
self.path = Path(self.directory_config['path'])
self.cache_dir = self.path / 'cache' / self.model_name
if not self.path.exists() or not self.path.is_dir():
raise RuntimeError(f'Invalid path: {self.path}')
if not self.enable_ar_bucket:
self.ars = np.array([1.0])
elif ars := self.directory_config.get('ar_buckets', self.dataset_config.get('ar_buckets', None)):
self.ars = self._process_user_provided_ars(ars)
else:
min_ar = self.directory_config.get('min_ar', self.dataset_config['min_ar'])
max_ar = self.directory_config.get('max_ar', self.dataset_config['max_ar'])
num_ar_buckets = self.directory_config.get('num_ar_buckets', self.dataset_config['num_ar_buckets'])
self.ars = np.geomspace(min_ar, max_ar, num=num_ar_buckets)
frame_buckets = self.directory_config.get('frame_buckets', self.dataset_config.get('frame_buckets', [1]))
if 1 not in frame_buckets:
# always have an image bucket for convenience
frame_buckets.append(1)
frame_buckets.sort()
self.frame_buckets = np.array(frame_buckets)
def cache_metadata(self, regenerate_cache=False):
files = list(self.path.glob('*'))
# deterministic order
files.sort()
image_files = []
caption_files = []
for file in files:
if not file.is_file() or file.suffix == '.txt' or file.suffix == '.npz':
continue
image_file = file
caption_file = image_file.with_suffix('.txt')
if not os.path.exists(caption_file):
logger.warning(f'Image file {image_file} does not have corresponding caption file.')
caption_file = ''
image_files.append(str(image_file))
caption_files.append(str(caption_file))
assert len(image_files) > 0, f'Directory {self.path} had no images/videos!'
metadata_dataset = datasets.Dataset.from_dict({'image_file': image_files, 'caption_file': caption_files})
# Shuffle the data. Use a fixed seed, so the dataset is identical on all processes.
# Processes other than rank 0 will then load it from cache.
metadata_dataset = metadata_dataset.shuffle(seed=0)
metadata_map_fn = self._metadata_map_fn(self.ars, self.frame_buckets)
fingerprint = Hasher.hash([metadata_dataset._fingerprint, metadata_map_fn])
print('caching metadata')
metadata_dataset = metadata_dataset.map(
metadata_map_fn,
cache_file_name=str(self.cache_dir / f'metadata/metadata_{fingerprint}.arrow'),
load_from_cache_file=(not regenerate_cache),
batched=True,
batch_size=1,
num_proc=NUM_PROC,
remove_columns=metadata_dataset.column_names,
)
grouped_metadata = defaultdict(lambda: defaultdict(list))
for example in metadata_dataset:
ar_bucket = example['ar_bucket']
ar_bucket = (ar_bucket[0], int(ar_bucket[1]))
d = grouped_metadata[ar_bucket]
for k, v in example.items():
d[k].append(v)
self.ar_buckets = []
for ar_bucket, metadata in grouped_metadata.items():
metadata = datasets.Dataset.from_dict(metadata)
self.ar_buckets.append(
ARBucketDataset(
ar_bucket,
self.resolutions,
metadata,
self.directory_config,
self.model_name,
)
)
def _set_defaults(self, directory_config, dataset_config):
directory_config.setdefault('enable_ar_bucket', dataset_config.get('enable_ar_bucket', False))
directory_config.setdefault('resolutions', dataset_config['resolutions'])
directory_config.setdefault('shuffle_tags', dataset_config.get('shuffle_tags', False))
directory_config.setdefault('caption_prefix', dataset_config.get('caption_prefix', ''))
def _metadata_map_fn(self, ars, frame_buckets):
log_ars = np.log(ars)
def fn(example):
# batch size always 1
caption_file = example['caption_file'][0]
image_file = example['image_file'][0]
if not caption_file:
caption = ''
else:
with open(caption_file) as f:
caption = f.read().strip()
if self.directory_config['shuffle_tags']:
tags = [tag.strip() for tag in caption.split(',')]
random.shuffle(tags)
caption = ', '.join(tags)
caption = self.directory_config['caption_prefix'] + caption
empty_return = {'image_file': [], 'caption': [], 'ar_bucket': [], 'is_video': []}
image_file = Path(image_file)
try:
if image_file.suffix in VIDEO_EXTENSIONS:
# 100% accurate frame count, but much slower.
# frames = 0
# for frame in imageio.v3.imiter(image_file):
# frames += 1
# height, width = frame.shape[:2]
# TODO: this is an estimate of frame count. What happens if variable frame rate? Is
# it still close enough?
meta = imageio.v3.immeta(image_file)
height, width = meta['size']
frames = int(meta['fps'] * meta['duration'])
else:
pil_img = Image.open(image_file)
width, height = pil_img.size
frames = 1
except Exception:
logger.warning(f'Image file {image_file} could not be opened. Skipping.')
return empty_return
is_video = (frames > 1)
log_ar = np.log(width / height)
# Best AR bucket is the one with the smallest AR difference in log space.
i = np.argmin(np.abs(log_ar - log_ars))
# find closest frame bucket where the number of frames is greater than or equal to the bucket
diffs = frames - frame_buckets
positive_diffs = diffs[diffs >= 0]
if len(positive_diffs) == 0:
# video not long enough to find any valid frame bucket
print(f'video with frames={frames} is being skipped because it is too short')
return empty_return
j = np.argmin(positive_diffs)
if is_video and frame_buckets[j] == 1:
# don't let video be mapped to the image frame bucket
print(f'video with frames={frames} is being skipped because it is too short')
return empty_return
ar_bucket = (ars[i], frame_buckets[j])
return {'image_file': [str(image_file)], 'caption': [caption], 'ar_bucket': [ar_bucket], 'is_video': [is_video]}
return fn
def _process_user_provided_ars(self, ars):
ar_buckets = set()
for ar in ars:
if isinstance(ar, (tuple, list)):
assert len(ar) == 2
ar = round(ar[0] / ar[1], 6)
ar_buckets.add(ar)
ar_buckets = list(ar_buckets)
ar_buckets.sort()
return np.array(ar_buckets)
def _process_user_provided_resolutions(self, resolutions):
result = set()
for res in resolutions:
if isinstance(res, (tuple, list)):
assert len(res) == 2
res = round(math.sqrt(res[0] * res[1]), 6)
result.add(res)
result = list(result)
result.sort()
return result
def get_size_bucket_datasets(self):
result = []
for ar_bucket_dataset in self.ar_buckets:
result.extend(ar_bucket_dataset.get_size_bucket_datasets())
return result
def cache_latents(self, map_fn, regenerate_cache=False, caching_batch_size=1):
print(f'caching latents: {self.path}')
for ds in self.ar_buckets:
ds.cache_latents(map_fn, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
def cache_text_embeddings(self, map_fn, i, regenerate_cache=False, caching_batch_size=1):
for ds in self.ar_buckets:
ds.cache_text_embeddings(map_fn, i, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
# Outermost dataset object that the caller uses. Contains multiple ConcatenatedBatchedDataset. Responsible
# for returning the correct batch for the process's data parallel rank. Calls model.prepare_inputs so the
# returned tuple of tensors is whatever the model needs.
class Dataset:
def __init__(self, dataset_config, model_name):
super().__init__()
self.dataset_config = dataset_config
self.model_name = model_name
self.post_init_called = False
self.eval_quantile = None
self.directory_datasets = []
for directory_config in dataset_config['directory']:
directory_dataset = DirectoryDataset(directory_config, dataset_config, model_name)
self.directory_datasets.append(directory_dataset)
def post_init(self, data_parallel_rank, data_parallel_world_size, per_device_batch_size, gradient_accumulation_steps):
self.data_parallel_rank = data_parallel_rank
self.data_parallel_world_size = data_parallel_world_size
self.batch_size = per_device_batch_size * gradient_accumulation_steps
self.global_batch_size = self.data_parallel_world_size * self.batch_size
# group same size_bucket together
datasets_by_size_bucket = defaultdict(list)
for directory_dataset in self.directory_datasets:
for size_bucket_dataset in directory_dataset.get_size_bucket_datasets():
datasets_by_size_bucket[size_bucket_dataset.size_bucket].append(size_bucket_dataset)
self.buckets = []
for datasets in datasets_by_size_bucket.values():
print(len(datasets))
self.buckets.append(ConcatenatedBatchedDataset(datasets))
for bucket in self.buckets:
bucket.post_init(self.global_batch_size)
iteration_order = []
for i, bucket in enumerate(self.buckets):
iteration_order.extend([i]*(len(bucket)))
shuffle_with_seed(iteration_order, 0)
cumulative_sums = [0] * len(self.buckets)
for k, dataset_idx in enumerate(iteration_order):
iteration_order[k] = (dataset_idx, cumulative_sums[dataset_idx])
cumulative_sums[dataset_idx] += 1
self.iteration_order = iteration_order
if DEBUG:
print(f'Dataset iteration_order: {self.iteration_order}')
self.post_init_called = True
if subsample_ratio := self.dataset_config.get('subsample_ratio', None):
new_len = int(len(self) * subsample_ratio)
self.iteration_order = self.iteration_order[:new_len]
def set_eval_quantile(self, quantile):
self.eval_quantile = quantile
def __len__(self):
assert self.post_init_called
return len(self.iteration_order)
def __getitem__(self, idx):
assert self.post_init_called
i, j = self.iteration_order[idx]
examples = self.buckets[i][j]
start_idx = self.data_parallel_rank*self.batch_size
examples_for_this_dp_rank = examples[start_idx:start_idx+self.batch_size]
if DEBUG:
print((start_idx, start_idx+self.batch_size))
batch = self._collate(examples_for_this_dp_rank)
return batch
# collates a list of dictionaries of tensors into a single dictionary of batched tensors
def _collate(self, examples):
ret = {}
for key in examples[0].keys():
ret[key] = torch.stack([example[key] for example in examples])
return ret
def cache_metadata(self, regenerate_cache=False):
for ds in self.directory_datasets:
ds.cache_metadata(regenerate_cache=regenerate_cache)
def cache_latents(self, map_fn, regenerate_cache=False, caching_batch_size=1):
for ds in self.directory_datasets:
ds.cache_latents(map_fn, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
def cache_text_embeddings(self, map_fn, i, regenerate_cache=False, caching_batch_size=1):
for ds in self.directory_datasets:
ds.cache_text_embeddings(map_fn, i, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
def _cache_fn(datasets, queue, preprocess_media_file_fn, num_text_encoders, regenerate_cache, caching_batch_size):
# Dataset map() starts a bunch of processes. Make sure torch uses a limited number of threads
# to avoid CPU contention.
# TODO: if we ever change Datasets map to use spawn instead of fork, this might not work.
#torch.set_num_threads(os.cpu_count() // NUM_PROC)
# HF Datasets map can randomly hang if this is greater than one (???)
# See https://github.com/pytorch/pytorch/issues/10996
# Alternatively, we could try fixing this by using spawn instead of fork.
torch.set_num_threads(1)
for ds in datasets:
ds.cache_metadata(regenerate_cache=regenerate_cache)
def latents_map_fn(example, indices):
first_size_bucket = example['size_bucket'][0]
tensors = []
te_idx = []
for idx, path, size_bucket in zip(indices, example['image_file'], example['size_bucket']):
assert size_bucket == first_size_bucket
items = preprocess_media_file_fn(path, size_bucket)
tensors.extend(items)
te_idx.extend([idx] * len(items))
if len(tensors) == 0:
return {'latents': [], 'te_idx': []}
caching_batch_size = len(example['image_file'])
results = defaultdict(list)
for i in range(0, len(tensors), caching_batch_size):
batched = torch.stack(tensors[i:i+caching_batch_size])
parent_conn, child_conn = mp.Pipe(duplex=False)
queue.put((0, batched, child_conn))
result = parent_conn.recv() # dict
for k, v in result.items():
results[k].append(v)
# concatenate the list of tensors at each key into one batched tensor
for k, v in results.items():
results[k] = torch.cat(v)
results['te_idx'] = te_idx
return results
for ds in datasets:
ds.cache_latents(latents_map_fn, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
for text_encoder_idx in range(num_text_encoders):
def text_embedding_map_fn(example):
parent_conn, child_conn = mp.Pipe(duplex=False)
queue.put((text_encoder_idx+1, example['caption'], example['is_video'], child_conn))
result = parent_conn.recv() # dict
return result
for ds in datasets:
ds.cache_text_embeddings(text_embedding_map_fn, text_encoder_idx+1, regenerate_cache=regenerate_cache, caching_batch_size=caching_batch_size)
# signal that we're done
queue.put(None)
# Helper class to make caching multiple datasets more efficient by moving
# models to GPU as few times as needed.
class DatasetManager:
def __init__(self, model, regenerate_cache=False, caching_batch_size=1):
self.model = model
self.vae = self.model.get_vae()
self.text_encoders = self.model.get_text_encoders()
self.submodels = [self.vae] + list(self.text_encoders)
self.call_vae_fn = self.model.get_call_vae_fn(self.vae)
self.call_text_encoder_fns = [self.model.get_call_text_encoder_fn(text_encoder) for text_encoder in self.text_encoders]
self.regenerate_cache = regenerate_cache
self.caching_batch_size = caching_batch_size
self.datasets = []
def register(self, dataset):
self.datasets.append(dataset)
# Some notes for myself:
# Use a manager queue, since that can be pickled and unpickled, and sent to other processes.
# IMPORTANT: we use multiprocess library (not Python multiprocessing!) just like HF Datasets does.
# After hours of debugging and looking up related issues, I have concluded multiprocessing is outright bugged
# for this use case. Something about making a manager queue and sending it to the caching process, and then
# further sending it to map() workers via the pickled map function, is broken. It gets through a lot of the caching,
# but eventually, inevitably, queue.put() will fail with BrokenPipeError. Switching from multiprocessing to multiprocess,
# which has basically the same API, and everything works perfectly. ¯\_(ツ)_/¯
def cache(self):
if is_main_process():
manager = mp.Manager()
queue = [manager.Queue()]
else:
queue = [None]
torch.distributed.broadcast_object_list(queue, src=0, group=dist.get_world_group())
queue = queue[0]
# start up a process to run through the dataset caching flow
if is_main_process():
process = mp.Process(
target=_cache_fn,
args=(
self.datasets,
queue,
self.model.get_preprocess_media_file_fn(),
len(self.text_encoders),
self.regenerate_cache,
self.caching_batch_size,
)
)
process.start()
# loop on the original processes (one per GPU) to handle tasks requiring GPU models (VAE, text encoders)
while True:
task = queue.get()
if task is None:
# Propagate None so all worker processes break out of this loop.
# This is safe because it's a FIFO queue. The first None always comes after all work items.
queue.put(None)
break
self._handle_task(task)
# Free memory in all unneeded submodels. This is easier than trying to delete every reference.
# TODO: check if this is actually freeing memory.
for model in self.submodels:
model.to('meta')
dist.barrier()
if is_main_process():
process.join()
# Now load all datasets from cache.
for ds in self.datasets:
ds.cache_metadata()
ds.cache_latents(None)
for i in range(1, len(self.text_encoders)+1):
ds.cache_text_embeddings(None, i)
@torch.no_grad()
def _handle_task(self, task):
id = task[0]
# moved needed submodel to cuda, and everything else to cpu
if next(self.submodels[id].parameters()).device.type != 'cuda':
for i, submodel in enumerate(self.submodels):
if i != id:
submodel.to('cpu')
self.submodels[id].to('cuda')
if id == 0:
tensor, pipe = task[1:]
results = self.call_vae_fn(tensor)
elif id > 0:
caption, is_video, pipe = task[1:]
results = self.call_text_encoder_fns[id-1](caption, is_video=is_video)
else:
raise RuntimeError()
# Need to move to CPU here. If we don't, we get this error:
# RuntimeError: Cannot re-initialize CUDA in forked subprocess. To use CUDA with multiprocessing, you must use the 'spawn' start method
# I think this is because HF Datasets uses the multiprocess library (different from Python multiprocessing!) so it will always use fork.
results = {k: v.to('cpu') for k, v in results.items()}
pipe.send(results)
def split_batch(batch, pieces):
example_tuple = batch
split_size = example_tuple[0].size(0) // pieces
split_examples = zip(*(torch.split(tensor, split_size) for tensor in example_tuple))
# Deepspeed works with a tuple of (features, labels), even if we don't provide a loss_fn to PipelineEngine,
# and instead compute the loss ourselves in the model. It's okay to just return None for the labels here.
return [(ex, None) for ex in split_examples]
# DataLoader that divides batches into microbatches for gradient accumulation steps when doing
# pipeline parallel training. Iterates indefinitely (deepspeed requirement). Keeps track of epoch.
# Updates epoch as soon as the final batch is returned (notably different from qlora-pipe).
class PipelineDataLoader:
def __init__(self, dataset, gradient_accumulation_steps, model, num_dataloader_workers=2):
self.model = model
self.dataset = dataset
self.gradient_accumulation_steps = gradient_accumulation_steps
self.num_dataloader_workers = num_dataloader_workers
self.iter_called = False
self.eval_quantile = None
self.epoch = 1
self.num_batches_pulled = 0
self.next_micro_batch = None
self.recreate_dataloader = False
# Be careful to only create the DataLoader some bounded number of times: https://github.com/pytorch/pytorch/issues/91252
self._create_dataloader()
self.data = self._pull_batches_from_dataloader()
def reset(self):
self.epoch = 1
self.num_batches_pulled = 0
self.next_micro_batch = None
self.data = self._pull_batches_from_dataloader()
def set_eval_quantile(self, quantile):
self.eval_quantile = quantile
def __iter__(self):
self.iter_called = True
return self
def __len__(self):
return len(self.dataset) * self.gradient_accumulation_steps
def __next__(self):
if self.next_micro_batch == None:
self.next_micro_batch = next(self.data)
ret = self.next_micro_batch
try:
self.next_micro_batch = next(self.data)
except StopIteration:
if self.recreate_dataloader:
self._create_dataloader()
self.recreate_dataloader = False
self.data = self._pull_batches_from_dataloader()
self.num_batches_pulled = 0
self.next_micro_batch = next(self.data)
self.epoch += 1
return ret
def _create_dataloader(self, skip_first_n_batches=None):
if skip_first_n_batches is not None:
sampler = SkipFirstNSampler(skip_first_n_batches, len(self.dataset))
else:
sampler = None
self.dataloader = torch.utils.data.DataLoader(
self.dataset,
pin_memory=True,
batch_size=None,
sampler=sampler,
num_workers=self.num_dataloader_workers,
persistent_workers=(self.num_dataloader_workers > 0),
)
def _pull_batches_from_dataloader(self):
for batch in self.dataloader:
batch = self.model.prepare_inputs(batch, timestep_quantile=self.eval_quantile)
self.num_batches_pulled += 1
for micro_batch in split_batch(batch, self.gradient_accumulation_steps):
yield micro_batch
# Only the first and last stages in the pipeline pull from the dataloader. Parts of the code need
# to know the epoch, so we synchronize the epoch so the processes that don't use the dataloader
# know the current epoch.
def sync_epoch(self):
process_group = dist.get_world_group()
result = [None] * dist.get_world_size(process_group)
torch.distributed.all_gather_object(result, self.epoch, group=process_group)
max_epoch = -1
for epoch in result:
max_epoch = max(epoch, max_epoch)
self.epoch = max_epoch
def state_dict(self):
return {
'epoch': self.epoch,
'num_batches_pulled': self.num_batches_pulled,
}
def load_state_dict(self, state_dict):
assert not self.iter_called
self.epoch = state_dict['epoch']
# -1 because by preloading the next micro_batch, it's always going to have one more batch
# pulled than the actual number of batches iterated by the caller.
self.num_batches_pulled = state_dict['num_batches_pulled'] - 1
self._create_dataloader(skip_first_n_batches=self.num_batches_pulled)
self.data = self._pull_batches_from_dataloader()
# Recreate the dataloader after the first pass so that it won't skip
# batches again (we only want it to skip batches the first time).
self.recreate_dataloader = True
class SkipFirstNSampler(torch.utils.data.Sampler):
def __init__(self, n, dataset_length):
super().__init__()
self.n = n
self.dataset_length = dataset_length
def __len__(self):
return self.dataset_length
def __iter__(self):
for i in range(self.n, self.dataset_length):
yield i
if __name__ == '__main__':
from utils import common
common.is_main_process = lambda: True
from contextlib import contextmanager
@contextmanager
def _zero_first():
yield
common.zero_first = _zero_first
from utils import dataset as dataset_util
dataset_util.DEBUG = True
from models import flux
model = flux.CustomFluxPipeline.from_pretrained('/data2/imagegen_models/FLUX.1-dev', torch_dtype=torch.bfloat16)
model.model_config = {'guidance': 1.0, 'dtype': torch.bfloat16}
import toml
dataset_manager = dataset_util.DatasetManager(model)
with open('/home/anon/code/diffusion-pipe-configs/datasets/tiny1.toml') as f:
dataset_config = toml.load(f)
train_data = dataset_util.Dataset(dataset_config, model)
dataset_manager.register(train_data)
dataset_manager.cache()
train_data.post_init(data_parallel_rank=0, data_parallel_world_size=1, per_device_batch_size=1, gradient_accumulation_steps=2)
print(f'Dataset length: {len(train_data)}')
for item in train_data:
pass
+330
View File
@@ -0,0 +1,330 @@
import os
import sys
import warnings
warnings.filterwarnings('ignore')
import numpy as np
import torch, random
import torch.distributed as dist
from PIL import Image
import pandas as pd
import librosa
import json
import cv2
def get_image_from_path(image_path):
"""
读取图片文件并返回一个 PIL Image 对象。
:param image_path: 图片文件的路径 (例如 PNG, JPG 等)
:return: PIL Image 对象,如果读取失败则返回 None
"""
try:
# 使用 Pillow 的 Image.open() 直接打开图片文件
img = Image.open(image_path)
# 为了确保输出格式与原函数一致 (RGB),我们进行转换。
# 这么做可以处理带透明度通道的 PNG 图片 (RGBA -> RGB)
# 或者其他色彩模式的图片。
return img.convert('RGB')
except FileNotFoundError:
print(f"错误: 无法找到文件 {image_path}")
return None
except Exception as e:
print(f"错误: 打开图片文件时发生未知错误 {image_path}: {e}")
return None
def _center_crop_to_aspect_ratio(img: Image.Image, target_ratio: float) -> Image.Image:
"""
辅助函数:将图片通过中心裁剪的方式调整到目标长宽比。
target_ratio: 宽度 / 高度
"""
w, h = img.size
current_ratio = w / h
if abs(current_ratio - target_ratio) < 1e-4: # 如果比例已经很接近,则无需裁剪
return img
if current_ratio > target_ratio: # 当前图片比目标更“宽”,需要裁掉左右两边
new_w = int(target_ratio * h)
left = (w - new_w) // 2
right = left + new_w
top, bottom = 0, h
else: # 当前图片比目标更“高”,需要裁掉上下两边
new_h = int(w / target_ratio)
top = (h - new_h) // 2
bottom = top + new_h
left, right = 0, w
img_cropped = img.crop((left, top, right, bottom))
print(f"INFO: 图片从 {img.size} 裁剪到 {img_cropped.size} 以匹配目标长宽比 {target_ratio:.3f}")
return img_cropped
def process_images_final(img1, img2, reference_aspect_from = 'small'):
"""
处理两张图片,严格满足以下所有条件:
1. 所有最终图片的宽高都是32的倍数。
2. 一张图的短边固定为256px,另一张固定为704px。
3. 两张图最终的长宽比【完全相等】。
4. 通过中心裁剪来统一长宽比。
:param img1: 第一个Pillow图片对象。
:param img2: 第二个Pillow图片对象。
:param reference_aspect_from: 以哪张图的长宽比为基准, 'small' 或 'large'。
:return: 一个包含两张处理后图片的元组 (img1_processed, img2_processed)。
"""
# 1. 识别小图和大图
w1, h1 = img1.size
w2, h2 = img2.size
short1, short2 = min(w1, h1), min(w2, h2)
if short1 < short2:
img_small_orig, img_large_orig, is_img1_small = img1, img2, True
else:
img_small_orig, img_large_orig, is_img1_small = img2, img1, False
# 2. 确定基准长宽比并裁剪非基准图
if reference_aspect_from == 'small':
ref_img, other_img = img_small_orig, img_large_orig
else:
ref_img, other_img = img_large_orig, img_small_orig
ref_w, ref_h = ref_img.size
target_aspect_ratio = ref_w / ref_h
other_img_cropped = _center_crop_to_aspect_ratio(other_img, target_aspect_ratio)
if reference_aspect_from == 'small':
img_small_src, img_large_src = ref_img, other_img_cropped
else:
img_small_src, img_large_src = other_img_cropped, ref_img
# 3. 核心计算:基于严格比例推导尺寸
is_horizontal = target_aspect_ratio > 1
# 3.1 计算小图的理想长边
if is_horizontal:
ideal_long_small = 256 * target_aspect_ratio
else:
ideal_long_small = 256 / target_aspect_ratio
# 3.2 【关键修正】将小图理想长边调整到最接近的128的倍数
# 这是为了确保大图的长边也能成为32的倍数 (128 = 32 * 4)
final_long_small = int(round(ideal_long_small / 128.0) * 128)
final_long_small = max(128, final_long_small) # 确保不为0
# 3.3 计算大图的最终长边,它现在必然是32的倍数
final_long_large = int(final_long_small * (704 / 256.0))
# 4. 组合最终尺寸
if is_horizontal:
final_size_small = (final_long_small, 256)
final_size_large = (final_long_large, 704)
else:
final_size_small = (256, final_long_small)
final_size_large = (704, final_long_large)
print(f"INFO: 小图理想长边 {ideal_long_small:.2f} -> 调整到128的倍数 -> {final_long_small}")
print(f"FINAL: 小图尺寸: {final_size_small}, 大图尺寸: {final_size_large}")
# 验证长宽比是否相等
final_ratio_small = final_size_small[0] / final_size_small[1]
final_ratio_large = final_size_large[0] / final_size_large[1]
# 5. 缩放图片到最终尺寸
img_small_final = img_small_src.resize(final_size_small, Image.Resampling.LANCZOS)
img_large_final = img_large_src.resize(final_size_large, Image.Resampling.LANCZOS)
# 6. 按原始输入顺序返回结果
if is_img1_small:
return img_small_final, img_large_final
else:
return img_large_final, img_small_final
def resize_images(img1: Image.Image, img2: Image.Image):
"""
根据特定规则缩放两个Pillow图片对象。(已修正版本)
"""
w1, h1 = img1.size
w2, h2 = img2.size
short1, long1 = (w1, h1) if w1 < h1 else (h1, w1)
short2, long2 = (w2, h2) if w2 < h2 else (h2, w2)
# 情况一:两个短边都是256 (此部分逻辑正确,无需修改)
if short1 == 256 and short2 == 256:
print("检测到情况一:两张图片的短边均为256。")
if long1 < long2:
target_size = img1.size
img2_resized = img2.resize(target_size, Image.Resampling.LANCZOS)
return img1, img2_resized
elif long2 < long1:
target_size = img2.size
img1_resized = img1.resize(target_size, Image.Resampling.LANCZOS)
return img1_resized, img2
else:
return img1, img2
# 情况二:一个短边是256,另一个是512 (此部分已修正)
elif (short1 == 256 and short2 == 512) or (short1 == 512 and short2 == 256):
print("检测到情况二:一张图片短边为256,另一张为512。")
if short1 == 256:
img_small, img_large = img1, img2
long_small, long_large_orig = long1, long2
else:
img_small, img_large = img2, img1
long_small, long_large_orig = long2, long1
# 计算大图的目标长边
target_long_large = long_small * 2
# 如果大图当前的长边已经是目标长度,则无需处理
if target_long_large == long_large_orig:
return img1, img2
w_large, h_large = img_large.size
# --- 核心修正逻辑 ---
# 直接将新的长边与固定的短边(512)组合,而不是通过比例计算
if w_large > h_large: # 如果宽度是长边, 高度就是固定的512
target_size_large = (target_long_large, h_large)
else: # 如果高度是长边, 宽度就是固定的512
target_size_large = (w_large, target_long_large)
# --- 修正结束 ---
print(f"小图尺寸: {img_small.size}, 大图将从 {img_large.size} 缩放到 {target_size_large}")
img_large_resized = img_large.resize(target_size_large, Image.Resampling.LANCZOS)
if short1 == 256:
return img_small, img_large_resized
else:
return img_large_resized, img_small
else:
raise ValueError("输入的图片尺寸不符合预设的两种情况。")
def resize_short_side(img, short_side=512):
w, h = img.size
if h < w:
new_h = short_side
new_w = int(w * short_side / h)
else:
new_w = short_side
new_h = int(h * short_side / w)
new_h = int(new_h // 32 * 32)
new_w = int(new_w // 32 * 32)
img = img.resize((new_w, new_h), Image.LANCZOS)
return img
def get_first_frame_from_video(video_path):
"""
读取视频文件的第一帧并返回一个 PIL Image 对象。
:param video_path: 视频文件的路径
:return: PIL Image 对象,如果读取失败则返回 None
"""
# 打开视频文件
cap = cv2.VideoCapture(video_path)
# 检查视频是否成功打开
if not cap.isOpened():
print(f"错误: 无法打开视频文件 {video_path}")
return None
# 读取第一帧
# cap.read() 返回一个元组 (布尔值, 帧)
# 布尔值表示是否成功读取,帧是图像数据
ret, frame = cap.read()
# 释放视频捕获对象,这很重要
cap.release()
if ret:
# OpenCV 读取的图像格式是 BGR (蓝, 绿, 红)
# PIL 和大多数其他库使用的格式是 RGB (红, 绿, 蓝)
# 所以我们需要进行颜色空间转换
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# 将 NumPy 数组格式的帧转换为 PIL Image 对象
return Image.fromarray(frame_rgb)
else:
print("错误: 无法从视频中读取第一帧")
return None
def get_all_frames_from_video(video_path: str):
"""
读取视频文件的所有帧并返回一个 PIL Image 对象的列表。
:param video_path: 视频文件的路径
:return: PIL Image 对象的列表,如果读取失败或视频为空则返回 None
"""
# 打开视频文件
cap = cv2.VideoCapture(video_path)
# 检查视频是否成功打开
if not cap.isOpened():
print(f"错误: 无法打开视频文件 {video_path}")
return None
frames = []
while True:
# 读取一帧
# cap.read() 返回一个元组 (布尔值, 帧)
# 布尔值表示是否成功读取,帧是图像数据
ret, frame = cap.read()
# 如果 ret 为 False,表示视频已结束或读取时发生错误
if not ret:
break
# OpenCV 读取的图像格式是 BGR (蓝, 绿, 红)
# PIL 和大多数其他库使用的格式是 RGB (红, 绿, 蓝)
# 所以我们需要进行颜色空间转换
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
# 将 NumPy 数组格式的帧转换为 PIL Image 对象并添加到列表中
frames.append(Image.fromarray(frame_rgb))
# 释放视频捕获对象,这很重要
cap.release()
if not frames:
print(f"错误: 未能从视频 {video_path} 中读取任何帧。")
return None
return frames
import torchvision.transforms.functional as TF
def read_obj_tensor_from_path(input_path=None, target_size=None, background_color=(255, 255, 255)):
if input_path is None:
final_background_color = background_color + (255,)
background = Image.new('RGBA', target_size, final_background_color)
elif not os.path.exists(input_path):
final_background_color = background_color + (255,)
background = Image.new('RGBA', target_size, final_background_color)
else:
with Image.open(input_path) as img:
# 确保图片是RGBA模式,以更好地处理透明度
img = img.convert("RGBA")
original_width, original_height = img.size
target_width, target_height = target_size
# 1. 计算缩放比例,确保图片能完整放入目标框内
ratio = min(target_width / original_width, target_height / original_height)
# 2. 计算缩放后的新尺寸
new_width = int(original_width * ratio)
new_height = int(original_height * ratio)
# 3. 高质量缩放图片
resized_img = img.resize((new_width, new_height), Image.Resampling.LANCZOS)
# 4. 创建一个新的纯色背景画布
# 注意:颜色需要是RGBA格式,所以白色是 (255, 255, 255, 255)
# 最后一个值255代表完全不透明
final_background_color = background_color + (255,)
background = Image.new('RGBA', target_size, final_background_color)
# 5. 计算粘贴位置,使其居中
paste_x = (target_width - new_width) // 2
paste_y = (target_height - new_height) // 2
# 6. 将缩放后的图片粘贴到背景画布上
# 第三个参数 `resized_img` 作为蒙版,可以正确处理PNG的透明通道
background.paste(resized_img, (paste_x, paste_y), resized_img)
return background.convert('RGB')
+77
View File
@@ -0,0 +1,77 @@
# copy/pasted from pytorch lightning
# https://github.com/Lightning-AI/lightning/blob/0d52f4577310b5a1624bed4d23d49e37fb05af9e/src/lightning_fabric/utilities/seed.py
# and
# https://github.com/Lightning-AI/lightning/blob/98f7696d1681974d34fad59c03b4b58d9524ed13/src/pytorch_lightning/utilities/seed.py
# Copyright The Lightning team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from contextlib import contextmanager
from typing import Generator, Dict, Any
import torch
import numpy as np
from random import getstate as python_get_rng_state
from random import setstate as python_set_rng_state
def _collect_rng_states(include_cuda: bool = True) -> Dict[str, Any]:
"""Collect the global random state of :mod:`torch`, :mod:`torch.cuda`, :mod:`numpy` and Python."""
states = {
"torch": torch.get_rng_state(),
"numpy": np.random.get_state(),
"python": python_get_rng_state(),
}
if include_cuda:
try:
states["torch.cuda"] = torch.cuda.get_rng_state_all()
except RuntimeError:
# CUDA initialization failure.
pass
return states
def _set_rng_states(rng_state_dict: Dict[str, Any]) -> None:
"""Set the global random state of :mod:`torch`, :mod:`torch.cuda`, :mod:`numpy` and Python in the current
process."""
torch.set_rng_state(rng_state_dict["torch"])
# torch.cuda rng_state is only included since v1.8.
if "torch.cuda" in rng_state_dict:
torch.cuda.set_rng_state_all(rng_state_dict["torch.cuda"])
np.random.set_state(rng_state_dict["numpy"])
version, state, gauss = rng_state_dict["python"]
python_set_rng_state((version, tuple(state), gauss))
@contextmanager
def isolate_rng(include_cuda: bool = True) -> Generator[None, None, None]:
"""A context manager that resets the global random state on exit to what it was before entering.
It supports isolating the states for PyTorch, Numpy, and Python built-in random number generators.
Args:
include_cuda: Whether to allow this function to also control the `torch.cuda` random number generator.
Set this to ``False`` when using the function in a forked process where CUDA re-initialization is
prohibited.
Example:
>>> import torch
>>> torch.manual_seed(1) # doctest: +ELLIPSIS
<torch._C.Generator object at ...>
>>> with isolate_rng():
... [torch.rand(1) for _ in range(3)]
[tensor([0.7576]), tensor([0.2793]), tensor([0.4031])]
>>> torch.rand(1)
tensor([0.7576])
"""
states = _collect_rng_states(include_cuda)
yield
_set_rng_states(states)
+185
View File
@@ -0,0 +1,185 @@
# Copyright (c) 2024 NVIDIA CORPORATION.
# Licensed under the MIT license.
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
# LICENSE is in incl_licenses directory.
import math
import os
import random
import torch
import torch.utils.data
import numpy as np
import librosa
from librosa.filters import mel as librosa_mel_fn
import pathlib
from tqdm import tqdm
from typing import List, Tuple, Optional
# from env import AttrDict
MAX_WAV_VALUE = 32767.0 # NOTE: 32768.0 -1 to prevent int16 overflow (results in popping sound in corner cases)
def dynamic_range_compression(x, C=1, clip_val=1e-5):
return np.log(np.clip(x, a_min=clip_val, a_max=None) * C)
def dynamic_range_decompression(x, C=1):
return np.exp(x) / C
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5):
return torch.log(torch.clamp(x, min=clip_val) * C)
def dynamic_range_decompression_torch(x, C=1):
return torch.exp(x) / C
def spectral_normalize_torch(magnitudes):
return dynamic_range_compression_torch(magnitudes)
def spectral_de_normalize_torch(magnitudes):
return dynamic_range_decompression_torch(magnitudes)
mel_basis_cache = {}
hann_window_cache = {}
def mel_spectrogram(
y: torch.Tensor,
n_fft: int,
num_mels: int,
sampling_rate: int,
hop_size: int,
win_size: int,
fmin: int,
fmax: int = None,
center: bool = False,
) -> torch.Tensor:
"""
Calculate the mel spectrogram of an input signal.
Args:
y (torch.Tensor): Input signal.
n_fft (int): FFT size.
num_mels (int): Number of mel bins.
sampling_rate (int): Sampling rate of the input signal.
hop_size (int): Hop size for STFT.
win_size (int): Window size for STFT.
fmin (int): Minimum frequency for mel filterbank.
center (bool): Whether to pad the input to center the frames. Default is False.
Returns:
torch.Tensor: Mel spectrogram.
"""
# if torch.min(y) < -1.0:
# print(f"[WARNING] Min value of input waveform signal is {torch.min(y)}")
# if torch.max(y) > 1.0:
# print(f"[WARNING] Max value of input waveform signal is {torch.max(y)}")
device = y.device
key = f"{n_fft}_{num_mels}_{sampling_rate}_{hop_size}_{win_size}_{fmin}_{fmax}_{device}"
if key not in mel_basis_cache:
mel = librosa_mel_fn(
sr=sampling_rate, n_fft=n_fft, n_mels=num_mels, fmin=fmin, fmax=fmax
)
mel_basis_cache[key] = torch.from_numpy(mel).float().to(device)
hann_window_cache[key] = torch.hann_window(win_size).to(device)
mel_basis = mel_basis_cache[key]
hann_window = hann_window_cache[key]
padding = (n_fft - hop_size) // 2
y = torch.nn.functional.pad(
y.unsqueeze(1), (padding, padding), mode="reflect"
).squeeze(1)
spec = torch.stft(
y,
n_fft,
hop_length=hop_size,
win_length=win_size,
window=hann_window,
center=center,
pad_mode="reflect",
normalized=False,
onesided=True,
return_complex=True,
)
spec = torch.sqrt(torch.view_as_real(spec).pow(2).sum(-1) + 1e-9)
mel_spec = torch.matmul(mel_basis, spec)
mel_spec = spectral_normalize_torch(mel_spec)
return mel_spec
def get_mel_spectrogram(wav, sampling_rate, hop_size):
"""
Generate mel spectrogram from a waveform using given hyperparameters.
Args:
wav (torch.Tensor): Input waveform.
h: Hyperparameters object with attributes n_fft, num_mels, sampling_rate, hop_size, win_size, fmin, fmax.
Returns:
torch.Tensor: Mel spectrogram.
"""
# return mel_spectrogram(
# wav,
# h.n_fft,
# h.num_mels,
# h.sampling_rate,
# h.hop_size,
# h.win_size,
# h.fmin,
# h.fmax,
# )
return mel_spectrogram(
wav,
1024,
80,
sampling_rate,
hop_size,
1024,
0,
8000
)
#1024 80 22000 220 1024 0 8000
def get_dataset_filelist(a):
training_files = []
validation_files = []
list_unseen_validation_files = []
with open(a.input_training_file, "r", encoding="utf-8") as fi:
training_files = [
os.path.join(a.input_wavs_dir, x.split("|")[0] + ".wav")
for x in fi.read().split("\n")
if len(x) > 0
]
print(f"first training file: {training_files[0]}")
with open(a.input_validation_file, "r", encoding="utf-8") as fi:
validation_files = [
os.path.join(a.input_wavs_dir, x.split("|")[0] + ".wav")
for x in fi.read().split("\n")
if len(x) > 0
]
print(f"first validation file: {validation_files[0]}")
for i in range(len(a.list_input_unseen_validation_file)):
with open(a.list_input_unseen_validation_file[i], "r", encoding="utf-8") as fi:
unseen_validation_files = [
os.path.join(a.list_input_unseen_wavs_dir[i], x.split("|")[0] + ".wav")
for x in fi.read().split("\n")
if len(x) > 0
]
print(
f"first unseen {i}th validation fileset: {unseen_validation_files[0]}"
)
list_unseen_validation_files.append(unseen_validation_files)
return training_files, validation_files, list_unseen_validation_files
+104
View File
@@ -0,0 +1,104 @@
import os
import sys
import json
import jieba
from pypinyin import Style, lazy_pinyin
from tqdm import tqdm
import multiprocessing
import pandas as pd
# convert_char_to_pinyin: convert char to pinyin
def convert_char_to_pinyin(text_list, polyphone=True):
if jieba.dt.initialized is False:
jieba.default_logger.setLevel(50) # CRITICAL
jieba.initialize()
final_text_list = []
custom_trans = str.maketrans(
{";": ",", "“": '"', "”": '"', "‘": "'", "’": "'"}
) # add custom trans here, to address oov
def is_chinese(c):
return (
"\u3100" <= c <= "\u9fff" # common chinese characters
)
# convert char to pinyin
for text in text_list:
char_list = []
text = text.translate(custom_trans)
for seg in jieba.cut(text):
seg_byte_len = len(bytes(seg, "UTF-8"))
if seg_byte_len == len(seg): # if pure alphabets and symbols
if char_list and seg_byte_len > 1 and char_list[-1] not in " :'\"":
char_list.append(" ")
char_list.extend(seg)
elif polyphone and seg_byte_len == 3 * len(seg): # if pure east asian characters
seg_ = lazy_pinyin(seg, style=Style.TONE3, tone_sandhi=True)
for i, c in enumerate(seg):
if is_chinese(c):
char_list.append(" ")
char_list.append(seg_[i])
else: # if mixed characters, alphabets and symbols
for c in seg:
if ord(c) < 256:
char_list.extend(c)
elif is_chinese(c):
char_list.append(" ")
char_list.extend(lazy_pinyin(c, style=Style.TONE3, tone_sandhi=True))
else:
char_list.append(c)
final_text_list.append(char_list)
return final_text_list
# run: run
def run(base_dir, raw_data, start_index, end_index, thread):
output_data = []
for idx in range(start_index, end_index):
if (idx-start_index) % 100 == 0:
print('processed {:d}/{:d}, ranging from [{:d}, {:d}], thread number {:d}'.format(
idx+1-start_index, end_index-start_index, start_index, end_index, thread))
try:
data = raw_data[idx]
text = open(f"{base_dir}/{data['text_file']}", "r").readline()
out_pinyin = convert_char_to_pinyin([text])
out_pinyin = out_pinyin[0]
output_data.append(
{
"base_dir": base_dir,
"audio_file": data["audio_file"],
"text_file": data["text_file"],
"pinyin": out_pinyin,
}
)
except:
pass
return output_data
# main: main
def main(base_dir, raw_data, output_file, num_proc=32):
processor_list = []
pool = multiprocessing.Pool(processes=num_proc)
for thr in range(num_proc):
start, end = len(raw_data) // num_proc * thr, len(raw_data) // num_proc * (thr + 1)
if thr == num_proc - 1:
end = len(raw_data)
processor_list.append(pool.apply_async(run, (base_dir, raw_data, start, end, thr, )))
pool.close()
pool.join()
total_file_info = []
for proc in processor_list:
total_file_info += proc.get()
df = pd.DataFrame(total_file_info)
df.to_csv(output_file, index=False)
# prepare_libritts: prepare libritts
if __name__ == "__main__":
input_json = "/apdcephfs_cq8/share_1367250/zixiangzhou/projects/VideoChat/data_pipeline/LibriTTS/LibriTTS_test.json"
with open(input_json, "r") as f:
raw_data = json.load(f)
main("/apdcephfs_jn2/share_302243908", raw_data, "LibriTTS_test.csv", num_proc=64)
+366
View File
@@ -0,0 +1,366 @@
import torch
import numpy as np
from torch.nn import functional as F
from einops.einops import rearrange
def cam2pixel(cam_coord, f, c):
x = cam_coord[:, 0] / cam_coord[:, 2] * f[0] + c[0]
y = cam_coord[:, 1] / cam_coord[:, 2] * f[1] + c[1]
z = cam_coord[:, 2]
return np.stack((x, y, z), 1)
def pixel2cam(pixel_coord, f, c):
x = (pixel_coord[:, 0] - c[0]) / f[0] * pixel_coord[:, 2]
y = (pixel_coord[:, 1] - c[1]) / f[1] * pixel_coord[:, 2]
z = pixel_coord[:, 2]
return np.stack((x, y, z), 1)
def world2cam(world_coord, R, t):
cam_coord = np.dot(R, world_coord.transpose(1, 0)).transpose(1, 0) + t.reshape(1, 3)
return cam_coord
def cam2world(cam_coord, R, t):
world_coord = np.dot(np.linalg.inv(R), (cam_coord - t.reshape(1, 3)).transpose(1, 0)).transpose(1, 0)
return world_coord
def rigid_transform_3D(A, B):
n, dim = A.shape
centroid_A = np.mean(A, axis=0)
centroid_B = np.mean(B, axis=0)
H = np.dot(np.transpose(A - centroid_A), B - centroid_B) / n
U, s, V = np.linalg.svd(H)
R = np.dot(np.transpose(V), np.transpose(U))
if np.linalg.det(R) < 0:
s[-1] = -s[-1]
V[2] = -V[2]
R = np.dot(np.transpose(V), np.transpose(U))
varP = np.var(A, axis=0).sum()
c = 1 / varP * np.sum(s)
t = -np.dot(c * R, np.transpose(centroid_A)) + np.transpose(centroid_B)
return c, R, t
def rigid_align(A, B):
c, R, t = rigid_transform_3D(A, B)
A2 = np.transpose(np.dot(c * R, np.transpose(A))) + t
return A2
def transform_joint_to_other_db(src_joint, src_name, dst_name):
src_joint_num = len(src_name)
dst_joint_num = len(dst_name)
new_joint = np.zeros(((dst_joint_num,) + src_joint.shape[1:]), dtype=np.float32)
for src_idx in range(len(src_name)):
name = src_name[src_idx]
if name in dst_name:
dst_idx = dst_name.index(name)
new_joint[dst_idx] = src_joint[src_idx]
return new_joint
def rotation_matrix_to_angle_axis(rotation_matrix):
"""Convert 3x4 rotation matrix to Rodrigues vector
Args:
rotation_matrix (Tensor): rotation matrix.
Returns:
Tensor: Rodrigues vector transformation.
Shape:
- Input: :math:`(N, 3, 4)`
- Output: :math:`(N, 3)`
Example:
>>> input = torch.rand(2, 3, 4) # Nx4x4
>>> output = tgm.rotation_matrix_to_angle_axis(input) # Nx3
"""
# todo add check that matrix is a valid rotation matrix
quaternion = rotation_matrix_to_quaternion(rotation_matrix)
return quaternion_to_angle_axis(quaternion)
def quaternion_to_angle_axis(quaternion: torch.Tensor) -> torch.Tensor:
"""Convert quaternion vector to angle axis of rotation.
Adapted from ceres C++ library: ceres-solver/include/ceres/rotation.h
Args:
quaternion (torch.Tensor): tensor with quaternions.
Return:
torch.Tensor: tensor with angle axis of rotation.
Shape:
- Input: :math:`(*, 4)` where `*` means, any number of dimensions
- Output: :math:`(*, 3)`
Example:
>>> quaternion = torch.rand(2, 4) # Nx4
>>> angle_axis = tgm.quaternion_to_angle_axis(quaternion) # Nx3
"""
if not torch.is_tensor(quaternion):
raise TypeError("Input type is not a torch.Tensor. Got {}".format(
type(quaternion)))
if not quaternion.shape[-1] == 4:
raise ValueError("Input must be a tensor of shape Nx4 or 4. Got {}"
.format(quaternion.shape))
# unpack input and compute conversion
q1: torch.Tensor = quaternion[..., 1]
q2: torch.Tensor = quaternion[..., 2]
q3: torch.Tensor = quaternion[..., 3]
sin_squared_theta: torch.Tensor = q1 * q1 + q2 * q2 + q3 * q3
sin_theta: torch.Tensor = torch.sqrt(sin_squared_theta)
cos_theta: torch.Tensor = quaternion[..., 0]
two_theta: torch.Tensor = 2.0 * torch.where(
cos_theta < 0.0,
torch.atan2(-sin_theta, -cos_theta),
torch.atan2(sin_theta, cos_theta))
k_pos: torch.Tensor = two_theta / sin_theta
k_neg: torch.Tensor = 2.0 * torch.ones_like(sin_theta)
k: torch.Tensor = torch.where(sin_squared_theta > 0.0, k_pos, k_neg)
angle_axis: torch.Tensor = torch.zeros_like(quaternion)[..., :3]
angle_axis[..., 0] += q1 * k
angle_axis[..., 1] += q2 * k
angle_axis[..., 2] += q3 * k
return angle_axis
def rotation_matrix_to_quaternion(rotation_matrix, eps=1e-6):
"""Convert 3x4 rotation matrix to 4d quaternion vector
This algorithm is based on algorithm described in
https://github.com/KieranWynn/pyquaternion/blob/master/pyquaternion/quaternion.py#L201
Args:
rotation_matrix (Tensor): the rotation matrix to convert.
Return:
Tensor: the rotation in quaternion
Shape:
- Input: :math:`(N, 3, 4)`
- Output: :math:`(N, 4)`
Example:
>>> input = torch.rand(4, 3, 4) # Nx3x4
>>> output = tgm.rotation_matrix_to_quaternion(input) # Nx4
"""
if not torch.is_tensor(rotation_matrix):
raise TypeError("Input type is not a torch.Tensor. Got {}".format(
type(rotation_matrix)))
input_shape = rotation_matrix.shape
if len(input_shape) == 2:
rotation_matrix = rotation_matrix.unsqueeze(0)
if len(rotation_matrix.shape) > 3:
raise ValueError(
"Input size must be a three dimensional tensor. Got {}".format(
rotation_matrix.shape))
if not rotation_matrix.shape[-2:] == (3, 4):
raise ValueError(
"Input size must be a N x 3 x 4 tensor. Got {}".format(
rotation_matrix.shape))
rmat_t = torch.transpose(rotation_matrix, 1, 2)
mask_d2 = rmat_t[:, 2, 2] < eps
mask_d0_d1 = rmat_t[:, 0, 0] > rmat_t[:, 1, 1]
mask_d0_nd1 = rmat_t[:, 0, 0] < -rmat_t[:, 1, 1]
t0 = 1 + rmat_t[:, 0, 0] - rmat_t[:, 1, 1] - rmat_t[:, 2, 2]
q0 = torch.stack([rmat_t[:, 1, 2] - rmat_t[:, 2, 1],
t0, rmat_t[:, 0, 1] + rmat_t[:, 1, 0],
rmat_t[:, 2, 0] + rmat_t[:, 0, 2]], -1)
t0_rep = t0.repeat(4, 1).t()
t1 = 1 - rmat_t[:, 0, 0] + rmat_t[:, 1, 1] - rmat_t[:, 2, 2]
q1 = torch.stack([rmat_t[:, 2, 0] - rmat_t[:, 0, 2],
rmat_t[:, 0, 1] + rmat_t[:, 1, 0],
t1, rmat_t[:, 1, 2] + rmat_t[:, 2, 1]], -1)
t1_rep = t1.repeat(4, 1).t()
t2 = 1 - rmat_t[:, 0, 0] - rmat_t[:, 1, 1] + rmat_t[:, 2, 2]
q2 = torch.stack([rmat_t[:, 0, 1] - rmat_t[:, 1, 0],
rmat_t[:, 2, 0] + rmat_t[:, 0, 2],
rmat_t[:, 1, 2] + rmat_t[:, 2, 1], t2], -1)
t2_rep = t2.repeat(4, 1).t()
t3 = 1 + rmat_t[:, 0, 0] + rmat_t[:, 1, 1] + rmat_t[:, 2, 2]
q3 = torch.stack([t3, rmat_t[:, 1, 2] - rmat_t[:, 2, 1],
rmat_t[:, 2, 0] - rmat_t[:, 0, 2],
rmat_t[:, 0, 1] - rmat_t[:, 1, 0]], -1)
t3_rep = t3.repeat(4, 1).t()
mask_c0 = mask_d2.float() * mask_d0_d1.float()
mask_c1 = mask_d2.float() * (1 - mask_d0_d1.float())
mask_c2 = (1 - mask_d2.float()) * mask_d0_nd1.float()
mask_c3 = (1 - mask_d2.float()) * (1 - mask_d0_nd1.float())
mask_c0 = mask_c0.view(-1, 1).type_as(q0)
mask_c1 = mask_c1.view(-1, 1).type_as(q1)
mask_c2 = mask_c2.view(-1, 1).type_as(q2)
mask_c3 = mask_c3.view(-1, 1).type_as(q3)
q = q0 * mask_c0 + q1 * mask_c1 + q2 * mask_c2 + q3 * mask_c3
q /= torch.sqrt(t0_rep * mask_c0 + t1_rep * mask_c1 + # noqa
t2_rep * mask_c2 + t3_rep * mask_c3)
q *= 0.5
if len(input_shape) == 2:
q = q.squeeze(0)
return q
def rot6d_to_axis_angle(x):
batch_size = x.shape[0]
x = x.view(-1, 3, 2)
a1 = x[:, :, 0]
a2 = x[:, :, 1]
b1 = F.normalize(a1)
b2 = F.normalize(a2 - torch.einsum('bi,bi->b', b1, a2).unsqueeze(-1) * b1)
b3 = torch.cross(b1, b2)
rot_mat = torch.stack((b1, b2, b3), dim=-1) # 3x3 rotation matrix
rot_mat = torch.cat([rot_mat, torch.zeros((batch_size, 3, 1)).cuda().float()], 2) # 3x4 rotation matrix
axis_angle = rotation_matrix_to_angle_axis(rot_mat).reshape(-1, 3) # axis-angle
axis_angle[torch.isnan(axis_angle)] = 0.0
return axis_angle
def rot6d_to_rotmat(x):
"""Convert 6D rotation representation to 3x3 rotation matrix.
Based on Zhou et al., "On the Continuity of Rotation Representations in Neural Networks", CVPR 2019
Input:
(B,6) Batch of 6-D rotation representations
Output:
(B,3,3) Batch of corresponding rotation matrices
"""
if x.shape[-1] == 6:
batch_size = x.shape[0]
if len(x.shape) == 3:
num = x.shape[1]
x = rearrange(x, 'b n d -> (b n) d', d=6)
else:
num = 1
x = rearrange(x, 'b (k l) -> b k l', k=3, l=2)
# x = x.view(-1,3,2)
a1 = x[:, :, 0]
a2 = x[:, :, 1]
b1 = F.normalize(a1)
b2 = F.normalize(a2 - torch.einsum('bi,bi->b', b1, a2).unsqueeze(-1) * b1)
b3 = torch.cross(b1, b2, dim=-1)
mat = torch.stack((b1, b2, b3), dim=-1)
if num > 1:
mat = rearrange(mat, '(b n) h w-> b n h w', b=batch_size, n=num, h=3, w=3)
else:
x = x.view(-1,3,2)
a1 = x[:, :, 0]
a2 = x[:, :, 1]
b1 = F.normalize(a1)
b2 = F.normalize(a2 - torch.einsum('bi,bi->b', b1, a2).unsqueeze(-1) * b1)
b3 = torch.cross(b1, b2, dim=-1)
mat = torch.stack((b1, b2, b3), dim=-1)
return mat
def batch_rodrigues(theta):
"""Convert axis-angle representation to rotation matrix.
Args:
theta: size = [B, 3]
Returns:
Rotation matrix corresponding to the quaternion -- size = [B, 3, 3]
"""
l1norm = torch.norm(theta + 1e-8, p = 2, dim = 1)
angle = torch.unsqueeze(l1norm, -1)
normalized = torch.div(theta, angle)
angle = angle * 0.5
v_cos = torch.cos(angle)
v_sin = torch.sin(angle)
quat = torch.cat([v_cos, v_sin * normalized], dim = 1)
return quat_to_rotmat(quat)
def quat_to_rotmat(quat):
"""Convert quaternion coefficients to rotation matrix.
Args:
quat: size = [B, 4] 4 <===>(w, x, y, z)
Returns:
Rotation matrix corresponding to the quaternion -- size = [B, 3, 3]
"""
norm_quat = quat
norm_quat = norm_quat/norm_quat.norm(p=2, dim=1, keepdim=True)
w, x, y, z = norm_quat[:,0], norm_quat[:,1], norm_quat[:,2], norm_quat[:,3]
B = quat.size(0)
w2, x2, y2, z2 = w.pow(2), x.pow(2), y.pow(2), z.pow(2)
wx, wy, wz = w*x, w*y, w*z
xy, xz, yz = x*y, x*z, y*z
rotMat = torch.stack([w2 + x2 - y2 - z2, 2*xy - 2*wz, 2*wy + 2*xz,
2*wz + 2*xy, w2 - x2 + y2 - z2, 2*yz - 2*wx,
2*xz - 2*wy, 2*wx + 2*yz, w2 - x2 - y2 + z2], dim=1).view(B, 3, 3)
return rotMat
def sample_joint_features(img_feat, joint_xy):
height, width = img_feat.shape[2:]
x = joint_xy[:, :, 0] / (width - 1) * 2 - 1
y = joint_xy[:, :, 1] / (height - 1) * 2 - 1
grid = torch.stack((x, y), 2)[:, :, None, :]
img_feat = F.grid_sample(img_feat, grid, align_corners=True)[:, :, :, 0] # batch_size, channel_dim, joint_num
img_feat = img_feat.permute(0, 2, 1).contiguous() # batch_size, joint_num, channel_dim
return img_feat
def soft_argmax_2d(heatmap2d):
batch_size = heatmap2d.shape[0]
height, width = heatmap2d.shape[2:]
heatmap2d = heatmap2d.reshape((batch_size, -1, height * width))
heatmap2d = F.softmax(heatmap2d, 2)
heatmap2d = heatmap2d.reshape((batch_size, -1, height, width))
accu_x = heatmap2d.sum(dim=(2))
accu_y = heatmap2d.sum(dim=(3))
accu_x = accu_x * torch.arange(width).float().cuda()[None, None, :]
accu_y = accu_y * torch.arange(height).float().cuda()[None, None, :]
accu_x = accu_x.sum(dim=2, keepdim=True)
accu_y = accu_y.sum(dim=2, keepdim=True)
coord_out = torch.cat((accu_x, accu_y), dim=2)
return coord_out
def soft_argmax_3d(heatmap3d):
batch_size = heatmap3d.shape[0]
depth, height, width = heatmap3d.shape[2:]
heatmap3d = heatmap3d.reshape((batch_size, -1, depth * height * width))
heatmap3d = F.softmax(heatmap3d, 2)
heatmap3d = heatmap3d.reshape((batch_size, -1, depth, height, width))
accu_x = heatmap3d.sum(dim=(2, 3))
accu_y = heatmap3d.sum(dim=(2, 4))
accu_z = heatmap3d.sum(dim=(3, 4))
accu_x = accu_x * torch.arange(width).float().cuda()[None, None, :]
accu_y = accu_y * torch.arange(height).float().cuda()[None, None, :]
accu_z = accu_z * torch.arange(depth).float().cuda()[None, None, :]
accu_x = accu_x.sum(dim=2, keepdim=True)
accu_y = accu_y.sum(dim=2, keepdim=True)
accu_z = accu_z.sum(dim=2, keepdim=True)
coord_out = torch.cat((accu_x, accu_y, accu_z), dim=2)
return coord_out
+12
View File
@@ -0,0 +1,12 @@
from . import configs, distributed, modules
from .modules.clip import CLIPModel
from .modules.vae2_2 import Wan2_2_VAE
from .modules.t5 import T5EncoderModel
from .modules.model_tia2mv_rope_back import WanModel as WanModelTIA2MVROPEBack
from .utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from .modules.fc_model import AudioEmbedding
from .tia2mv_obj_back_id_prefix import WanTIA2MVRefBackIDPrefix
# import modules
+59
View File
@@ -0,0 +1,59 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import copy
import os
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
from .wan_i2v_14B import i2v_14B
from .wan_t2v_1_3B import t2v_1_3B
from .wan_t2v_14B import t2v_14B
from .wan_i2v_1_3B import i2v_1_3B
from .wan_i2v_1_3B_audio import i2v_1_3B_audio
from .wan_i2v_1_3B_avideo import i2v_1_3B_avideo
from .wan_ti2v_5B import ti2v_5B
from .wan_i2v_A14B import i2v_A14B
# the config of t2i_14B is the same as t2v_14B
t2i_14B = copy.deepcopy(t2v_14B)
t2i_14B.__name__ = 'Config: Wan T2I 14B'
# the config of flf2v_14B is the same as i2v_14B
# the config of flf2v_14B is the same as i2v_14B
flf2v_14B = copy.deepcopy(i2v_14B)
flf2v_14B.__name__ = 'Config: Wan FLF2V 14B'
flf2v_14B.sample_neg_prompt = "镜头切换," + flf2v_14B.sample_neg_prompt
WAN_CONFIGS = {
't2v-14B': t2v_14B,
't2v-1.3B': t2v_1_3B,
'i2v-14B': i2v_14B,
'i2v-1.3B': i2v_1_3B,
'i2v-1.3B-audio': i2v_1_3B_audio,
'i2v-1.3B-avideo': i2v_1_3B_avideo,
't2i-14B': t2i_14B,
'ti2v-5B': ti2v_5B, #wan2.2
'ti2v-A14B': i2v_A14B, #wan2.2
'flf2v-14B': flf2v_14B
}
SIZE_CONFIGS = {
'720*1280': (720, 1280),
'1280*720': (1280, 720),
'480*832': (480, 832),
'832*480': (832, 480),
'1024*1024': (1024, 1024),
}
MAX_AREA_CONFIGS = {
'720*1280': 720 * 1280,
'1280*720': 1280 * 720,
'480*832': 480 * 832,
'832*480': 832 * 480,
}
SUPPORTED_SIZES = {
't2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
't2v-1.3B': ('480*832', '832*480'),
'i2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
'flf2v-14B': ('720*1280', '1280*720', '480*832', '832*480'),
't2i-14B': tuple(SIZE_CONFIGS.keys()),
}
+22
View File
@@ -0,0 +1,22 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
#------------------------ Wan shared config ------------------------#
wan_shared_cfg = EasyDict()
# t5
wan_shared_cfg.t5_model = 'umt5_xxl'
wan_shared_cfg.t5_dtype = torch.bfloat16
wan_shared_cfg.text_len = 512
# transformer
wan_shared_cfg.param_dtype = torch.bfloat16
# inference
wan_shared_cfg.num_train_timesteps = 1000
wan_shared_cfg.sample_fps = 16
wan_shared_cfg.sample_neg_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,' \
'整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,' \
'画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,' \
'手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
+36
View File
@@ -0,0 +1,36 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan I2V 14B ------------------------#
i2v_14B = EasyDict(__name__='Config: Wan I2V 14B')
i2v_14B.update(wan_shared_cfg)
i2v_14B.sample_neg_prompt = "镜头晃动," + i2v_14B.sample_neg_prompt
i2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_14B.t5_tokenizer = 'google/umt5-xxl'
# clip
i2v_14B.clip_model = 'clip_xlm_roberta_vit_h_14'
i2v_14B.clip_dtype = torch.float16
i2v_14B.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
i2v_14B.clip_tokenizer = 'xlm-roberta-large'
# vae
i2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_14B.vae_stride = (4, 8, 8)
# transformer
i2v_14B.patch_size = (1, 2, 2)
i2v_14B.dim = 5120
i2v_14B.ffn_dim = 13824
i2v_14B.freq_dim = 256
i2v_14B.num_heads = 40
i2v_14B.num_layers = 40
i2v_14B.window_size = (-1, -1)
i2v_14B.qk_norm = True
i2v_14B.cross_attn_norm = True
i2v_14B.eps = 1e-6
+37
View File
@@ -0,0 +1,37 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan T2V 1.3B ------------------------#
i2v_1_3B = EasyDict(__name__='Config: Wan T2V 1.3B')
i2v_1_3B.update(wan_shared_cfg)
# t5
i2v_1_3B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_1_3B.t5_tokenizer = 'google/umt5-xxl'
# clip
i2v_1_3B.clip_model = 'clip_xlm_roberta_vit_h_14'
i2v_1_3B.clip_dtype = torch.float16
i2v_1_3B.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
i2v_1_3B.clip_tokenizer = 'xlm-roberta-large'
# vae
i2v_1_3B.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_1_3B.vae_stride = (4, 8, 8)
# transformer
i2v_1_3B.patch_size = (1, 2, 2)
i2v_1_3B.dim = 1536
i2v_1_3B.ffn_dim = 8960
i2v_1_3B.freq_dim = 256
i2v_1_3B.num_heads = 12
i2v_1_3B.num_layers = 30
i2v_1_3B.window_size = (-1, -1)
i2v_1_3B.qk_norm = True
i2v_1_3B.cross_attn_norm = True
i2v_1_3B.eps = 1e-6
i2v_1_3B.text_len = 512
+45
View File
@@ -0,0 +1,45 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan T2V 1.3B ------------------------#
i2v_1_3B_audio = EasyDict(__name__='Config: Wan I2V 1.3B Audio')
i2v_1_3B_audio.update(wan_shared_cfg)
# t5
i2v_1_3B_audio.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_1_3B_audio.t5_tokenizer = 'google/umt5-xxl'
# clip
i2v_1_3B_audio.clip_model = 'clip_xlm_roberta_vit_h_14'
i2v_1_3B_audio.clip_dtype = torch.float16
i2v_1_3B_audio.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
i2v_1_3B_audio.clip_tokenizer = 'xlm-roberta-large'
# vae
i2v_1_3B_audio.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_1_3B_audio.vae_stride = (4, 8, 8)
# audio
i2v_1_3B_audio.audio_embedder = "audio/fc_model.safetensors"
i2v_1_3B_audio.audio_encoder = "audio/whisper-tiny"
i2v_1_3B_audio.bigvgan = "audio/code2wav_bigvgan_model"
i2v_1_3B_audio.audio_mean = -5.1753
i2v_1_3B_audio.audio_std = 2.1544
i2v_1_3B_audio.sampling_rate = 24000
# transformer
i2v_1_3B_audio.patch_size = (1, 2, 2)
i2v_1_3B_audio.dim = 1536
i2v_1_3B_audio.ffn_dim = 8960
i2v_1_3B_audio.freq_dim = 256
i2v_1_3B_audio.num_heads = 12
i2v_1_3B_audio.num_layers = 30
i2v_1_3B_audio.window_size = (-1, -1)
i2v_1_3B_audio.qk_norm = True
i2v_1_3B_audio.cross_attn_norm = True
i2v_1_3B_audio.eps = 1e-6
i2v_1_3B_audio.text_len = 512
+45
View File
@@ -0,0 +1,45 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan T2V 1.3B ------------------------#
i2v_1_3B_avideo = EasyDict(__name__='Config: Wan I2V 1.3B Audio')
i2v_1_3B_avideo.update(wan_shared_cfg)
# t5
i2v_1_3B_avideo.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_1_3B_avideo.t5_tokenizer = 'google/umt5-xxl'
# clip
i2v_1_3B_avideo.clip_model = 'clip_xlm_roberta_vit_h_14'
i2v_1_3B_avideo.clip_dtype = torch.float16
i2v_1_3B_avideo.clip_checkpoint = 'models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth'
i2v_1_3B_avideo.clip_tokenizer = 'xlm-roberta-large'
# vae
i2v_1_3B_avideo.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_1_3B_avideo.vae_stride = (4, 8, 8)
# audio
i2v_1_3B_avideo.audio_embedder = "audio/fc_model.safetensors"
i2v_1_3B_avideo.audio_encoder = "audio/whisper-tiny"
i2v_1_3B_avideo.bigvgan = "audio/code2wav_bigvgan_model"
i2v_1_3B_avideo.audio_mean = -5.1753
i2v_1_3B_avideo.audio_std = 2.1544
i2v_1_3B_avideo.sampling_rate = 24000
# transformer
i2v_1_3B_avideo.patch_size = (1, 2, 2)
i2v_1_3B_avideo.dim = 1536
i2v_1_3B_avideo.ffn_dim = 8960
i2v_1_3B_avideo.freq_dim = 256
i2v_1_3B_avideo.num_heads = 12
i2v_1_3B_avideo.num_layers = 30
i2v_1_3B_avideo.window_size = (-1, -1)
i2v_1_3B_avideo.qk_norm = True
i2v_1_3B_avideo.cross_attn_norm = True
i2v_1_3B_avideo.eps = 1e-6
i2v_1_3B_avideo.text_len = 512
+38
View File
@@ -0,0 +1,38 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan I2V A14B ------------------------#
i2v_A14B = EasyDict(__name__='Config: Wan I2V A14B')
i2v_A14B.update(wan_shared_cfg)
i2v_A14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
i2v_A14B.t5_tokenizer = 'google/umt5-xxl'
# vae
i2v_A14B.vae_checkpoint = 'Wan2.1_VAE.pth'
i2v_A14B.vae_stride = (4, 8, 8)
# transformer
i2v_A14B.patch_size = (1, 2, 2)
i2v_A14B.dim = 5120
i2v_A14B.ffn_dim = 13824
i2v_A14B.freq_dim = 256
i2v_A14B.num_heads = 40
i2v_A14B.num_layers = 40
i2v_A14B.window_size = (-1, -1)
i2v_A14B.qk_norm = True
i2v_A14B.cross_attn_norm = True
i2v_A14B.eps = 1e-6
i2v_A14B.low_noise_checkpoint = 'low_noise_model'
i2v_A14B.high_noise_checkpoint = 'high_noise_model'
# inference
i2v_A14B.sample_shift = 5.0
i2v_A14B.sample_steps = 40
i2v_A14B.boundary = 0.900
i2v_A14B.sample_guide_scale = (3.5, 3.5) # low noise, high noise
+29
View File
@@ -0,0 +1,29 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan T2V 14B ------------------------#
t2v_14B = EasyDict(__name__='Config: Wan T2V 14B')
t2v_14B.update(wan_shared_cfg)
# t5
t2v_14B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
t2v_14B.t5_tokenizer = 'google/umt5-xxl'
# vae
t2v_14B.vae_checkpoint = 'Wan2.1_VAE.pth'
t2v_14B.vae_stride = (4, 8, 8)
# transformer
t2v_14B.patch_size = (1, 2, 2)
t2v_14B.dim = 5120
t2v_14B.ffn_dim = 13824
t2v_14B.freq_dim = 256
t2v_14B.num_heads = 40
t2v_14B.num_layers = 40
t2v_14B.window_size = (-1, -1)
t2v_14B.qk_norm = True
t2v_14B.cross_attn_norm = True
t2v_14B.eps = 1e-6
+29
View File
@@ -0,0 +1,29 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan T2V 1.3B ------------------------#
t2v_1_3B = EasyDict(__name__='Config: Wan T2V 1.3B')
t2v_1_3B.update(wan_shared_cfg)
# t5
t2v_1_3B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
t2v_1_3B.t5_tokenizer = 'google/umt5-xxl'
# vae
t2v_1_3B.vae_checkpoint = 'Wan2.1_VAE.pth'
t2v_1_3B.vae_stride = (4, 8, 8)
# transformer
t2v_1_3B.patch_size = (1, 2, 2)
t2v_1_3B.dim = 1536
t2v_1_3B.ffn_dim = 8960
t2v_1_3B.freq_dim = 256
t2v_1_3B.num_heads = 12
t2v_1_3B.num_layers = 30
t2v_1_3B.window_size = (-1, -1)
t2v_1_3B.qk_norm = True
t2v_1_3B.cross_attn_norm = True
t2v_1_3B.eps = 1e-6
+57
View File
@@ -0,0 +1,57 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from easydict import EasyDict
from .shared_config import wan_shared_cfg
#------------------------ Wan TI2V 5B ------------------------#
ti2v_5B = EasyDict(__name__='Config: Wan TI2V 5B')
ti2v_5B.update(wan_shared_cfg)
# t5
ti2v_5B.t5_checkpoint = 'models_t5_umt5-xxl-enc-bf16.pth'
ti2v_5B.t5_tokenizer = 'google/umt5-xxl'
# vae
ti2v_5B.vae_checkpoint = 'Wan2.2_VAE.pth'
ti2v_5B.motion_vae_checkpoint = 'Wan2.1_VAE.pth'
ti2v_5B.vae_stride = (4, 16, 16)
# transformer
ti2v_5B.patch_size = (1, 2, 2)
ti2v_5B.dim = 3072
ti2v_5B.ffn_dim = 14336
ti2v_5B.freq_dim = 256
ti2v_5B.text_dim = 4096
ti2v_5B.in_dim = 48 # Wan2.2 TI2V 5B uses 48-dim input (from Wan2.2_VAE)
ti2v_5B.out_dim = 48 # Wan2.2 TI2V 5B uses 48-dim output (to Wan2.2_VAE)
ti2v_5B.num_heads = 24
ti2v_5B.num_layers = 30
ti2v_5B.window_size = (-1, -1)
ti2v_5B.qk_norm = True
ti2v_5B.cross_attn_norm = True
ti2v_5B.eps = 1e-6
# Wan2.2 specific features
ti2v_5B.seperated_timestep = True
ti2v_5B.fuse_vae_embedding_in_latents = True
ti2v_5B.require_clip_embedding = False
ti2v_5B.require_vae_embedding = False
ti2v_5B.audio_dim = 1536
ti2v_5B.audio_ffn_dim = 8960
ti2v_5B.audio_num_heads = 12
ti2v_5B.audio_num_layers = 30
# inference
ti2v_5B.sample_fps = 24
ti2v_5B.sample_shift = 5.0
ti2v_5B.sample_steps = 50
ti2v_5B.sample_guide_scale = 5.0
ti2v_5B.frame_num = 121
ti2v_5B.audio_embedder = "audio/fc_model.safetensors"
ti2v_5B.audio_encoder = "audio/whisper-tiny"
ti2v_5B.bigvgan = "audio/code2wav_bigvgan_model"
ti2v_5B.audio_mean = -5.1753
ti2v_5B.audio_std = 2.1544
ti2v_5B.sampling_rate = 24000
View File
+43
View File
@@ -0,0 +1,43 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import gc
from functools import partial
import torch
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
from torch.distributed.utils import _free_storage
#FSDP
# shard_model: shard model FSDP
def shard_model(
model,
device_id,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
process_group=None,
sharding_strategy=ShardingStrategy.FULL_SHARD,
sync_module_states=True,
):
# shard model FSDP
model = FSDP(
module=model,
process_group=process_group,
sharding_strategy=sharding_strategy,
auto_wrap_policy=partial(
lambda_auto_wrap_policy, lambda_fn=lambda m: m in model.blocks),
mixed_precision=MixedPrecision(
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
buffer_dtype=buffer_dtype),
device_id=device_id,
sync_module_states=sync_module_states)
return model
# free_model: free model
def free_model(model):
for m in model.modules():
if isinstance(m, FSDP):
_free_storage(m._handle.flat_param.data)
del model
gc.collect()
torch.cuda.empty_cache()
+164
View File
@@ -0,0 +1,164 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
import torch.cuda.amp as amp
from xfuser.core.distributed import (get_sequence_parallel_rank,
get_sequence_parallel_world_size,
get_sp_group)
from xfuser.core.long_ctx_attention import xFuserLongContextAttention
from ..modules.model import sinusoidal_embedding_1d
def pad_freqs(original_tensor, target_len):
seq_len, s1, s2 = original_tensor.shape
pad_size = target_len - seq_len
padding_tensor = torch.ones(
pad_size,
s1,
s2,
dtype=original_tensor.dtype,
device=original_tensor.device)
padded_tensor = torch.cat([original_tensor, padding_tensor], dim=0)
return padded_tensor
def usp_dit_forward(
self,
x,
t,
context,
seq_len,
clip_fea=None,
y=None,
):
"""
x: A list of videos each with shape [C, T, H, W].
t: [B].
context: A list of text embeddings each with shape [L, C].
"""
if self.model_type == 'i2v':
assert clip_fea is not None and y is not None
# params
device = self.patch_embedding.weight.device
if self.freqs.device != device:
self.freqs = self.freqs.to(device)
if y is not None:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
# embeddings
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
assert seq_lens.max() <= seq_len
x = torch.cat([
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1)
for u in x
])
# time embeddings
with amp.autocast(dtype=torch.float32):
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).float())
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
assert e.dtype == torch.float32 and e0.dtype == torch.float32
# context
context_lens = None
context = self.text_embedding(
torch.stack([
torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]))
if clip_fea is not None:
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context = torch.concat([context_clip, context], dim=1)
# arguments
kwargs = dict(
e=e0,
seq_lens=seq_lens,
grid_sizes=grid_sizes,
freqs=self.freqs,
context=context,
context_lens=context_lens)
# Context Parallel
x = torch.chunk(
x, get_sequence_parallel_world_size(),
dim=1)[get_sequence_parallel_rank()]
# print('x shape:', x.shape)
iii = 0
for block in self.blocks:
iii+=1
# print('block :', iii, get_sequence_parallel_rank())
x = block(x, **kwargs)
# print('block ok :', iii, get_sequence_parallel_rank())
# head
x = self.head(x, e)
# Context Parallel
x = get_sp_group().all_gather(x, dim=1)
# unpatchify
x = self.unpatchify(x, grid_sizes)
return [u.float() for u in x]
def usp_attn_forward(self,
x,
seq_lens,
grid_sizes,
freqs,
dtype=torch.bfloat16):
# print('attn forward')
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
half_dtypes = (torch.float16, torch.bfloat16)
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# query, key, value function
def qkv_fn(x):
# print('com q fn:', x.shape)
qq = self.q(x)
# print('qq is ok:', qq.shape)
q = self.norm_q(qq).view(b, s, n, d)
# print('q is ok:', q.shape)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
return q, k, v
# print('com qkv:', x.shape)
q, k, v = qkv_fn(x)
# print('qkv is ok', q.shape, k.shape, v.shape)
q = rope_apply(q, grid_sizes, freqs)
k = rope_apply(k, grid_sizes, freqs)
# TODO: We should use unpaded q,k,v for attention.
# k_lens = seq_lens // get_sequence_parallel_world_size()
# if k_lens is not None:
# q = torch.cat([u[:l] for u, l in zip(q, k_lens)]).unsqueeze(0)
# k = torch.cat([u[:l] for u, l in zip(k, k_lens)]).unsqueeze(0)
# v = torch.cat([u[:l] for u, l in zip(v, k_lens)]).unsqueeze(0)
# print('self attn coming')
x = xFuserLongContextAttention()(
None,
query=half(q),
key=half(k),
value=half(v),
window_size=self.window_size)
# print('self attn is ok')
# TODO: padding after attention.
# x = torch.cat([x, x.new_zeros(b, s - x.size(1), n, d)], dim=1)
# output
x = x.flatten(2)
x = self.o(x)
return x
+20
View File
@@ -0,0 +1,20 @@
# import modules
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
from .attention import flash_attention
from .model_tia2mv_rope_back import WanModel
from .t5 import T5Decoder, T5Encoder, T5EncoderModel, T5Model
from .tokenizers import HuggingfaceTokenizer
from .vae2_2 import Wan2_2_VAE
__all__ = [
'Wan2_2_VAE',
'WanModel',
'T5Model',
'T5Encoder',
'T5Decoder',
'T5EncoderModel',
'HuggingfaceTokenizer',
'flash_attention',
]
+175
View File
@@ -0,0 +1,175 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except ModuleNotFoundError:
FLASH_ATTN_3_AVAILABLE = False
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except ModuleNotFoundError:
FLASH_ATTN_2_AVAILABLE = False
import warnings
__all__ = [
'flash_attention',
'attention',
]
def flash_attention(
q,
k,
v,
q_lens=None,
k_lens=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
version=None,
):
"""
q: [B, Lq, Nq, C1].
k: [B, Lk, Nk, C1].
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
q_lens: [B].
k_lens: [B].
dropout_p: float. Dropout probability.
softmax_scale: float. The scaling of QK^T before applying softmax.
causal: bool. Whether to apply causal attention mask.
window_size: (left right). If not (-1, -1), apply sliding window local attention.
deterministic: bool. If True, slightly slower and uses more memory.
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
"""
half_dtypes = (torch.float16, torch.bfloat16)
assert dtype in half_dtypes
assert q.device.type == 'cuda' and q.size(-1) <= 256
# params
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# preprocess query
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor(
[lq] * b, dtype=torch.int32).to(
device=q.device, non_blocking=True)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
# preprocess key, value
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor(
[lk] * b, dtype=torch.int32).to(
device=k.device, non_blocking=True)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
warnings.warn(
'Flash attention 3 is not available, use flash attention 2 instead.'
)
# apply attention
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
# Note: dropout_p, window_size are not supported in FA3 now.
x = flash_attn_interface.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
seqused_q=None,
seqused_k=None,
max_seqlen_q=lq,
max_seqlen_k=lk,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic)[0].unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
max_seqlen_q=lq,
max_seqlen_k=lk,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic).unflatten(0, (b, lq))
# output
return x.type(out_dtype)
def attention(
q,
k,
v,
q_lens=None,
k_lens=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
fa_version=None,
):
if FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE:
return flash_attention(
q=q,
k=k,
v=v,
q_lens=q_lens,
k_lens=k_lens,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
q_scale=q_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic,
dtype=dtype,
version=fa_version,
)
else:
attn_mask = None
q = q.transpose(1, 2).to(dtype)
k = k.transpose(1, 2).to(dtype)
v = v.transpose(1, 2).to(dtype)
out = torch.nn.functional.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
out = out.transpose(1, 2).contiguous()
return out
+65
View File
@@ -0,0 +1,65 @@
import sys
sys.path.append('/apdcephfs_cq10/share_1367250/terohu/code/ARwan')
import logging
import os
from argparse import ArgumentParser
from datetime import timedelta
from pathlib import Path
from utils.wanx_audio_dataset import VideoAudioTextLoader
import pandas as pd
import tensordict as td
import torch
import torch.distributed as distributed
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from mmaudio.ext.autoencoder import AutoEncoderModule
from mmaudio.ext.mel_converter import get_mel_converter
import torchaudio
log = logging.getLogger()
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
# 16k
SAMPLE_RATE = 16_000
NUM_SAMPLES = 16_000 * 8
tod_vae_ckpt = '/apdcephfs_cq10/share_1367250/terohu/code/MMAudio/ext_weights/v1-16.pth'
bigvgan_vocoder_ckpt = '/apdcephfs_cq10/share_1367250/terohu/code/MMAudio/ext_weights/best_netG.pt'
mode = '16k'
# 44k
"""
NOTE: 352800 (8*44100) is not divisible by (STFT hop size * VAE downsampling ratio) which is 1024.
353280 is the next integer divisible by 1024.
"""
@torch.inference_mode()
def main():
# 16k
tod = AutoEncoderModule(vae_ckpt_path=tod_vae_ckpt,
vocoder_ckpt_path=bigvgan_vocoder_ckpt,
mode=mode).eval().cuda()
# 44k
# mel_converter = get_mel_converter(mode).eval().cuda()
# waveforms=1
# mel = mel_converter(waveforms)
# dist = tod.encode(mel)
a_mean = dist.mean.detach().cpu().transpose(1, 2)
a_std = dist.std.detach().cpu().transpose(1, 2)
latent=a_mean+a_std*torch.randn_like(a_mean).to(a_mean)
mel=tod.decode(latent.transpose(1, 2))
audios=tod.vocode(mel)
audio = audios.float().cpu()[0]
torchaudio.save('tmp.wav', audio, 16000)
+542
View File
@@ -0,0 +1,542 @@
# Modified from ``https://github.com/openai/CLIP'' and ``https://github.com/mlfoundations/open_clip''
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import logging
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as T
from .attention import flash_attention
from .tokenizers import HuggingfaceTokenizer
from .xlm_roberta import XLMRoberta
__all__ = [
'XLMRobertaCLIP',
'clip_xlm_roberta_vit_h_14',
'CLIPModel',
]
# pos_interpolate: interpolate position embeddings
def pos_interpolate(pos, seq_len):
if pos.size(1) == seq_len:
return pos
else:
src_grid = int(math.sqrt(pos.size(1)))
tar_grid = int(math.sqrt(seq_len))
n = pos.size(1) - src_grid * src_grid
return torch.cat([
pos[:, :n],
F.interpolate(
pos[:, n:].float().reshape(1, src_grid, src_grid, -1).permute(
0, 3, 1, 2),
size=(tar_grid, tar_grid),
mode='bicubic',
align_corners=False).flatten(2).transpose(1, 2)
],
dim=1)
# QuickGELU: Quick GELU activation function
class QuickGELU(nn.Module):
def forward(self, x):
return x * torch.sigmoid(1.702 * x)
# LayerNorm: Layer normalization
class LayerNorm(nn.LayerNorm):
def forward(self, x):
return super().forward(x.float()).type_as(x)
class SelfAttention(nn.Module):
def __init__(self,
dim,
num_heads,
causal=False,
attn_dropout=0.0,
proj_dropout=0.0):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.causal = causal
self.attn_dropout = attn_dropout
self.proj_dropout = proj_dropout
# layers
self.to_qkv = nn.Linear(dim, dim * 3)
self.proj = nn.Linear(dim, dim)
def forward(self, x):
"""
x: [B, L, C].
"""
b, s, c, n, d = *x.size(), self.num_heads, self.head_dim
# compute query, key, value
q, k, v = self.to_qkv(x).view(b, s, 3, n, d).unbind(2)
# compute attention
p = self.attn_dropout if self.training else 0.0
x = flash_attention(q, k, v, dropout_p=p, causal=self.causal, version=2)
x = x.reshape(b, s, c)
# output
x = self.proj(x)
x = F.dropout(x, self.proj_dropout, self.training)
return x
class SwiGLU(nn.Module):
def __init__(self, dim, mid_dim):
super().__init__()
self.dim = dim
self.mid_dim = mid_dim
# layers
self.fc1 = nn.Linear(dim, mid_dim)
self.fc2 = nn.Linear(dim, mid_dim)
self.fc3 = nn.Linear(mid_dim, dim)
def forward(self, x):
x = F.silu(self.fc1(x)) * self.fc2(x)
x = self.fc3(x)
return x
class AttentionBlock(nn.Module):
def __init__(self,
dim,
mlp_ratio,
num_heads,
post_norm=False,
causal=False,
activation='quick_gelu',
attn_dropout=0.0,
proj_dropout=0.0,
norm_eps=1e-5):
assert activation in ['quick_gelu', 'gelu', 'swi_glu']
super().__init__()
self.dim = dim
self.mlp_ratio = mlp_ratio
self.num_heads = num_heads
self.post_norm = post_norm
self.causal = causal
self.norm_eps = norm_eps
# layers
self.norm1 = LayerNorm(dim, eps=norm_eps)
self.attn = SelfAttention(dim, num_heads, causal, attn_dropout,
proj_dropout)
self.norm2 = LayerNorm(dim, eps=norm_eps)
if activation == 'swi_glu':
self.mlp = SwiGLU(dim, int(dim * mlp_ratio))
else:
self.mlp = nn.Sequential(
nn.Linear(dim, int(dim * mlp_ratio)),
QuickGELU() if activation == 'quick_gelu' else nn.GELU(),
nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))
def forward(self, x):
if self.post_norm:
x = x + self.norm1(self.attn(x))
x = x + self.norm2(self.mlp(x))
else:
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x
class AttentionPool(nn.Module):
def __init__(self,
dim,
mlp_ratio,
num_heads,
activation='gelu',
proj_dropout=0.0,
norm_eps=1e-5):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.mlp_ratio = mlp_ratio
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.proj_dropout = proj_dropout
self.norm_eps = norm_eps
# layers
gain = 1.0 / math.sqrt(dim)
self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))
self.to_q = nn.Linear(dim, dim)
self.to_kv = nn.Linear(dim, dim * 2)
self.proj = nn.Linear(dim, dim)
self.norm = LayerNorm(dim, eps=norm_eps)
self.mlp = nn.Sequential(
nn.Linear(dim, int(dim * mlp_ratio)),
QuickGELU() if activation == 'quick_gelu' else nn.GELU(),
nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(proj_dropout))
def forward(self, x):
"""
x: [B, L, C].
"""
b, s, c, n, d = *x.size(), self.num_heads, self.head_dim
# compute query, key, value
q = self.to_q(self.cls_embedding).view(1, 1, n, d).expand(b, -1, -1, -1)
k, v = self.to_kv(x).view(b, s, 2, n, d).unbind(2)
# compute attention
x = flash_attention(q, k, v, version=2)
x = x.reshape(b, 1, c)
# output
x = self.proj(x)
x = F.dropout(x, self.proj_dropout, self.training)
# mlp
x = x + self.mlp(self.norm(x))
return x[:, 0]
class VisionTransformer(nn.Module):
def __init__(self,
image_size=224,
patch_size=16,
dim=768,
mlp_ratio=4,
out_dim=512,
num_heads=12,
num_layers=12,
pool_type='token',
pre_norm=True,
post_norm=False,
activation='quick_gelu',
attn_dropout=0.0,
proj_dropout=0.0,
embedding_dropout=0.0,
norm_eps=1e-5):
if image_size % patch_size != 0:
print(
'[WARNING] image_size is not divisible by patch_size',
flush=True)
assert pool_type in ('token', 'token_fc', 'attn_pool')
out_dim = out_dim or dim
super().__init__()
self.image_size = image_size
self.patch_size = patch_size
self.num_patches = (image_size // patch_size)**2
self.dim = dim
self.mlp_ratio = mlp_ratio
self.out_dim = out_dim
self.num_heads = num_heads
self.num_layers = num_layers
self.pool_type = pool_type
self.post_norm = post_norm
self.norm_eps = norm_eps
# embeddings
gain = 1.0 / math.sqrt(dim)
self.patch_embedding = nn.Conv2d(
3,
dim,
kernel_size=patch_size,
stride=patch_size,
bias=not pre_norm)
if pool_type in ('token', 'token_fc'):
self.cls_embedding = nn.Parameter(gain * torch.randn(1, 1, dim))
self.pos_embedding = nn.Parameter(gain * torch.randn(
1, self.num_patches +
(1 if pool_type in ('token', 'token_fc') else 0), dim))
self.dropout = nn.Dropout(embedding_dropout)
# transformer
self.pre_norm = LayerNorm(dim, eps=norm_eps) if pre_norm else None
self.transformer = nn.Sequential(*[
AttentionBlock(dim, mlp_ratio, num_heads, post_norm, False,
activation, attn_dropout, proj_dropout, norm_eps)
for _ in range(num_layers)
])
self.post_norm = LayerNorm(dim, eps=norm_eps)
# head
if pool_type == 'token':
self.head = nn.Parameter(gain * torch.randn(dim, out_dim))
elif pool_type == 'token_fc':
self.head = nn.Linear(dim, out_dim)
elif pool_type == 'attn_pool':
self.head = AttentionPool(dim, mlp_ratio, num_heads, activation,
proj_dropout, norm_eps)
def forward(self, x, interpolation=False, use_31_block=False):
b = x.size(0)
# embeddings
x = self.patch_embedding(x).flatten(2).permute(0, 2, 1)
if self.pool_type in ('token', 'token_fc'):
x = torch.cat([self.cls_embedding.expand(b, -1, -1), x], dim=1)
if interpolation:
e = pos_interpolate(self.pos_embedding, x.size(1))
else:
e = self.pos_embedding
x = self.dropout(x + e)
if self.pre_norm is not None:
x = self.pre_norm(x)
# transformer
if use_31_block:
x = self.transformer[:-1](x)
return x
else:
x = self.transformer(x)
return x
class XLMRobertaWithHead(XLMRoberta):
def __init__(self, **kwargs):
self.out_dim = kwargs.pop('out_dim')
super().__init__(**kwargs)
# head
mid_dim = (self.dim + self.out_dim) // 2
self.head = nn.Sequential(
nn.Linear(self.dim, mid_dim, bias=False), nn.GELU(),
nn.Linear(mid_dim, self.out_dim, bias=False))
def forward(self, ids):
# xlm-roberta
x = super().forward(ids)
# average pooling
mask = ids.ne(self.pad_id).unsqueeze(-1).to(x)
x = (x * mask).sum(dim=1) / mask.sum(dim=1)
# head
x = self.head(x)
return x
class XLMRobertaCLIP(nn.Module):
def __init__(self,
embed_dim=1024,
image_size=224,
patch_size=14,
vision_dim=1280,
vision_mlp_ratio=4,
vision_heads=16,
vision_layers=32,
vision_pool='token',
vision_pre_norm=True,
vision_post_norm=False,
activation='gelu',
vocab_size=250002,
max_text_len=514,
type_size=1,
pad_id=1,
text_dim=1024,
text_heads=16,
text_layers=24,
text_post_norm=True,
text_dropout=0.1,
attn_dropout=0.0,
proj_dropout=0.0,
embedding_dropout=0.0,
norm_eps=1e-5):
super().__init__()
self.embed_dim = embed_dim
self.image_size = image_size
self.patch_size = patch_size
self.vision_dim = vision_dim
self.vision_mlp_ratio = vision_mlp_ratio
self.vision_heads = vision_heads
self.vision_layers = vision_layers
self.vision_pre_norm = vision_pre_norm
self.vision_post_norm = vision_post_norm
self.activation = activation
self.vocab_size = vocab_size
self.max_text_len = max_text_len
self.type_size = type_size
self.pad_id = pad_id
self.text_dim = text_dim
self.text_heads = text_heads
self.text_layers = text_layers
self.text_post_norm = text_post_norm
self.norm_eps = norm_eps
# models
self.visual = VisionTransformer(
image_size=image_size,
patch_size=patch_size,
dim=vision_dim,
mlp_ratio=vision_mlp_ratio,
out_dim=embed_dim,
num_heads=vision_heads,
num_layers=vision_layers,
pool_type=vision_pool,
pre_norm=vision_pre_norm,
post_norm=vision_post_norm,
activation=activation,
attn_dropout=attn_dropout,
proj_dropout=proj_dropout,
embedding_dropout=embedding_dropout,
norm_eps=norm_eps)
self.textual = XLMRobertaWithHead(
vocab_size=vocab_size,
max_seq_len=max_text_len,
type_size=type_size,
pad_id=pad_id,
dim=text_dim,
out_dim=embed_dim,
num_heads=text_heads,
num_layers=text_layers,
post_norm=text_post_norm,
dropout=text_dropout)
self.log_scale = nn.Parameter(math.log(1 / 0.07) * torch.ones([]))
def forward(self, imgs, txt_ids):
"""
imgs: [B, 3, H, W] of torch.float32.
- mean: [0.48145466, 0.4578275, 0.40821073]
- std: [0.26862954, 0.26130258, 0.27577711]
txt_ids: [B, L] of torch.long.
Encoded by data.CLIPTokenizer.
"""
xi = self.visual(imgs)
xt = self.textual(txt_ids)
return xi, xt
def param_groups(self):
groups = [{
'params': [
p for n, p in self.named_parameters()
if 'norm' in n or n.endswith('bias')
],
'weight_decay': 0.0
}, {
'params': [
p for n, p in self.named_parameters()
if not ('norm' in n or n.endswith('bias'))
]
}]
return groups
def _clip(pretrained=False,
pretrained_name=None,
model_cls=XLMRobertaCLIP,
return_transforms=False,
return_tokenizer=False,
tokenizer_padding='eos',
dtype=torch.float32,
device='cpu',
**kwargs):
# init a model on device
with torch.device(device):
model = model_cls(**kwargs)
# set device
model = model.to(dtype=dtype, device=device)
output = (model,)
# init transforms
if return_transforms:
# mean and std
if 'siglip' in pretrained_name.lower():
mean, std = [0.5, 0.5, 0.5], [0.5, 0.5, 0.5]
else:
mean = [0.48145466, 0.4578275, 0.40821073]
std = [0.26862954, 0.26130258, 0.27577711]
# transforms
transforms = T.Compose([
T.Resize((model.image_size, model.image_size),
interpolation=T.InterpolationMode.BICUBIC),
T.ToTensor(),
T.Normalize(mean=mean, std=std)
])
output += (transforms,)
return output[0] if len(output) == 1 else output
def clip_xlm_roberta_vit_h_14(
pretrained=False,
pretrained_name='open-clip-xlm-roberta-large-vit-huge-14',
**kwargs):
cfg = dict(
embed_dim=1024,
image_size=224,
patch_size=14,
vision_dim=1280,
vision_mlp_ratio=4,
vision_heads=16,
vision_layers=32,
vision_pool='token',
activation='gelu',
vocab_size=250002,
max_text_len=514,
type_size=1,
pad_id=1,
text_dim=1024,
text_heads=16,
text_layers=24,
text_post_norm=True,
text_dropout=0.1,
attn_dropout=0.0,
proj_dropout=0.0,
embedding_dropout=0.0)
cfg.update(**kwargs)
return _clip(pretrained, pretrained_name, XLMRobertaCLIP, **cfg)
class CLIPModel:
def __init__(self, dtype, device, checkpoint_path, tokenizer_path):
self.dtype = dtype
self.device = device
self.checkpoint_path = checkpoint_path
self.tokenizer_path = tokenizer_path
# init model
self.model, self.transforms = clip_xlm_roberta_vit_h_14(
pretrained=False,
return_transforms=True,
return_tokenizer=False,
dtype=dtype,
device=device)
self.model = self.model.eval().requires_grad_(False)
logging.info(f'loading {checkpoint_path}')
self.model.load_state_dict(
torch.load(checkpoint_path, map_location='cpu'))
# init tokenizer
self.tokenizer = HuggingfaceTokenizer(
name=tokenizer_path,
seq_len=self.model.max_text_len - 2,
clean='whitespace')
def visual(self, videos):
# preprocess
size = (self.model.image_size,) * 2
videos = torch.cat([
F.interpolate(
u.transpose(0, 1),
size=size,
mode='bicubic',
align_corners=False) for u in videos
])
videos = self.transforms.transforms[-1](videos.mul_(0.5).add_(0.5))
# forward
with torch.amp.autocast('cuda',dtype=self.dtype):
out = self.model.visual(videos, use_31_block=True)
return out
+56
View File
@@ -0,0 +1,56 @@
import os
import torch.nn as nn
import torch
# AudioMLP: audio mlp
class AudioMLP(nn.Module):
def __init__(self, input_dim=512, hidden_dim=512, output_dim=384):
super().__init__()
self.embed_to_whispers = nn.ModuleList([
self._build_mlp(input_dim=input_dim, hidden_dim=hidden_dim, output_dim=output_dim),
self._build_mlp(input_dim=input_dim, hidden_dim=hidden_dim, output_dim=output_dim),
self._build_mlp(input_dim=input_dim, hidden_dim=hidden_dim, output_dim=output_dim),
self._build_mlp(input_dim=input_dim, hidden_dim=hidden_dim, output_dim=output_dim),
self._build_mlp(input_dim=input_dim, hidden_dim=hidden_dim, output_dim=output_dim)
])
# _build_mlp: build mlp
def _build_mlp(self, input_dim, hidden_dim, output_dim):
return nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.BatchNorm1d(hidden_dim),
nn.Linear(hidden_dim, output_dim)
)
# forward: forward pass
def forward(self, x):
return [net(x) for net in self.embed_to_whispers] # 返回5个输出张量的列表
# AudioEmbedding: audio embedding
class AudioEmbedding(nn.Module):
def __init__(self, input_dim=384 * 5, hidden_dim=2048, output_dim=8194):
super().__init__()
self.whispers_to_token = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.BatchNorm1d(hidden_dim),
nn.Linear(hidden_dim, hidden_dim * 2),
nn.ReLU(),
nn.BatchNorm1d(hidden_dim * 2),
nn.Linear(hidden_dim * 2, output_dim)
)
self.codec_embed = nn.Embedding(8194, 512)
# forward: forward pass
def forward(self, x):
l=0
if len(x.shape)==3:
b,l,d=x.shape
x=x.view(-1,d)
audio_token = self.whispers_to_token(x) # [(B T), 8192]
indices = audio_token.argmax(dim=-1) # [(B T)]
audio_embedding = self.codec_embed(indices) # [(B T), embed_dim]
if l!=0:
audio_embedding=audio_embedding.view(b,l,-1)
return audio_embedding
File diff suppressed because it is too large Load Diff
+525
View File
@@ -0,0 +1,525 @@
# Modified from transformers.models.t5.modeling_t5
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
# Modified from transformers.models.t5.modeling_t5
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import logging
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from .tokenizers import HuggingfaceTokenizer
__all__ = [
'T5Model',
'T5Encoder',
'T5Decoder',
'T5EncoderModel',
]
# Modified from transformers.models.t5.modeling_t5
def fp16_clamp(x):
if x.dtype == torch.float16 and torch.isinf(x).any():
clamp = torch.finfo(x.dtype).max - 1000
x = torch.clamp(x, min=-clamp, max=clamp)
return x
# Modified from transformers.models.t5.modeling_t5
def init_weights(m):
if isinstance(m, T5LayerNorm):
nn.init.ones_(m.weight)
elif isinstance(m, T5Model):
nn.init.normal_(m.token_embedding.weight, std=1.0)
elif isinstance(m, T5FeedForward):
nn.init.normal_(m.gate[0].weight, std=m.dim**-0.5)
nn.init.normal_(m.fc1.weight, std=m.dim**-0.5)
nn.init.normal_(m.fc2.weight, std=m.dim_ffn**-0.5)
elif isinstance(m, T5Attention):
nn.init.normal_(m.q.weight, std=(m.dim * m.dim_attn)**-0.5)
nn.init.normal_(m.k.weight, std=m.dim**-0.5)
nn.init.normal_(m.v.weight, std=m.dim**-0.5)
nn.init.normal_(m.o.weight, std=(m.num_heads * m.dim_attn)**-0.5)
elif isinstance(m, T5RelativeEmbedding):
nn.init.normal_(
m.embedding.weight, std=(2 * m.num_buckets * m.num_heads)**-0.5)
# Modified from transformers.models.t5.modeling_t5
class GELU(nn.Module):
def forward(self, x):
return 0.5 * x * (1.0 + torch.tanh(
math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
# Modified from transformers.models.t5.modeling_t5
class T5LayerNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super(T5LayerNorm, self).__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) +
self.eps)
if self.weight.dtype in [torch.float16, torch.bfloat16]:
x = x.type_as(self.weight)
return self.weight * x
class T5Attention(nn.Module):
def __init__(self, dim, dim_attn, num_heads, dropout=0.1):
assert dim_attn % num_heads == 0
super(T5Attention, self).__init__()
self.dim = dim
self.dim_attn = dim_attn
self.num_heads = num_heads
self.head_dim = dim_attn // num_heads
# layers
self.q = nn.Linear(dim, dim_attn, bias=False)
self.k = nn.Linear(dim, dim_attn, bias=False)
self.v = nn.Linear(dim, dim_attn, bias=False)
self.o = nn.Linear(dim_attn, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x, context=None, mask=None, pos_bias=None):
"""
x: [B, L1, C].
context: [B, L2, C] or None.
mask: [B, L2] or [B, L1, L2] or None.
"""
# check inputs
context = x if context is None else context
b, n, c = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.q(x).view(b, -1, n, c)
k = self.k(context).view(b, -1, n, c)
v = self.v(context).view(b, -1, n, c)
# attention bias
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
if pos_bias is not None:
attn_bias += pos_bias
if mask is not None:
assert mask.ndim in [2, 3]
mask = mask.view(b, 1, 1,
-1) if mask.ndim == 2 else mask.unsqueeze(1)
attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
# compute attention (T5 does not use scaling)
attn = torch.einsum('binc,bjnc->bnij', q, k) + attn_bias
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
x = torch.einsum('bnij,bjnc->binc', attn, v)
# output
x = x.reshape(b, -1, n * c)
x = self.o(x)
x = self.dropout(x)
return x
class T5FeedForward(nn.Module):
def __init__(self, dim, dim_ffn, dropout=0.1):
super(T5FeedForward, self).__init__()
self.dim = dim
self.dim_ffn = dim_ffn
# layers
self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
x = self.fc1(x) * self.gate(x)
x = self.dropout(x)
x = self.fc2(x)
x = self.dropout(x)
return x
class T5SelfAttention(nn.Module):
def __init__(self,
dim,
dim_attn,
dim_ffn,
num_heads,
num_buckets,
shared_pos=True,
dropout=0.1):
super(T5SelfAttention, self).__init__()
self.dim = dim
self.dim_attn = dim_attn
self.dim_ffn = dim_ffn
self.num_heads = num_heads
self.num_buckets = num_buckets
self.shared_pos = shared_pos
# layers
self.norm1 = T5LayerNorm(dim)
self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
self.norm2 = T5LayerNorm(dim)
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
num_buckets, num_heads, bidirectional=True)
def forward(self, x, mask=None, pos_bias=None):
e = pos_bias if self.shared_pos else self.pos_embedding(
x.size(1), x.size(1))
x = fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
x = fp16_clamp(x + self.ffn(self.norm2(x)))
return x
class T5CrossAttention(nn.Module):
def __init__(self,
dim,
dim_attn,
dim_ffn,
num_heads,
num_buckets,
shared_pos=True,
dropout=0.1):
super(T5CrossAttention, self).__init__()
self.dim = dim
self.dim_attn = dim_attn
self.dim_ffn = dim_ffn
self.num_heads = num_heads
self.num_buckets = num_buckets
self.shared_pos = shared_pos
# layers
self.norm1 = T5LayerNorm(dim)
self.self_attn = T5Attention(dim, dim_attn, num_heads, dropout)
self.norm2 = T5LayerNorm(dim)
self.cross_attn = T5Attention(dim, dim_attn, num_heads, dropout)
self.norm3 = T5LayerNorm(dim)
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(
num_buckets, num_heads, bidirectional=False)
def forward(self,
x,
mask=None,
encoder_states=None,
encoder_mask=None,
pos_bias=None):
e = pos_bias if self.shared_pos else self.pos_embedding(
x.size(1), x.size(1))
x = fp16_clamp(x + self.self_attn(self.norm1(x), mask=mask, pos_bias=e))
x = fp16_clamp(x + self.cross_attn(
self.norm2(x), context=encoder_states, mask=encoder_mask))
x = fp16_clamp(x + self.ffn(self.norm3(x)))
return x
class T5RelativeEmbedding(nn.Module):
def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
super(T5RelativeEmbedding, self).__init__()
self.num_buckets = num_buckets
self.num_heads = num_heads
self.bidirectional = bidirectional
self.max_dist = max_dist
# layers
self.embedding = nn.Embedding(num_buckets, num_heads)
def forward(self, lq, lk):
device = self.embedding.weight.device
# rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
# torch.arange(lq).unsqueeze(1).to(device)# rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
# torch.arange(lq).unsqueeze(1).to(device)
rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
torch.arange(lq, device=device).unsqueeze(1)
rel_pos = self._relative_position_bucket(rel_pos)
rel_pos_embeds = self.embedding(rel_pos)
rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
0) # [1, N, Lq, Lk]
return rel_pos_embeds.contiguous()
def _relative_position_bucket(self, rel_pos):
# preprocess
if self.bidirectional:
num_buckets = self.num_buckets // 2
rel_buckets = (rel_pos > 0).long() * num_buckets
rel_pos = torch.abs(rel_pos)
else:
num_buckets = self.num_buckets
rel_buckets = 0
rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
# embeddings for small and large positions
# embeddings for small and large positions
max_exact = num_buckets // 2
rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
math.log(self.max_dist / max_exact) *
(num_buckets - max_exact)).long()
rel_pos_large = torch.min(
rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
return rel_buckets
class T5Encoder(nn.Module):
def __init__(self,
vocab,
dim,
dim_attn,
dim_ffn,
num_heads,
num_layers,
num_buckets,
shared_pos=True,
dropout=0.1):
super(T5Encoder, self).__init__()
self.dim = dim
self.dim_attn = dim_attn
self.dim_ffn = dim_ffn
self.num_heads = num_heads
self.num_layers = num_layers
self.num_buckets = num_buckets
self.shared_pos = shared_pos
# layers
self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
else nn.Embedding(vocab, dim)
self.pos_embedding = T5RelativeEmbedding(
num_buckets, num_heads, bidirectional=True) if shared_pos else None
self.dropout = nn.Dropout(dropout)
self.blocks = nn.ModuleList([
T5SelfAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
shared_pos, dropout) for _ in range(num_layers)
])
self.norm = T5LayerNorm(dim)
# initialize weights
self.apply(init_weights)
def forward(self, ids, mask=None):
x = self.token_embedding(ids)
x = self.dropout(x)
e = self.pos_embedding(x.size(1),
x.size(1)) if self.shared_pos else None
for block in self.blocks:
x = block(x, mask, pos_bias=e)
x = self.norm(x)
x = self.dropout(x)
return x
class T5Decoder(nn.Module):
def __init__(self,
vocab,
dim,
dim_attn,
dim_ffn,
num_heads,
num_layers,
num_buckets,
shared_pos=True,
dropout=0.1):
super(T5Decoder, self).__init__()
self.dim = dim
self.dim_attn = dim_attn
self.dim_ffn = dim_ffn
self.num_heads = num_heads
self.num_layers = num_layers
self.num_buckets = num_buckets
self.shared_pos = shared_pos
# layers
self.token_embedding = vocab if isinstance(vocab, nn.Embedding) \
else nn.Embedding(vocab, dim)
self.pos_embedding = T5RelativeEmbedding(
num_buckets, num_heads, bidirectional=False) if shared_pos else None
self.dropout = nn.Dropout(dropout)
self.blocks = nn.ModuleList([
T5CrossAttention(dim, dim_attn, dim_ffn, num_heads, num_buckets,
shared_pos, dropout) for _ in range(num_layers)
])
self.norm = T5LayerNorm(dim)
# initialize weights
self.apply(init_weights)
def forward(self, ids, mask=None, encoder_states=None, encoder_mask=None):
b, s = ids.size()
# causal mask
if mask is None:
mask = torch.tril(torch.ones(1, s, s).to(ids.device))
elif mask.ndim == 2:
mask = torch.tril(mask.unsqueeze(1).expand(-1, s, -1))
# layers
x = self.token_embedding(ids)
x = self.dropout(x)
e = self.pos_embedding(x.size(1),
x.size(1)) if self.shared_pos else None
for block in self.blocks:
x = block(x, mask, encoder_states, encoder_mask, pos_bias=e)
x = self.norm(x)
x = self.dropout(x)
return x
class T5Model(nn.Module):
def __init__(self,
vocab_size,
dim,
dim_attn,
dim_ffn,
num_heads,
encoder_layers,
decoder_layers,
num_buckets,
shared_pos=True,
dropout=0.1):
super(T5Model, self).__init__()
self.vocab_size = vocab_size
self.dim = dim
self.dim_attn = dim_attn
self.dim_ffn = dim_ffn
self.num_heads = num_heads
self.encoder_layers = encoder_layers
self.decoder_layers = decoder_layers
self.num_buckets = num_buckets
# layers
self.token_embedding = nn.Embedding(vocab_size, dim)
self.encoder = T5Encoder(self.token_embedding, dim, dim_attn, dim_ffn,
num_heads, encoder_layers, num_buckets,
shared_pos, dropout)
self.decoder = T5Decoder(self.token_embedding, dim, dim_attn, dim_ffn,
num_heads, decoder_layers, num_buckets,
shared_pos, dropout)
self.head = nn.Linear(dim, vocab_size, bias=False)
# initialize weights
self.apply(init_weights)
def forward(self, encoder_ids, encoder_mask, decoder_ids, decoder_mask):
x = self.encoder(encoder_ids, encoder_mask)
x = self.decoder(decoder_ids, decoder_mask, x, encoder_mask)
x = self.head(x)
return x
def _t5(name,
encoder_only=False,
decoder_only=False,
return_tokenizer=False,
tokenizer_kwargs={},
dtype=torch.float32,
device='cpu',
**kwargs):
# sanity check
assert not (encoder_only and decoder_only)
# params
if encoder_only:
model_cls = T5Encoder
kwargs['vocab'] = kwargs.pop('vocab_size')
kwargs['num_layers'] = kwargs.pop('encoder_layers')
_ = kwargs.pop('decoder_layers')
elif decoder_only:
model_cls = T5Decoder
kwargs['vocab'] = kwargs.pop('vocab_size')
kwargs['num_layers'] = kwargs.pop('decoder_layers')
_ = kwargs.pop('encoder_layers')
else:
model_cls = T5Model
# init model
with torch.device(device):
model = model_cls(**kwargs)
# set device
model = model.to(dtype=dtype, device=device)
# init tokenizer
if return_tokenizer:
from .tokenizers import HuggingfaceTokenizer
tokenizer = HuggingfaceTokenizer(f'google/{name}', **tokenizer_kwargs)
return model, tokenizer
else:
return model
def umt5_xxl(**kwargs):
cfg = dict(
vocab_size=256384,
dim=4096,
dim_attn=4096,
dim_ffn=10240,
num_heads=64,
encoder_layers=24,
decoder_layers=24,
num_buckets=32,
shared_pos=False,
dropout=0.1)
cfg.update(**kwargs)
return _t5('umt5-xxl', **cfg)
class T5EncoderModel:
def __init__(
self,
text_len,
dtype=torch.bfloat16,
device=torch.cuda.current_device(),
checkpoint_path=None,
tokenizer_path=None,
shard_fn=None,
):
self.text_len = text_len
self.dtype = dtype
self.device = device
self.checkpoint_path = checkpoint_path
self.tokenizer_path = tokenizer_path
# init model
model = umt5_xxl(
encoder_only=True,
return_tokenizer=False,
dtype=dtype,
device=device).eval().requires_grad_(False)
logging.info(f'loading {checkpoint_path}')
model.load_state_dict(torch.load(checkpoint_path, map_location='cpu'))
self.model = model
if shard_fn is not None:
self.model = shard_fn(self.model, sync_module_states=False)
else:
self.model.to(self.device)
# init tokenizer
# init tokenizer
self.tokenizer = HuggingfaceTokenizer(
name=tokenizer_path, seq_len=text_len, clean='whitespace')
def __call__(self, texts, device):
ids, mask = self.tokenizer(
texts, return_mask=True, add_special_tokens=True)
ids = ids.to(device)
mask = mask.to(device)
seq_lens = mask.gt(0).sum(dim=1).long()
context = self.model(ids, mask)
return [u[:v] for u, v in zip(context, seq_lens)]
def text_embedding(self, ids, mask, device):
ids = ids.to(device)
mask = mask.to(device)
seq_lens = mask.gt(0).sum(dim=1).long()
context = self.model(ids, mask)
return [u[:v] for u, v in zip(context, seq_lens)]
+82
View File
@@ -0,0 +1,82 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import html
import string
import ftfy
import regex as re
from transformers import AutoTokenizer
__all__ = ['HuggingfaceTokenizer']
# Modified from transformers.models.t5.tokenization_t5_fast
def basic_clean(text):
text = ftfy.fix_text(text)
text = html.unescape(html.unescape(text))
return text.strip()
def whitespace_clean(text):
text = re.sub(r'\s+', ' ', text)
text = text.strip()
return text
# Modified from transformers.models.t5.tokenization_t5_fast
def canonicalize(text, keep_punctuation_exact_string=None):
text = text.replace('_', ' ')
if keep_punctuation_exact_string:
text = keep_punctuation_exact_string.join(
part.translate(str.maketrans('', '', string.punctuation))
for part in text.split(keep_punctuation_exact_string))
else:
text = text.translate(str.maketrans('', '', string.punctuation))
text = text.lower()
text = re.sub(r'\s+', ' ', text)
return text.strip()
# Modified from transformers.models.t5.tokenization_t5_fast
class HuggingfaceTokenizer:
def __init__(self, name, seq_len=None, clean=None, **kwargs):
assert clean in (None, 'whitespace', 'lower', 'canonicalize')
self.name = name
self.seq_len = seq_len
self.clean = clean
# init tokenizer
self.tokenizer = AutoTokenizer.from_pretrained(name, **kwargs)
self.vocab_size = self.tokenizer.vocab_size
def __call__(self, sequence, **kwargs):
return_mask = kwargs.pop('return_mask', False)
# arguments
_kwargs = {'return_tensors': 'pt'}
if self.seq_len is not None:
_kwargs.update({
'padding': 'max_length',
'truncation': True,
'max_length': self.seq_len
})
_kwargs.update(**kwargs)
# tokenization
if isinstance(sequence, str):
sequence = [sequence]
if self.clean:
sequence = [self._clean(u) for u in sequence]
ids = self.tokenizer(sequence, **_kwargs)
# output
if return_mask:
return ids.input_ids, ids.attention_mask
else:
return ids.input_ids
# Modified from transformers.models.t5.tokenization_t5_fast
def _clean(self, text):
if self.clean == 'whitespace':
text = whitespace_clean(basic_clean(text))
elif self.clean == 'lower':
text = whitespace_clean(basic_clean(text)).lower()
elif self.clean == 'canonicalize':
text = canonicalize(basic_clean(text))
return text
File diff suppressed because it is too large Load Diff
+170
View File
@@ -0,0 +1,170 @@
# Modified from transformers.models.xlm_roberta.modeling_xlm_roberta
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
import torch.nn as nn
import torch.nn.functional as F
__all__ = ['XLMRoberta', 'xlm_roberta_large']
class SelfAttention(nn.Module):
def __init__(self, dim, num_heads, dropout=0.1, eps=1e-5):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
# layers
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask):
"""
x: [B, L, C].
"""
b, s, c, n, d = *x.size(), self.num_heads, self.head_dim
# compute query, key, value
q = self.q(x).reshape(b, s, n, d).permute(0, 2, 1, 3)
k = self.k(x).reshape(b, s, n, d).permute(0, 2, 1, 3)
v = self.v(x).reshape(b, s, n, d).permute(0, 2, 1, 3)
# compute attention
p = self.dropout.p if self.training else 0.0
x = F.scaled_dot_product_attention(q, k, v, mask, p)
x = x.permute(0, 2, 1, 3).reshape(b, s, c)
# output
x = self.o(x)
x = self.dropout(x)
return x
class AttentionBlock(nn.Module):
def __init__(self, dim, num_heads, post_norm, dropout=0.1, eps=1e-5):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.post_norm = post_norm
self.eps = eps
# layers
self.attn = SelfAttention(dim, num_heads, dropout, eps)
self.norm1 = nn.LayerNorm(dim, eps=eps)
self.ffn = nn.Sequential(
nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim),
nn.Dropout(dropout))
self.norm2 = nn.LayerNorm(dim, eps=eps)
def forward(self, x, mask):
if self.post_norm:
x = self.norm1(x + self.attn(x, mask))
x = self.norm2(x + self.ffn(x))
else:
x = x + self.attn(self.norm1(x), mask)
x = x + self.ffn(self.norm2(x))
return x
class XLMRoberta(nn.Module):
"""
XLMRobertaModel with no pooler and no LM head.
"""
def __init__(self,
vocab_size=250002,
max_seq_len=514,
type_size=1,
pad_id=1,
dim=1024,
num_heads=16,
num_layers=24,
post_norm=True,
dropout=0.1,
eps=1e-5):
super().__init__()
self.vocab_size = vocab_size
self.max_seq_len = max_seq_len
self.type_size = type_size
self.pad_id = pad_id
self.dim = dim
self.num_heads = num_heads
self.num_layers = num_layers
self.post_norm = post_norm
self.eps = eps
# embeddings
self.token_embedding = nn.Embedding(vocab_size, dim, padding_idx=pad_id)
self.type_embedding = nn.Embedding(type_size, dim)
self.pos_embedding = nn.Embedding(max_seq_len, dim, padding_idx=pad_id)
self.dropout = nn.Dropout(dropout)
# blocks
self.blocks = nn.ModuleList([
AttentionBlock(dim, num_heads, post_norm, dropout, eps)
for _ in range(num_layers)
])
# norm layer
self.norm = nn.LayerNorm(dim, eps=eps)
def forward(self, ids):
"""
ids: [B, L] of torch.LongTensor.
"""
b, s = ids.shape
mask = ids.ne(self.pad_id).long()
# embeddings
x = self.token_embedding(ids) + \
self.type_embedding(torch.zeros_like(ids)) + \
self.pos_embedding(self.pad_id + torch.cumsum(mask, dim=1) * mask)
if self.post_norm:
x = self.norm(x)
x = self.dropout(x)
# blocks
mask = torch.where(
mask.view(b, 1, 1, s).gt(0), 0.0,
torch.finfo(x.dtype).min)
for block in self.blocks:
x = block(x, mask)
# output
if not self.post_norm:
x = self.norm(x)
return x
def xlm_roberta_large(pretrained=False,
return_tokenizer=False,
device='cpu',
**kwargs):
"""
XLMRobertaLarge adapted from Huggingface.
"""
# params
cfg = dict(
vocab_size=250002,
max_seq_len=514,
type_size=1,
pad_id=1,
dim=1024,
num_heads=16,
num_layers=24,
post_norm=True,
dropout=0.1,
eps=1e-5)
cfg.update(**kwargs)
# init a model on device
with torch.device(device):
model = XLMRoberta(**cfg)
return model
File diff suppressed because it is too large Load Diff
+7
View File
@@ -0,0 +1,7 @@
# import modules
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
__all__ = [
'HuggingfaceTokenizer', 'get_sampling_sigmas', 'retrieve_timesteps',
'FlowDPMSolverMultistepScheduler', 'FlowUniPCMultistepScheduler'
]
+780
View File
@@ -0,0 +1,780 @@
# Convert unipc for flow matching
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import (KarrasDiffusionSchedulers,
SchedulerMixin,
SchedulerOutput)
from diffusers.utils import deprecate, is_scipy_available
if is_scipy_available():
import scipy.stats
class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
"""
`UniPCMultistepScheduler` is a training-free framework designed for the fast sampling of diffusion models.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
solver_order (`int`, default `2`):
The UniPC order which can be any positive integer. The effective order of accuracy is `solver_order + 1`
due to the UniC. It is recommended to use `solver_order=2` for guided sampling, and `solver_order=3` for
unconditional sampling.
prediction_type (`str`, defaults to "flow_prediction"):
Prediction type of the scheduler function; must be `flow_prediction` for this scheduler, which predicts
the flow of the diffusion process.
thresholding (`bool`, defaults to `False`):
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
as Stable Diffusion.
dynamic_thresholding_ratio (`float`, defaults to 0.995):
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
sample_max_value (`float`, defaults to 1.0):
The threshold value for dynamic thresholding. Valid only when `thresholding=True` and `predict_x0=True`.
predict_x0 (`bool`, defaults to `True`):
Whether to use the updating algorithm on the predicted x0.
solver_type (`str`, default `bh2`):
Solver type for UniPC. It is recommended to use `bh1` for unconditional sampling when steps < 10, and `bh2`
otherwise.
lower_order_final (`bool`, default `True`):
Whether to use lower-order solvers in the final steps. Only valid for < 15 inference steps. This can
stabilize the sampling of DPMSolver for steps < 15, especially for steps <= 10.
disable_corrector (`list`, default `[]`):
Decides which step to disable the corrector to mitigate the misalignment between `epsilon_theta(x_t, c)`
and `epsilon_theta(x_t^c, c)` which can influence convergence for a large guidance scale. Corrector is
usually disabled during the first few steps.
solver_p (`SchedulerMixin`, default `None`):
Any other scheduler that if specified, the algorithm becomes `solver_p + UniC`.
use_karras_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`,
the sigmas are determined according to a sequence of noise levels {σi}.
use_exponential_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use exponential sigmas for step sizes in the noise schedule during the sampling process.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
steps_offset (`int`, defaults to 0):
An offset added to the inference steps, as required by some model families.
final_sigmas_type (`str`, defaults to `"zero"`):
The final `sigma` value for the noise schedule during the sampling process. If `"sigma_min"`, the final
sigma is the same as the last sigma in the training schedule. If `zero`, the final sigma is set to 0.
"""
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
solver_order: int = 2,
prediction_type: str = "flow_prediction",
shift: Optional[float] = 1.0,
use_dynamic_shifting=False,
thresholding: bool = False,
dynamic_thresholding_ratio: float = 0.995,
sample_max_value: float = 1.0,
predict_x0: bool = True,
solver_type: str = "bh2",
lower_order_final: bool = True,
disable_corrector: List[int] = [],
solver_p: SchedulerMixin = None,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
):
if solver_type not in ["bh1", "bh2"]:
if solver_type in ["midpoint", "heun", "logrho"]:
self.register_to_config(solver_type="bh2")
else:
raise NotImplementedError(
f"{solver_type} is not implemented for {self.__class__}")
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps = None
alphas = np.linspace(1, 1 / num_train_timesteps,
num_train_timesteps)[::-1].copy()
sigmas = 1.0 - alphas
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
self.sigmas = sigmas
self.timesteps = sigmas * num_train_timesteps
self.model_outputs = [None] * solver_order
self.timestep_list = [None] * solver_order
self.lower_order_nums = 0
self.disable_corrector = disable_corrector
self.solver_p = solver_p
self.last_sample = None
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def set_timesteps(
self,
num_inference_steps: Union[int, None] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[Union[float, None]] = None,
shift: Optional[Union[float, None]] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
Total number of the spacing of the time steps.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError(
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is None:
sigmas = np.linspace(self.sigma_max, self.sigma_min,
num_inference_steps +
1).copy()[:-1] # pyright: ignore
if self.config.use_dynamic_shifting:
sigmas = self.time_shift(mu, 1.0, sigmas) # pyright: ignore
else:
if shift is None:
shift = self.config.shift
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
if self.config.final_sigmas_type == "sigma_min":
sigma_last = ((1 - self.alphas_cumprod[0]) /
self.alphas_cumprod[0])**0.5
elif self.config.final_sigmas_type == "zero":
sigma_last = 0
else:
raise ValueError(
f"`final_sigmas_type` must be one of 'zero', or 'sigma_min', but got {self.config.final_sigmas_type}"
)
timesteps = sigmas * self.config.num_train_timesteps
sigmas = np.concatenate([sigmas, [sigma_last]
]).astype(np.float32) # pyright: ignore
self.sigmas = torch.from_numpy(sigmas)
self.timesteps = torch.from_numpy(timesteps).to(
device=device, dtype=torch.int64)
self.num_inference_steps = len(timesteps)
self.model_outputs = [
None,
] * self.config.solver_order
self.lower_order_nums = 0
self.last_sample = None
if self.solver_p:
self.solver_p.set_timesteps(self.num_inference_steps, device=device)
# add an index counter for schedulers that allow duplicated timesteps
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
"""
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
photorealism as well as better image-text alignment, especially when using very large guidance weights."
https://arxiv.org/abs/2205.11487
"""
dtype = sample.dtype
batch_size, channels, *remaining_dims = sample.shape
if dtype not in (torch.float32, torch.float64):
sample = sample.float(
) # upcast for quantile calculation, and clamp not implemented for cpu half
# Flatten sample for doing quantile calculation along each image
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
s = torch.quantile(
abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
s = torch.clamp(
s, min=1, max=self.config.sample_max_value
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
s = s.unsqueeze(
1) # (batch_size, 1) because clamp will broadcast along dim=0
sample = torch.clamp(
sample, -s, s
) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
sample = sample.reshape(batch_size, channels, *remaining_dims)
sample = sample.to(dtype)
return sample
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._sigma_to_t
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def _sigma_to_alpha_sigma_t(self, sigma):
return 1 - sigma, sigma
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.set_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
def convert_model_output(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
**kwargs,
) -> torch.Tensor:
r"""
Convert the model output to the corresponding type the UniPC algorithm needs.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
Returns:
`torch.Tensor`:
The converted model output.
"""
timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
"missing `sample` as a required keyward argument")
sigma = self.sigmas[self.step_index]
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
if self.predict_x0:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
x0_pred = self._threshold_sample(x0_pred)
return x0_pred
else:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
epsilon = sample - (1 - sigma_t) * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
x0_pred = self._threshold_sample(x0_pred)
epsilon = model_output + x0_pred
return epsilon
def multistep_uni_p_bh_update(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: int = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model at the current timestep.
prev_timestep (`int`):
The previous discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
order (`int`):
The order of UniP at this timestep (corresponds to the *p* in UniPC-p).
Returns:
`torch.Tensor`:
The sample tensor at the previous timestep.
"""
prev_timestep = args[0] if len(args) > 0 else kwargs.pop(
"prev_timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
" missing `sample` as a required keyward argument")
if order is None:
if len(args) > 2:
order = args[2]
else:
raise ValueError(
" missing `order` as a required keyward argument")
model_output_list = self.model_outputs
s0 = self.timestep_list[-1]
m0 = model_output_list[-1]
x = sample
if self.solver_p:
x_t = self.solver_p.step(model_output, s0, x).prev_sample
return x_t
sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[
self.step_index] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = sample.device
rks = []
D1s = []
for i in range(1, order):
si = self.step_index - i # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if len(D1s) > 0:
D1s = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version
if order == 2:
rhos_p = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
rhos_p = torch.linalg.solve(R[:-1, :-1],
b[:-1]).to(device).to(x.dtype)
else:
D1s = None
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - alpha_t * B_h * pred_res
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - sigma_t * B_h * pred_res
x_t = x_t.to(x.dtype)
return x_t
def multistep_uni_c_bh_update(
self,
this_model_output: torch.Tensor,
*args,
last_sample: torch.Tensor = None,
this_sample: torch.Tensor = None,
order: int = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniC (B(h) version).
Args:
this_model_output (`torch.Tensor`):
The model outputs at `x_t`.
this_timestep (`int`):
The current timestep `t`.
last_sample (`torch.Tensor`):
The generated sample before the last predictor `x_{t-1}`.
this_sample (`torch.Tensor`):
The generated sample after the last predictor `x_{t}`.
order (`int`):
The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`.
Returns:
`torch.Tensor`:
The corrected sample tensor at the current timestep.
"""
this_timestep = args[0] if len(args) > 0 else kwargs.pop(
"this_timestep", None)
if last_sample is None:
if len(args) > 1:
last_sample = args[1]
else:
raise ValueError(
" missing`last_sample` as a required keyward argument")
if this_sample is None:
if len(args) > 2:
this_sample = args[2]
else:
raise ValueError(
" missing`this_sample` as a required keyward argument")
if order is None:
if len(args) > 3:
order = args[3]
else:
raise ValueError(
" missing`order` as a required keyward argument")
model_output_list = self.model_outputs
m0 = model_output_list[-1]
x = last_sample
x_t = this_sample
model_t = this_model_output
sigma_t, sigma_s0 = self.sigmas[self.step_index], self.sigmas[
self.step_index - 1] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = this_sample.device
rks = []
D1s = []
for i in range(1, order):
si = self.step_index - (i + 1) # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if len(D1s) > 0:
D1s = torch.stack(D1s, dim=1)
else:
D1s = None
# for order 1, we use a simplified version
if order == 1:
rhos_c = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype)
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - alpha_t * B_h * (corr_res + rhos_c[-1] * D1_t)
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - sigma_t * B_h * (corr_res + rhos_c[-1] * D1_t)
x_t = x_t.to(x.dtype)
return x_t
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._init_step_index
def _init_step_index(self, timestep):
"""
Initialize the step_index counter for the scheduler.
"""
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def step(self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
generator=None) -> Union[SchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
Args:
model_output (`torch.Tensor`):
The direct output from learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`.
Returns:
[`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a
tuple is returned where the first element is the sample tensor.
"""
if self.num_inference_steps is None:
raise ValueError(
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
)
if self.step_index is None:
self._init_step_index(timestep)
use_corrector = (
self.step_index > 0 and
self.step_index - 1 not in self.disable_corrector and
self.last_sample is not None # pyright: ignore
)
model_output_convert = self.convert_model_output(
model_output, sample=sample)
if use_corrector:
sample = self.multistep_uni_c_bh_update(
this_model_output=model_output_convert,
last_sample=self.last_sample,
this_sample=sample,
order=self.this_order,
)
for i in range(self.config.solver_order - 1):
self.model_outputs[i] = self.model_outputs[i + 1]
self.timestep_list[i] = self.timestep_list[i + 1]
self.model_outputs[-1] = model_output_convert
self.timestep_list[-1] = timestep # pyright: ignore
if self.config.lower_order_final:
this_order = min(self.config.solver_order,
len(self.timesteps) -
self.step_index) # pyright: ignore
else:
this_order = self.config.solver_order
self.this_order = min(this_order,
self.lower_order_nums + 1) # warmup for multistep
assert self.this_order > 0
self.last_sample = sample
prev_sample = self.multistep_uni_p_bh_update(
model_output=model_output, # pass the original non-converted model output, in case solver-p is used
sample=sample,
order=self.this_order,
)
if self.lower_order_nums < self.config.solver_order:
self.lower_order_nums += 1
# upon completion increase step index by one
self._step_index += 1 # pyright: ignore
if not return_dict:
return (prev_sample,)
return SchedulerOutput(prev_sample=prev_sample)
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
"""
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
current timestep.
Args:
sample (`torch.Tensor`):
The input sample.
Returns:
`torch.Tensor`:
A scaled input sample.
"""
return sample
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(
device=original_samples.device, dtype=original_samples.dtype)
if original_samples.device.type == "mps" and torch.is_floating_point(
timesteps):
# mps does not support float64
schedule_timesteps = self.timesteps.to(
original_samples.device, dtype=torch.float32)
timesteps = timesteps.to(
original_samples.device, dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(original_samples.device)
timesteps = timesteps.to(original_samples.device)
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps)
for t in timesteps
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timesteps.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timesteps.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(original_samples.shape):
sigma = sigma.unsqueeze(-1)
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
noisy_samples = alpha_t * original_samples + sigma_t * noise
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
+572
View File
@@ -0,0 +1,572 @@
import os
from einops import rearrange
import torch
import torch.nn as nn
from xfuser.core.distributed import (
get_sequence_parallel_rank,
get_sequence_parallel_world_size,
get_sp_group,
)
from einops import rearrange, repeat
from functools import lru_cache
import imageio
import uuid
from tqdm import tqdm
import numpy as np
import subprocess
import soundfile as sf
VID_EXTENSIONS = (".mp4", ".avi", ".mov", ".mkv")
ASPECT_RATIO_627 = {
'0.26': ([320, 1216], 1), '0.38': ([384, 1024], 1), '0.50': ([448, 896], 1), '0.67': ([512, 768], 1),
'0.82': ([576, 704], 1), '1.00': ([640, 640], 1), '1.22': ([704, 576], 1), '1.50': ([768, 512], 1),
'1.86': ([832, 448], 1), '2.00': ([896, 448], 1), '2.50': ([960, 384], 1), '2.83': ([1088, 384], 1),
'3.60': ([1152, 320], 1), '3.80': ([1216, 320], 1), '4.00': ([1280, 320], 1)}
ASPECT_RATIO_960 = {
'0.22': ([448, 2048], 1), '0.29': ([512, 1792], 1), '0.36': ([576, 1600], 1), '0.45': ([640, 1408], 1),
'0.55': ([704, 1280], 1), '0.63': ([768, 1216], 1), '0.76': ([832, 1088], 1), '0.88': ([896, 1024], 1),
'1.00': ([960, 960], 1), '1.14': ([1024, 896], 1), '1.31': ([1088, 832], 1), '1.50': ([1152, 768], 1),
'1.58': ([1216, 768], 1), '1.82': ([1280, 704], 1), '1.91': ([1344, 704], 1), '2.20': ([1408, 640], 1),
'2.30': ([1472, 640], 1), '2.67': ([1536, 576], 1), '2.89': ([1664, 576], 1), '3.62': ([1856, 512], 1),
'3.75': ([1920, 512], 1)}
def torch_gc():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# 确保所有必要的库都已导入
import torch
import numpy as np
import imageio
from PIL import Image, ImageDraw, ImageFont
import textwrap
from tqdm import tqdm
import subprocess
import os
# 1. 确保所有必要的库都已导入
import torch
import numpy as np
import imageio
from PIL import Image, ImageDraw, ImageFont
import textwrap
from tqdm import tqdm
import subprocess
import os
def merge_audio_and_video(video_path, audio_path, output_path):
"""
使用 FFmpeg 将指定的音频和视频合并。
参数:
video_path (str): 原始视频文件路径。
audio_path (str): 要合并的音频文件路径。
output_path (str): 合成后视频的保存路径。
"""
# 构建 FFmpeg 命令
# -i video.mp4: 输入视频文件
# -i audio.wav: 输入音频文件
# -c:v copy: 直接复制视频流,不重新编码,速度快且无损画质
# -c:a aac: 将音频编码为 AAC 格式,这是 mp4 容器的常用格式
# -map 0:v:0: 映射第一个输入文件(视频)的视频流
# -map 1:a:0: 映射第二个输入文件(音频)的音频流
# -shortest: 当最短的输入流结束时,完成编码。这确保了如果音频比视频长,音频会被截断以匹配视频长度。
# -y: 如果输出文件已存在,则自动覆盖
new_path = '/apdcephfs_cq10/share_1367250/raylanzhang/ffmpeg2/ffmpeg/ffmpeg-7.0.2-i686-static'
# 获取当前的 PATH 环境变量
# 使用 os.environ.get('PATH', '') 来避免在 PATH 不存在时出错
current_path = os.environ.get('PATH', '')
# 构建新的 PATH,将新路径添加到最前面
# os.pathsep 是系统特定的路径分隔符 (在 Linux 和 macOS 上是 ':',在 Windows 上是 ';')
new_path_value = new_path + os.pathsep + current_path
# 设置新的 PATH 环境变量
os.environ['PATH'] = new_path_value
command = [
'ffmpeg',
'-i', video_path,
'-i', audio_path,
'-c:v', 'copy',
'-c:a', 'aac',
'-map', '0:v:0',
'-map', '1:a:0',
'-shortest', # <--- 关键参数
'-y',
output_path
]
try:
print(f"正在处理: {os.path.basename(video_path)} 和 {audio_path}")
# 执行命令,并隐藏 FFmpeg 的输出信息
subprocess.run(command, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
print(f"成功 -> {os.path.basename(output_path)}")
except subprocess.CalledProcessError as e:
print(f"处理失败: {os.path.basename(video_path)}")
# 打印错误信息以便调试
print("FFmpeg 错误信息:", e.stderr.decode())
except FileNotFoundError:
print("错误:找不到 'ffmpeg' 命令。请确保 FFmpeg 已正确安装并已添加到系统 PATH 环境变量中。")
return
def save_composite_video_with_audio(
main_video_tensor,
save_path,
motion_video_tensor=None,
vocal_audio_path=None,
text=None,
fps=25,
quality=7, # 默认质量稍作提高
font_path=None,
bg_color='black',
text_color='white'
):
"""
将一或两个视频与可选的文本和音频合并,并保存为一个最终的 MP4 文件。
该函数整合了视觉合成和音频添加的功能,并开启了调试模式:
1. 它首先将主视频、可选的第二个视频和可选的文本面板横向拼接,
生成一个无声的视觉合成视频。
2. 如果提供了音频文件,它会使用 FFmpeg 将音频剪辑到与视频等长,
然后将其混入视频中。FFmpeg 的所有输出都会被打印,以便调试。
3. 最后清理所有临时文件,只留下最终的带音频的视频。
参数:
main_video_tensor (torch.Tensor): 形状为 (C, T, H, W) 的主视频张量,像素值范围 [-1, 1]。
save_path (str): 最终视频的保存路径(.mp4 后缀会自动处理)。
motion_video_tensor (torch.Tensor, optional): 形状与主视频相同的第二个视频张量。
vocal_audio_path (str or list, optional): 输入的音频文件路径。如果为 None,则只输出无声视频。
text (str, optional): 要显示在视频右侧的文字。
fps (int): 视频的帧率。
quality (int): 视频质量,数值越高质量越好。
font_path (str, optional): .ttf 或 .otf 字体文件的路径。
bg_color (str): 文字区域的背景颜色。
text_color (str): 文字的颜色。
"""
# ------------------ 设置文件路径 ------------------
base_save_path = os.path.splitext(save_path)[0]
final_video_path = base_save_path + ".mp4"
temp_video_path = base_save_path + "_temp_video.mp4"
temp_audio_path = base_save_path + "_temp_audio.wav"
# =========================================================================
# PART 1: 视觉合成 (生成无声视频)
# =========================================================================
print("步骤 1/3: 正在生成视觉合成视频...")
# 内部辅助函数
def _save_video_from_frames(frames, path, fps, quality):
# 使用 imageio V3 API
writer = imageio.get_writer(path, fps=fps, quality=quality, macro_block_size=1)
for frame in tqdm(frames, desc=f"正在保存临时视频到 {path}"):
writer.append_data(np.array(frame))
writer.close()
def _process_video_tensor(video_tensor):
video_frames = (video_tensor + 1) / 2
video_frames = video_frames.permute(1, 2, 3, 0).cpu().numpy()
return np.clip(video_frames * 255, 0, 255).astype(np.uint8)
# 1.1 转换视频张量
video_frames_uint8 = _process_video_tensor(main_video_tensor)
_, T, H, W = main_video_tensor.shape
motion_frames_uint8 = None
if motion_video_tensor is not None:
motion_frames_uint8 = _process_video_tensor(motion_video_tensor)
# 1.2 创建文本面板
text_panel = None
if text:
margin = int(W * 0.1)
max_text_width, max_text_height = W - margin, H - margin
font_size = max(15, int(H / 15))
font = None
wrapped_text = text
while font_size > 8:
try:
font = ImageFont.truetype(font_path, font_size)
except (IOError, TypeError):
if font is None: font = ImageFont.load_default()
avg_char_width = font_size * 0.6
wrap_width = max(10, int(W / avg_char_width))
wrapped_text = textwrap.fill(text, width=wrap_width)
temp_draw = ImageDraw.Draw(Image.new('RGB', (W, H)))
bbox = temp_draw.multiline_textbbox((0, 0), wrapped_text, font=font)
text_width, text_height = bbox[2] - bbox[0], bbox[3] - bbox[1]
if text_width < max_text_width and text_height < max_text_height: break
font_size -= 1
text_panel = Image.new('RGB', (W, H), color=bg_color)
draw = ImageDraw.Draw(text_panel)
text_x, text_y = (W - text_width) / 2, (H - text_height) / 2
draw.multiline_text((text_x, text_y), wrapped_text, font=font, fill=text_color, align='center')
# 1.3 合成每一帧
num_columns = 1 + (motion_frames_uint8 is not None) + (text_panel is not None)
final_width = W * num_columns
composite_frames = []
for i in range(len(video_frames_uint8)):
composite_frame = Image.new('RGB', (final_width, H))
current_x = 0
composite_frame.paste(Image.fromarray(video_frames_uint8[i]), (current_x, 0))
current_x += W
if motion_frames_uint8 is not None:
composite_frame.paste(Image.fromarray(motion_frames_uint8[i]), (current_x, 0))
current_x += W
if text_panel is not None:
composite_frame.paste(text_panel, (current_x, 0))
composite_frames.append(composite_frame)
# 1.4 保存无声视频
_save_video_from_frames(composite_frames, temp_video_path, fps, quality)
print("无声视频已成功生成。")
# =========================================================================
# PART 2: 添加音频 (使用 FFmpeg)
# =========================================================================
if vocal_audio_path is None:
print("未提供音频文件,将无声视频重命名为最终文件。")
os.rename(temp_video_path, final_video_path)
print(f"最终视频已保存到: {final_video_path}")
return
print("步骤 2/3: 正在处理并添加音频...")
# 2.1 智能处理音频路径输入 (str 或 list)
actual_audio_file = None
if isinstance(vocal_audio_path, list):
if len(vocal_audio_path) > 0:
actual_audio_file = vocal_audio_path[0]
if len(vocal_audio_path) > 1:
print(f"警告: 提供了 {len(vocal_audio_path)} 个音频文件,只会使用第一个: {actual_audio_file}")
elif isinstance(vocal_audio_path, str):
actual_audio_file = vocal_audio_path
if not actual_audio_file:
print("警告: 提供的音频路径为空或格式不正确,将生成无声视频。")
os.rename(temp_video_path, final_video_path)
return
# 2.2 剪辑音频
duration = T / fps
try:
crop_command = [
"ffmpeg", "-y",
"-i", actual_audio_file,
"-t", f'{duration}',
"-acodec", "pcm_s16le", # 使用标准的WAV编码,兼容性好
temp_audio_path,
]
print("\n[调试信息] 将要执行音频剪辑命令:")
print(" ".join(crop_command))
# 执行命令并显示所有输出
subprocess.run(crop_command, check=True)
print(f"音频已成功剪辑到 {duration:.2f} 秒。")
except FileNotFoundError:
print("\n错误: 'ffmpeg' 命令未找到。请确保 FFmpeg 已正确安装并已添加到系统的 PATH 环境变量中。")
if os.path.exists(temp_video_path): os.remove(temp_video_path)
return
except subprocess.CalledProcessError as e:
print(f"\n错误:剪辑音频失败。FFmpeg 返回了非零退出码。请检查上面的 FFmpeg 输出信息以了解详情。")
print(f"Python 错误详情: {e}")
if os.path.exists(temp_video_path): os.remove(temp_video_path)
return
# 2.3 合并视频和音频
try:
merge_command = [
"ffmpeg", "-y",
"-i", temp_video_path,
"-i", temp_audio_path,
"-c:v", "copy", # 直接复制视频流,速度快
"-c:a", "aac", # 编码音频流
"-shortest", # 以最短的流为准结束
final_video_path,
]
print("\n[调试信息] 将要执行音视频合并命令:")
print(" ".join(merge_command))
# 执行命令并显示所有输出
subprocess.run(merge_command, check=True)
print("视频和音频合并成功。")
except FileNotFoundError:
print("\n错误: 'ffmpeg' 命令未找到。请确保 FFmpeg 已正确安装并已添加到系统的 PATH 环境变量中。")
except subprocess.CalledProcessError as e:
print(f"\n错误:合并视频和音频失败。FFmpeg 返回了非零退出码。请检查上面的 FFmpeg 输出信息以了解详情。")
print(f"Python 错误详情: {e}")
finally:
# =========================================================================
# PART 3: 清理临时文件
# =========================================================================
print("步骤 3/3: 正在清理临时文件...")
if os.path.exists(temp_video_path):
os.remove(temp_video_path)
if os.path.exists(temp_audio_path):
os.remove(temp_audio_path)
if os.path.exists(final_video_path):
print(f"任务完成!最终视频已成功保存到: {final_video_path}")
else:
print("任务失败,未生成最终文件。请检查上面的日志输出。")
def split_token_counts_and_frame_ids(T, token_frame, world_size, rank):
S = T * token_frame
split_sizes = [S // world_size + (1 if i < S % world_size else 0) for i in range(world_size)]
start = sum(split_sizes[:rank])
end = start + split_sizes[rank]
counts = [0] * T
for idx in range(start, end):
t = idx // token_frame
counts[t] += 1
counts_filtered = []
frame_ids = []
for t, c in enumerate(counts):
if c > 0:
counts_filtered.append(c)
frame_ids.append(t)
return counts_filtered, frame_ids
def normalize_and_scale(column, source_range, target_range, epsilon=1e-8):
source_min, source_max = source_range
new_min, new_max = target_range
normalized = (column - source_min) / (source_max - source_min + epsilon)
scaled = normalized * (new_max - new_min) + new_min
return scaled
@torch.compile
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None):
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
scale = 1.0 / visual_q.shape[-1] ** 0.5
visual_q = visual_q * scale
visual_q = visual_q.transpose(1, 2)
ref_k = ref_k.transpose(1, 2)
attn = visual_q @ ref_k.transpose(-2, -1)
if attn_bias is not None:
attn = attn + attn_bias
x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens
x_ref_attn_maps = []
ref_target_masks = ref_target_masks.to(visual_q.dtype)
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
for class_idx, ref_target_mask in enumerate(ref_target_masks):
torch_gc()
ref_target_mask = ref_target_mask[None, None, None, ...]
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
if mode == 'mean':
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens
elif mode == 'max':
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
x_ref_attn_maps.append(x_ref_attnmap)
del attn
del x_ref_attn_map_source
torch_gc()
return torch.concat(x_ref_attn_maps, dim=0)
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2, enable_sp=False):
"""Args:
query (torch.tensor): B M H K
key (torch.tensor): B M H K
shape (tuple): (N_t, N_h, N_w)
ref_target_masks: [B, N_h * N_w]
"""
N_t, N_h, N_w = shape
if enable_sp:
ref_k = get_sp_group().all_gather(ref_k, dim=1)
x_seqlens = N_h * N_w
ref_k = ref_k[:, :x_seqlens]
_, seq_lens, heads, _ = visual_q.shape
class_num, _ = ref_target_masks.shape
x_ref_attn_maps = torch.zeros(class_num, seq_lens).to(visual_q.device).to(visual_q.dtype)
split_chunk = heads // split_num
for i in range(split_num):
x_ref_attn_maps_perhead = calculate_x_ref_attn_map(
visual_q[:, :, i*split_chunk:(i+1)*split_chunk, :],
ref_k[:, :, i*split_chunk:(i+1)*split_chunk, :],
ref_target_masks
)
x_ref_attn_maps += x_ref_attn_maps_perhead
return x_ref_attn_maps / split_num
def rotate_half(x):
x = rearrange(x, "... (d r) -> ... d r", r=2)
x1, x2 = x.unbind(dim=-1)
x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, "... d r -> ... (d r)")
class RotaryPositionalEmbedding1D(nn.Module):
def __init__(self,
head_dim,
):
super().__init__()
self.head_dim = head_dim
self.base = 10000
@lru_cache(maxsize=32)
def precompute_freqs_cis_1d(self, pos_indices):
freqs = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2)[: (self.head_dim // 2)].float() / self.head_dim))
freqs = freqs.to(pos_indices.device)
freqs = torch.einsum("..., f -> ... f", pos_indices.float(), freqs)
freqs = repeat(freqs, "... n -> ... (n r)", r=2)
return freqs
def forward(self, x, pos_indices):
"""1D RoPE.
Args:
query (torch.tensor): [B, head, seq, head_dim]
pos_indices (torch.tensor): [seq,]
Returns:
query with the same shape as input.
"""
freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
x_ = x.float()
freqs_cis = freqs_cis.float().to(x.device)
cos, sin = freqs_cis.cos(), freqs_cis.sin()
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d')
x_ = (x_ * cos) + (rotate_half(x_) * sin)
return x_.type_as(x)
def save_video_ffmpeg(gen_video_samples, save_path, vocal_audio_list, fps=25, quality=5):
def save_video(frames, save_path, fps, quality=9, ffmpeg_params=None):
writer = imageio.get_writer(
save_path, fps=fps, quality=quality, ffmpeg_params=ffmpeg_params
)
for frame in tqdm(frames, desc="Saving video"):
frame = np.array(frame)
writer.append_data(frame)
writer.close()
save_path_tmp = save_path + "-temp.mp4"
video_audio = (gen_video_samples+1)/2 # C T H W
video_audio = video_audio.permute(1, 2, 3, 0).cpu().numpy()
video_audio = np.clip(video_audio * 255, 0, 255).astype(np.uint8)
save_video(video_audio, save_path_tmp, fps=fps, quality=quality)
# crop audio according to video length
_, T, _, _ = gen_video_samples.shape
duration = T / fps
save_path_crop_audio = save_path + "-cropaudio.wav"
final_command = [
"ffmpeg",
"-i",
vocal_audio_list[0],
"-t",
f'{duration}',
save_path_crop_audio,
]
subprocess.run(final_command, check=True)
# generate video with audio
save_path = save_path + ".mp4"
final_command = [
"ffmpeg",
"-y",
"-i",
save_path_tmp,
"-i",
save_path_crop_audio,
"-c:v",
"libx264",
"-c:a",
"aac",
"-shortest",
save_path,
]
subprocess.run(final_command, check=True)
os.remove(save_path_tmp)
os.remove(save_path_crop_audio)
class MomentumBuffer:
def __init__(self, momentum: float):
self.momentum = momentum
self.running_average = 0
def update(self, update_value: torch.Tensor):
new_average = self.momentum * self.running_average
self.running_average = update_value + new_average
def project(
v0: torch.Tensor, # [B, C, T, H, W]
v1: torch.Tensor, # [B, C, T, H, W]
):
dtype = v0.dtype
v0, v1 = v0.double(), v1.double()
v1 = torch.nn.functional.normalize(v1, dim=[-1, -2, -3, -4])
v0_parallel = (v0 * v1).sum(dim=[-1, -2, -3, -4], keepdim=True) * v1
v0_orthogonal = v0 - v0_parallel
return v0_parallel.to(dtype), v0_orthogonal.to(dtype)
def adaptive_projected_guidance(
diff: torch.Tensor, # [B, C, T, H, W]
pred_cond: torch.Tensor, # [B, C, T, H, W]
momentum_buffer: MomentumBuffer = None,
eta: float = 0.0,
norm_threshold: float = 55,
):
if momentum_buffer is not None:
momentum_buffer.update(diff)
diff = momentum_buffer.running_average
if norm_threshold > 0:
ones = torch.ones_like(diff)
diff_norm = diff.norm(p=2, dim=[-1, -2, -3, -4], keepdim=True)
print(f"diff_norm: {diff_norm}")
scale_factor = torch.minimum(ones, norm_threshold / diff_norm)
diff = diff * scale_factor
diff_parallel, diff_orthogonal = project(diff, pred_cond)
normalized_update = diff_orthogonal + eta * diff_parallel
return normalized_update
+118
View File
@@ -0,0 +1,118 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import argparse
import binascii
import os
import os.path as osp
import imageio
import torch
import torchvision
__all__ = ['cache_video', 'cache_image', 'str2bool']
def rand_name(length=8, suffix=''):
name = binascii.b2a_hex(os.urandom(length)).decode('utf-8')
if suffix:
if not suffix.startswith('.'):
suffix = '.' + suffix
name += suffix
return name
def cache_video(tensor,
save_file=None,
fps=30,
suffix='.mp4',
nrow=8,
normalize=True,
value_range=(-1, 1),
retry=5):
# cache file
cache_file = osp.join('/tmp', rand_name(
suffix=suffix)) if save_file is None else save_file
# save to cache
error = None
for _ in range(retry):
try:
# preprocess
tensor = tensor.clamp(min(value_range), max(value_range))
tensor = torch.stack([
torchvision.utils.make_grid(
u, nrow=nrow, normalize=normalize, value_range=value_range)
for u in tensor.unbind(2)
],
dim=1).permute(1, 2, 3, 0)
tensor = (tensor * 255).type(torch.uint8).cpu()
# write video
writer = imageio.get_writer(
cache_file, fps=fps, codec='libx264', quality=8)
for frame in tensor.numpy():
writer.append_data(frame)
writer.close()
return cache_file
except Exception as e:
error = e
continue
else:
print(f'cache_video failed, error: {error}', flush=True)
return None
def cache_image(tensor,
save_file,
nrow=8,
normalize=True,
value_range=(-1, 1),
retry=5):
# cache file
suffix = osp.splitext(save_file)[1]
if suffix.lower() not in [
'.jpg', '.jpeg', '.png', '.tiff', '.gif', '.webp'
]:
suffix = '.png'
# save to cache
error = None
for _ in range(retry):
try:
tensor = tensor.clamp(min(value_range), max(value_range))
torchvision.utils.save_image(
tensor,
save_file,
nrow=nrow,
normalize=normalize,
value_range=value_range)
return save_file
except Exception as e:
error = e
continue
def str2bool(v):
"""
Convert a string to a boolean.
Supported true values: 'yes', 'true', 't', 'y', '1'
Supported false values: 'no', 'false', 'f', 'n', '0'
Args:
v (str): String to convert.
Returns:
bool: Converted boolean value.
Raises:
argparse.ArgumentTypeError: If the value cannot be converted to boolean.
"""
if isinstance(v, bool):
return v
v_lower = v.lower()
if v_lower in ('yes', 'true', 't', 'y', '1'):
return True
elif v_lower in ('no', 'false', 'f', 'n', '0'):
return False
else:
raise argparse.ArgumentTypeError('Boolean value expected (True/False)')