@@ -0,0 +1 @@
|
||||
*.mp4 filter=lfs diff=lfs merge=lfs -text
|
||||
@@ -0,0 +1,8 @@
|
||||
__pyc*
|
||||
*/__pyc*
|
||||
*/*/__pyc*
|
||||
readme2*
|
||||
results
|
||||
temp
|
||||
hostfile*
|
||||
env.txt
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
|
After Width: | Height: | Size: 12 KiB |
|
After Width: | Height: | Size: 13 KiB |
|
After Width: | Height: | Size: 11 KiB |
|
After Width: | Height: | Size: 1.2 MiB |
|
After Width: | Height: | Size: 1.0 MiB |
|
After Width: | Height: | Size: 864 KiB |
|
After Width: | Height: | Size: 143 KiB |
|
After Width: | Height: | Size: 13 KiB |
|
After Width: | Height: | Size: 37 KiB |
|
After Width: | Height: | Size: 37 KiB |
@@ -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
|
||||
|
After Width: | Height: | Size: 387 KiB |
|
After Width: | Height: | Size: 292 KiB |
|
After Width: | Height: | Size: 258 KiB |
@@ -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()
|
||||
@@ -0,0 +1,3 @@
|
||||
|
||||
from .InteractAvatar_node import *
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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')
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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()),
|
||||
}
|
||||
@@ -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压缩残留,丑陋的,残缺的,多余的手指,' \
|
||||
'画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,' \
|
||||
'手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走'
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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',
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,7 @@
|
||||
# import modules
|
||||
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
|
||||
__all__ = [
|
||||
'HuggingfaceTokenizer', 'get_sampling_sigmas', 'retrieve_timesteps',
|
||||
'FlowDPMSolverMultistepScheduler', 'FlowUniPCMultistepScheduler'
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)')
|
||||