Merge pull request #1 from smthemex/Pr

init
This commit is contained in:
smthemex
2026-04-06 21:53:09 +08:00
committed by GitHub
57 changed files with 916010 additions and 1 deletions
+7
View File
@@ -0,0 +1,7 @@
__pycache__/
*.pyc
.pytest_cache/
test_outputs/
outputs/
+28
View File
@@ -0,0 +1,28 @@
{
"</think>": 151668,
"</tool_call>": 151658,
"</tool_response>": 151666,
"<think>": 151667,
"<tool_call>": 151657,
"<tool_response>": 151665,
"<|box_end|>": 151649,
"<|box_start|>": 151648,
"<|endoftext|>": 151643,
"<|file_sep|>": 151664,
"<|fim_middle|>": 151660,
"<|fim_pad|>": 151662,
"<|fim_prefix|>": 151659,
"<|fim_suffix|>": 151661,
"<|im_end|>": 151645,
"<|im_start|>": 151644,
"<|image_pad|>": 151655,
"<|object_ref_end|>": 151647,
"<|object_ref_start|>": 151646,
"<|quad_end|>": 151651,
"<|quad_start|>": 151650,
"<|repo_name|>": 151663,
"<|video_pad|>": 151656,
"<|vision_end|>": 151653,
"<|vision_pad|>": 151654,
"<|vision_start|>": 151652
}
+120
View File
@@ -0,0 +1,120 @@
{%- if tools %}
{{- '<|im_start|>system\n' }}
{%- if messages[0].role == 'system' %}
{%- if messages[0].content is string %}
{{- messages[0].content }}
{%- else %}
{%- for content in messages[0].content %}
{%- if 'text' in content %}
{{- content.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{{- '\n\n' }}
{%- endif %}
{{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
{%- for tool in tools %}
{{- "\n" }}
{{- tool | tojson }}
{%- endfor %}
{{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
{%- else %}
{%- if messages[0].role == 'system' %}
{{- '<|im_start|>system\n' }}
{%- if messages[0].content is string %}
{{- messages[0].content }}
{%- else %}
{%- for content in messages[0].content %}
{%- if 'text' in content %}
{{- content.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{{- '<|im_end|>\n' }}
{%- endif %}
{%- endif %}
{%- set image_count = namespace(value=0) %}
{%- set video_count = namespace(value=0) %}
{%- for message in messages %}
{%- if message.role == "user" %}
{{- '<|im_start|>' + message.role + '\n' }}
{%- if message.content is string %}
{{- message.content }}
{%- else %}
{%- for content in message.content %}
{%- if content.type == 'image' or 'image' in content or 'image_url' in content %}
{%- set image_count.value = image_count.value + 1 %}
{%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}
<|vision_start|><|image_pad|><|vision_end|>
{%- elif content.type == 'video' or 'video' in content %}
{%- set video_count.value = video_count.value + 1 %}
{%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}
<|vision_start|><|video_pad|><|vision_end|>
{%- elif 'text' in content %}
{{- content.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{{- '<|im_end|>\n' }}
{%- elif message.role == "assistant" %}
{{- '<|im_start|>' + message.role + '\n' }}
{%- if message.content is string %}
{{- message.content }}
{%- else %}
{%- for content_item in message.content %}
{%- if 'text' in content_item %}
{{- content_item.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{%- if message.tool_calls %}
{%- for tool_call in message.tool_calls %}
{%- if (loop.first and message.content) or (not loop.first) %}
{{- '\n' }}
{%- endif %}
{%- if tool_call.function %}
{%- set tool_call = tool_call.function %}
{%- endif %}
{{- '<tool_call>\n{"name": "' }}
{{- tool_call.name }}
{{- '", "arguments": ' }}
{%- if tool_call.arguments is string %}
{{- tool_call.arguments }}
{%- else %}
{{- tool_call.arguments | tojson }}
{%- endif %}
{{- '}\n</tool_call>' }}
{%- endfor %}
{%- endif %}
{{- '<|im_end|>\n' }}
{%- elif message.role == "tool" %}
{%- if loop.first or (messages[loop.index0 - 1].role != "tool") %}
{{- '<|im_start|>user' }}
{%- endif %}
{{- '\n<tool_response>\n' }}
{%- if message.content is string %}
{{- message.content }}
{%- else %}
{%- for content in message.content %}
{%- if content.type == 'image' or 'image' in content or 'image_url' in content %}
{%- set image_count.value = image_count.value + 1 %}
{%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}
<|vision_start|><|image_pad|><|vision_end|>
{%- elif content.type == 'video' or 'video' in content %}
{%- set video_count.value = video_count.value + 1 %}
{%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}
<|vision_start|><|video_pad|><|vision_end|>
{%- elif 'text' in content %}
{{- content.text }}
{%- endif %}
{%- endfor %}
{%- endif %}
{{- '\n</tool_response>' }}
{%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
{{- '<|im_end|>\n' }}
{%- endif %}
{%- endif %}
{%- endfor %}
{%- if add_generation_prompt %}
{{- '<|im_start|>assistant\n' }}
{%- endif %}
+64
View File
@@ -0,0 +1,64 @@
{
"architectures": [
"Qwen3VLForConditionalGeneration"
],
"dtype": "bfloat16",
"image_token_id": 151655,
"model_type": "qwen3_vl",
"text_config": {
"attention_bias": false,
"attention_dropout": 0.0,
"bos_token_id": 151643,
"dtype": "bfloat16",
"eos_token_id": 151645,
"head_dim": 128,
"hidden_act": "silu",
"hidden_size": 4096,
"initializer_range": 0.02,
"intermediate_size": 12288,
"max_position_embeddings": 262144,
"model_type": "qwen3_vl_text",
"num_attention_heads": 32,
"num_hidden_layers": 36,
"num_key_value_heads": 8,
"rms_norm_eps": 1e-06,
"rope_scaling": {
"mrope_interleaved": true,
"mrope_section": [
24,
20,
20
],
"rope_type": "default"
},
"rope_theta": 5000000,
"use_cache": true,
"vocab_size": 151936
},
"tie_word_embeddings": false,
"transformers_version": "4.57.1",
"video_token_id": 151656,
"vision_config": {
"deepstack_visual_indexes": [
8,
16,
24
],
"depth": 27,
"dtype": "bfloat16",
"hidden_act": "gelu_pytorch_tanh",
"hidden_size": 1152,
"in_channels": 3,
"initializer_range": 0.02,
"intermediate_size": 4304,
"model_type": "qwen3_vl",
"num_heads": 16,
"num_position_embeddings": 2304,
"out_hidden_size": 4096,
"patch_size": 16,
"spatial_merge_size": 2,
"temporal_patch_size": 2
},
"vision_end_token_id": 151653,
"vision_start_token_id": 151652
}
+13
View File
@@ -0,0 +1,13 @@
{
"bos_token_id": 151643,
"do_sample": true,
"eos_token_id": [
151645,
151643
],
"pad_token_id": 151643,
"temperature": 0.7,
"top_k": 20,
"top_p": 0.8,
"transformers_version": "4.57.1"
}
File diff suppressed because it is too large Load Diff
+39
View File
@@ -0,0 +1,39 @@
{
"crop_size": null,
"data_format": "channels_first",
"default_to_square": true,
"device": null,
"disable_grouping": null,
"do_center_crop": null,
"do_convert_rgb": true,
"do_normalize": true,
"do_pad": null,
"do_rescale": true,
"do_resize": true,
"image_mean": [
0.5,
0.5,
0.5
],
"image_processor_type": "Qwen2VLImageProcessorFast",
"image_std": [
0.5,
0.5,
0.5
],
"input_data_format": null,
"max_pixels": null,
"merge_size": 2,
"min_pixels": null,
"pad_size": null,
"patch_size": 16,
"processor_class": "Qwen3VLProcessor",
"resample": 3,
"rescale_factor": 0.00392156862745098,
"return_tensors": null,
"size": {
"longest_edge": 16777216,
"shortest_edge": 65536
},
"temporal_patch_size": 2
}
+31
View File
@@ -0,0 +1,31 @@
{
"additional_special_tokens": [
"<|im_start|>",
"<|im_end|>",
"<|object_ref_start|>",
"<|object_ref_end|>",
"<|box_start|>",
"<|box_end|>",
"<|quad_start|>",
"<|quad_end|>",
"<|vision_start|>",
"<|vision_end|>",
"<|vision_pad|>",
"<|image_pad|>",
"<|video_pad|>"
],
"eos_token": {
"content": "<|im_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false
},
"pad_token": {
"content": "<|endoftext|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false
}
}
File diff suppressed because it is too large Load Diff
+241
View File
@@ -0,0 +1,241 @@
{
"add_bos_token": false,
"add_prefix_space": false,
"added_tokens_decoder": {
"151643": {
"content": "<|endoftext|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151644": {
"content": "<|im_start|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151645": {
"content": "<|im_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151646": {
"content": "<|object_ref_start|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151647": {
"content": "<|object_ref_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151648": {
"content": "<|box_start|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151649": {
"content": "<|box_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151650": {
"content": "<|quad_start|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151651": {
"content": "<|quad_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151652": {
"content": "<|vision_start|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151653": {
"content": "<|vision_end|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151654": {
"content": "<|vision_pad|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151655": {
"content": "<|image_pad|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151656": {
"content": "<|video_pad|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": true
},
"151657": {
"content": "<tool_call>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151658": {
"content": "</tool_call>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151659": {
"content": "<|fim_prefix|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151660": {
"content": "<|fim_middle|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151661": {
"content": "<|fim_suffix|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151662": {
"content": "<|fim_pad|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151663": {
"content": "<|repo_name|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151664": {
"content": "<|file_sep|>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151665": {
"content": "<tool_response>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151666": {
"content": "</tool_response>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151667": {
"content": "<think>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
},
"151668": {
"content": "</think>",
"lstrip": false,
"normalized": false,
"rstrip": false,
"single_word": false,
"special": false
}
},
"additional_special_tokens": [
"<|im_start|>",
"<|im_end|>",
"<|object_ref_start|>",
"<|object_ref_end|>",
"<|box_start|>",
"<|box_end|>",
"<|quad_start|>",
"<|quad_end|>",
"<|vision_start|>",
"<|vision_end|>",
"<|vision_pad|>",
"<|image_pad|>",
"<|video_pad|>"
],
"bos_token": null,
"clean_up_tokenization_spaces": false,
"eos_token": "<|im_end|>",
"errors": "replace",
"extra_special_tokens": {},
"model_max_length": 262144,
"pad_token": "<|endoftext|>",
"processor_class": "Qwen3VLProcessor",
"split_special_tokens": false,
"tokenizer_class": "Qwen2Tokenizer",
"unk_token": null,
"use_fast": true
}
@@ -0,0 +1,41 @@
{
"crop_size": null,
"data_format": "channels_first",
"default_to_square": true,
"device": null,
"do_center_crop": null,
"do_convert_rgb": true,
"do_normalize": true,
"do_rescale": true,
"do_resize": true,
"do_sample_frames": true,
"fps": 2,
"image_mean": [
0.5,
0.5,
0.5
],
"image_std": [
0.5,
0.5,
0.5
],
"input_data_format": null,
"max_frames": 768,
"merge_size": 2,
"min_frames": 4,
"num_frames": null,
"pad_size": null,
"patch_size": 16,
"processor_class": "Qwen3VLProcessor",
"resample": 3,
"rescale_factor": 0.00392156862745098,
"return_metadata": false,
"size": {
"longest_edge": 25165824,
"shortest_edge": 4096
},
"temporal_patch_size": 2,
"video_metadata": null,
"video_processor_type": "Qwen3VLVideoProcessor"
}
File diff suppressed because one or more lines are too long
+256
View File
@@ -0,0 +1,256 @@
# !/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
import nodes
from .model_loader_utils import clear_comfyui_cache,save_lat_emb,read_lat_emb,tensor2pillist,phi2narry,tensor2pillist_upscale
from .inference import load_mmdit,infer_joyai,load_vae,get_latents,vae_decode
from .inference_und import get_conditioning,load_qwen3vl_model,encoder_input
MAX_SEED = np.iinfo(np.int32).max
node_cr_path = os.path.dirname(os.path.abspath(__file__))
device = torch.device(
"cuda:0") if torch.cuda.is_available() else torch.device(
"mps") if torch.backends.mps.is_available() else torch.device(
"cpu")
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 JoyAI_Image_SM_Model(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_SM_Model",
display_name="JoyAI_Image_SM_Model",
category="JoyAI_Image",
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")),
],
outputs=[
io.Model.Output(display_name="model"),
],
)
@classmethod
def execute(cls,dit,gguf) -> io.NodeOutput:
clear_comfyui_cache()
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
model= load_mmdit(dit_path,gguf_path,True)
return io.NodeOutput(model)
class JoyAI_Image_SM_VAE(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_SM_VAE",
display_name="JoyAI_Image_SM_VAE",
category="JoyAI_Image",
inputs=[
io.Combo.Input("vae",options= ["none"] + folder_paths.get_filename_list("vae") ),
],
outputs=[io.Vae.Output(display_name="vae"),],
)
@classmethod
def execute(cls,vae ) -> io.NodeOutput:
clear_comfyui_cache()
vae_path=folder_paths.get_full_path("vae", vae) if vae != "none" else None
vae=load_vae(vae_path,device,torch.bfloat16)
return io.NodeOutput(vae)
class JoyAI_Image_SM_Clip(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_SM_Clip",
display_name="JoyAI_Image_SM_Clip",
category="JoyAI_Image",
inputs=[
io.Combo.Input("clip",options= ["none"] + folder_paths.get_filename_list("clip") ),
io.Combo.Input("gguf",options= ["none"] + folder_paths.get_filename_list("gguf") ),
],
outputs=[io.Clip.Output(display_name="clip"),],
)
@classmethod
def execute(cls,clip,gguf ) -> io.NodeOutput:
clear_comfyui_cache()
safetensors_path=folder_paths.get_full_path("clip", clip) if clip != "none" else None
gguf_path=folder_paths.get_full_path("gguf", gguf) if gguf != "none" else None
clip=load_qwen3vl_model(safetensors_path,gguf_path)
return io.NodeOutput(clip)
class JoyAI_Vae_Decoder(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Vae_Decoder",
display_name="JoyAI_Vae_Decoder",
category="JoyAI_Image",
inputs=[
io.Vae.Input("vae"),
io.Latent.Input("latents",),
],
outputs=
[io.Image.Output(display_name="image"),],
)
@classmethod
def execute(cls,vae,latents ) -> io.NodeOutput:
clear_comfyui_cache()
image=vae_decode(vae,latents)
return io.NodeOutput(image)
class JoyAI_Image_LATENTS(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_LATENTS",
display_name="JoyAI_Image_LATENTS",
category="JoyAI_Image",
inputs=[
io.Image.Input("image"),
io.Int.Input("seed", default=0, min=0, max=MAX_SEED,display_mode=io.NumberDisplay.number),
io.Int.Input("width", default=1024, min=256, max=nodes.MAX_RESOLUTION,step=32,display_mode=io.NumberDisplay.number),
io.Int.Input("height", default=1024, min=256, max=nodes.MAX_RESOLUTION,step=32,display_mode=io.NumberDisplay.number),
io.Vae.Input("vae",optional=True),
],
outputs=[
io.Latent.Output(display_name="latent"),
],
)
@classmethod
def execute(cls,image,seed,width,height,vae=None,) -> io.NodeOutput:
clear_comfyui_cache()
# width=(width //32)*32 if width % 32 != 0 else width
# height=(height //32)*32 if height % 32 != 0 else height
images=tensor2pillist_upscale(image,width,height) if image is not None else None
if vae is None and images is not None:
raise Exception("When use image,you must provide a vae")
lat,_=get_latents(vae, images, height, width, device,seed,image, torch.bfloat16)
latent={"samples":lat,"width":width,"height":height,"images":images}
return io.NodeOutput(latent)
class JoyAI_Image_ENCODER(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_ENCODER",
display_name="JoyAI_Image_ENCODER",
category="JoyAI_Image",
inputs=[
io.Clip.Input("clip"),
io.String.Input("prompt",multiline=True,default="Turn the plate blue" ),
io.Combo.Input("infer_device",options= ["cuda","cpu"] ),
io.Boolean.Input("save_emb",default=False),
io.Image.Input("image",optional=True),
],
outputs=[
io.Conditioning.Output(display_name="positive"),
io.Conditioning.Output(display_name="negative"),
],
)
@classmethod
def execute(cls,clip,prompt, infer_device,save_emb,image=None) -> io.NodeOutput:
clear_comfyui_cache()
images=tensor2pillist(image) if image is not None else None
positive,negative=get_conditioning(clip,prompt, images,infer_device)
if save_emb:
save_lat_emb("embeds",positive,negative)
return io.NodeOutput(positive,negative)
class JoyAI_Image_Understand(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_Understand",
display_name="JoyAI_Image_Understand",
category="JoyAI_Image",
inputs=[
io.Clip.Input("clip"),
io.Image.Input("image"),
io.String.Input("prompt",multiline=True,default="Turn the plate blue" ),
io.Int.Input("max_new_tokens", default=2048, min=256, max=nodes.MAX_RESOLUTION,step=1,display_mode=io.NumberDisplay.number),
io.Float.Input("temperature", default=0.7, min=0, max=1,step=0.01,display_mode=io.NumberDisplay.number),
io.Float.Input("top_p", default=0.8, min=0, max=1,step=0.01,display_mode=io.NumberDisplay.number),
io.Int.Input("top_k", default=50, min=1, max=200,step=1,display_mode=io.NumberDisplay.number),
io.Combo.Input("infer_device",options= ["cuda","cpu"] ),
],
outputs=[
io.String.Output(display_name="response"),
],
)
@classmethod
def execute(cls,clip,image,prompt,max_new_tokens,temperature,top_p,top_k,infer_device,) -> io.NodeOutput:
clear_comfyui_cache()
images=tensor2pillist(image)
response=encoder_input(clip,prompt,images,max_new_tokens,top_p,top_k,temperature,infer_device)
return io.NodeOutput(response)
class JoyAI_Image_SM_KSampler(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="JoyAI_Image_SM_KSampler",
display_name="JoyAI_Image_SM_KSampler",
category="JoyAI_Image",
inputs=[
io.Model.Input("model"),
io.Latent.Input("latents",),
io.Int.Input("steps", default=20, min=1, max=nodes.MAX_RESOLUTION,step=1,display_mode=io.NumberDisplay.number),
io.Float.Input("guidance_scale", default=5.0, min=1, max=20,step=0.1,display_mode=io.NumberDisplay.number),
io.Boolean.Input("offload", default=True),
io.Int.Input("offload_block_num", default=1, min=1, max=40,step=1,display_mode=io.NumberDisplay.number),
io.Conditioning.Input("positive",optional=True),
io.Conditioning.Input("negative",optional=True),
],
outputs=[
io.Latent.Output(display_name="latent"),
],
)
@classmethod
def execute(cls, model,latents,steps,guidance_scale,offload,offload_block_num,positive=None,negative=None,) -> io.NodeOutput:
if positive is None:
positive,negative=read_lat_emb("embeds",device)
clear_comfyui_cache()
if not offload:
model.dit.to(device)
lat=infer_joyai(model,latents,positive,negative, steps, guidance_scale,offload,offload_block_num)
if not offload:
model.dit.to("cpu")
latent={"samples":lat}
return io.NodeOutput(latent)
class JoyAI_Image_SM_Extension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
JoyAI_Image_SM_Model,
JoyAI_Image_SM_VAE,
JoyAI_Image_SM_Clip,
JoyAI_Image_LATENTS,
JoyAI_Image_SM_KSampler,
JoyAI_Image_ENCODER,
JoyAI_Vae_Decoder,
JoyAI_Image_Understand,
]
async def comfy_entrypoint() -> JoyAI_Image_SM_Extension: # ComfyUI calls this to load your extension and its nodes.
return JoyAI_Image_SM_Extension()
+61 -1
View File
@@ -1,2 +1,62 @@
# ComfyUI_JoyAI_Image
Awakening Spatial Intelligence in Unified Multimodal Understanding and Generation
[JoyAI-Image](https://github.com/jd-opensource/JoyAI-Image):Awakening Spatial Intelligence in Unified Multimodal Understanding and Generation
Update
-----
* Need 64RAM+8VRAM
1.Installation
-----
* In the ./ComfyUI/custom_nodes directory, run the following:
```
git clone https://github.com/smthemex/ComfyUI_JoyAI_Image
```
2.requirements
----
```
pip install -r requirements.txt
```
3.checkpoints
----
* transformers/vae/clip [links](https://huggingface.co/jdopensource/JoyAI-Image-Edit)
* or [aliyun](https://pan.quark.cn/s/e20f511c921c)
* or [hg](https://huggingface.co/smthem/JoyAI-Image-Edit-merge-dit-gguf)
```
├── ComfyUI/models/
| ├── diffusion_models/
| ├──joy_image_transformer.safetensors
| ├── vae/
| ├──Wan2.1_VAE.pth
| ├── clips
| ├──JoyAI-Image-Und-merger_bf16.safetensors
```
4.Example
----
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example.png)
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example2.png)
![](https://github.com/smthemex/ComfyUI_JoyAI_Image/blob/main/example_workflows/example3.png)
5.Citation
----
```
@article{JoyAI2023,}
```
+2
View File
@@ -0,0 +1,2 @@
from .JoyAI_Image_node import *
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.1 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 681 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 537 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 9.6 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 598 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 795 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 342 KiB

+838
View File
@@ -0,0 +1,838 @@
{
"id": "2a1ae491-ed35-4bc8-b4dd-a5cee997013f",
"revision": 0,
"last_node_id": 32,
"last_link_id": 61,
"nodes": [
{
"id": 4,
"type": "JoyAI_Image_LATENTS",
"pos": [
20842.725512877667,
-492.32260878653443
],
"size": [
270,
150
],
"flags": {},
"order": 8,
"mode": 2,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 12
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": 25
}
],
"outputs": [
{
"name": "latent",
"type": "LATENT",
"links": [
1
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_LATENTS"
},
"widgets_values": [
1216048398,
"fixed",
1024,
1024
]
},
{
"id": 14,
"type": "JoyAI_Image_ENCODER",
"pos": [
20370.296397399088,
-457.29533838974066
],
"size": [
400,
200
],
"flags": {},
"order": 9,
"mode": 2,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 16
},
{
"name": "image",
"shape": 7,
"type": "IMAGE",
"link": 17
}
],
"outputs": [
{
"name": "positive",
"type": "CONDITIONING",
"links": [
18,
23
]
},
{
"name": "negative",
"type": "CONDITIONING",
"links": [
24
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_ENCODER"
},
"widgets_values": [
"Turn the plate blue",
"cuda",
true
]
},
{
"id": 1,
"type": "JoyAI_Image_SM_Model",
"pos": [
20804.444408542826,
-642.9905545506151
],
"size": [
343.41596747239964,
85.14639060425452
],
"flags": {},
"order": 0,
"mode": 2,
"inputs": [],
"outputs": [
{
"name": "model",
"type": "MODEL",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_SM_Model"
},
"widgets_values": [
"joy_image_transformer.safetensors",
"none"
]
},
{
"id": 10,
"type": "SaveImage",
"pos": [
21142.32249594224,
-302.6532458634289
],
"size": [
316.52843317599763,
352.14052057116317
],
"flags": {},
"order": 16,
"mode": 2,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 53
}
],
"outputs": [],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 16,
"type": "JoyAI_Vae_Decoder",
"pos": [
21178.79638603116,
-454.53169313079843
],
"size": [
257.58332745706355,
46
],
"flags": {},
"order": 14,
"mode": 2,
"inputs": [
{
"name": "vae",
"type": "VAE",
"link": 54
},
{
"name": "latents",
"type": "LATENT",
"link": 27
}
],
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": [
53
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Vae_Decoder"
},
"widgets_values": []
},
{
"id": 15,
"type": "JoyAI_Image_SM_VAE",
"pos": [
21175.02808913353,
-620.3113417158576
],
"size": [
270,
58
],
"flags": {},
"order": 1,
"mode": 2,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "VAE",
"links": [
25,
54
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_SM_VAE"
},
"widgets_values": [
"Wan2.1_VAE.pth"
]
},
{
"id": 5,
"type": "JoyAI_Image_SM_KSampler",
"pos": [
20827.07829554382,
-279.9061113465927
],
"size": [
282.4810546875,
190
],
"flags": {},
"order": 12,
"mode": 2,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 3
},
{
"name": "latents",
"type": "LATENT",
"link": 1
},
{
"name": "positive",
"shape": 7,
"type": "CONDITIONING",
"link": 23
},
{
"name": "negative",
"shape": 7,
"type": "CONDITIONING",
"link": 24
}
],
"outputs": [
{
"name": "latent",
"type": "LATENT",
"links": [
27
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_SM_KSampler"
},
"widgets_values": [
20,
7.5,
true,
2
]
},
{
"id": 13,
"type": "PreviewAny",
"pos": [
20372.948979592402,
-213.27126551002547
],
"size": [
400.40483637695434,
272.46636091078096
],
"flags": {},
"order": 11,
"mode": 2,
"inputs": [
{
"name": "source",
"type": "*",
"link": 18
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewAny"
},
"widgets_values": [
null,
null,
null
]
},
{
"id": 3,
"type": "JoyAI_Image_SM_Clip",
"pos": [
20469.48060913142,
-619.9307633964665
],
"size": [
270,
82
],
"flags": {},
"order": 2,
"mode": 2,
"inputs": [],
"outputs": [
{
"name": "clip",
"type": "CLIP",
"links": [
16
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_SM_Clip"
},
"widgets_values": [
"JoyAI-Image-Und-merger_bf16.safetensors",
"none"
]
},
{
"id": 11,
"type": "LoadImage",
"pos": [
20002.973180667486,
-547.0052369078016
],
"size": [
292.92662709497745,
479.8369035816668
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
12,
17
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"test_1.jpg",
"image"
]
},
{
"id": 17,
"type": "Note",
"pos": [
19504.319241910616,
-638.5414753541683
],
"size": [
454.2734004190024,
574.5500440018349
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"\n1、 移动物体示例,\nMove the <object> into the red box and finally remove the red box.\n\n实际提示词\nMove the apple into the red box and finally remove the red box.\n\n\n\n2、视角切换prompt示例:\nRotate the <object> to show the <view> side view.\n\n支持的视角切换 <view> 替换掉view\nfront\nright\nleft\nrear\nfront right\nfront left\nrear right\nrear left\n\n实际提示词\nRotate the chair to show the front side view.\nRotate the car to show the rear left side view.\n\n\n3、 镜头控制\nMove the camera.\n- Camera rotation: Yaw {y_rotation}°, Pitch {p_rotation}°.\n- Camera zoom: in/out/unchanged.\n- Keep the 3D scene static; only change the viewpoint.\n\n实际提示词:\nMove the camera.\n- Camera rotation: Yaw 45°, Pitch 0°.\n- Camera zoom: in.\n- Keep the 3D scene static; only change the viewpoint.\n\n或者\nMove the camera.\n- Camera rotation: Yaw -90°, Pitch 20°.\n- Camera zoom: unchanged.\n- Keep the 3D scene static; only change the viewpoint."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 31,
"type": "ImageBatch",
"pos": [
20425.85806087734,
451.9845542917651
],
"size": [
140,
46
],
"flags": {},
"order": 10,
"mode": 2,
"inputs": [
{
"name": "image1",
"type": "IMAGE",
"link": 59
},
{
"name": "image2",
"type": "IMAGE",
"link": 60
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
61
]
}
],
"properties": {
"Node name for S&R": "ImageBatch"
},
"widgets_values": []
},
{
"id": 28,
"type": "JoyAI_Image_SM_Clip",
"pos": [
20369.011635584688,
222.28848630003668
],
"size": [
270,
82
],
"flags": {},
"order": 5,
"mode": 2,
"inputs": [],
"outputs": [
{
"name": "clip",
"type": "CLIP",
"links": [
57
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_SM_Clip"
},
"widgets_values": [
"JoyAI-Image-Und-merger_bf16.safetensors",
"none"
]
},
{
"id": 26,
"type": "JoyAI_Image_Understand",
"pos": [
20670.151078083465,
282.192262702941
],
"size": [
400,
228
],
"flags": {},
"order": 13,
"mode": 2,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 57
},
{
"name": "image",
"type": "IMAGE",
"link": 61
}
],
"outputs": [
{
"name": "response",
"type": "STRING",
"links": [
55
]
}
],
"properties": {
"Node name for S&R": "JoyAI_Image_Understand"
},
"widgets_values": [
"Compare these two images.",
2048,
0.7,
0.8,
50,
"cuda"
]
},
{
"id": 27,
"type": "PreviewAny",
"pos": [
21099.17781626596,
227.7317034140748
],
"size": [
411.49710473692176,
320.16314870956194
],
"flags": {},
"order": 15,
"mode": 2,
"inputs": [
{
"name": "source",
"type": "*",
"link": 55
}
],
"outputs": [],
"properties": {
"Node name for S&R": "PreviewAny"
},
"widgets_values": [
null,
null,
null
]
},
{
"id": 32,
"type": "LoadImage",
"pos": [
20006.860719625834,
578.353819758959
],
"size": [
270,
314.00000000000006
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
60
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"test_3.png",
"image"
]
},
{
"id": 29,
"type": "LoadImage",
"pos": [
20006.25013320333,
159.22249838007153
],
"size": [
275.2101011260056,
358.77371330286906
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
59
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"test_1.jpg",
"image"
]
}
],
"links": [
[
1,
4,
0,
5,
1,
"LATENT"
],
[
3,
1,
0,
5,
0,
"MODEL"
],
[
12,
11,
0,
4,
0,
"IMAGE"
],
[
16,
3,
0,
14,
0,
"CLIP"
],
[
17,
11,
0,
14,
1,
"IMAGE"
],
[
18,
14,
0,
13,
0,
"CONDITIONING"
],
[
23,
14,
0,
5,
2,
"CONDITIONING"
],
[
24,
14,
1,
5,
3,
"CONDITIONING"
],
[
25,
15,
0,
4,
1,
"VAE"
],
[
27,
5,
0,
16,
1,
"LATENT"
],
[
53,
16,
0,
10,
0,
"IMAGE"
],
[
54,
15,
0,
16,
0,
"VAE"
],
[
55,
26,
0,
27,
0,
"STRING"
],
[
57,
28,
0,
26,
0,
"CLIP"
],
[
59,
29,
0,
31,
0,
"IMAGE"
],
[
60,
32,
0,
31,
1,
"IMAGE"
],
[
61,
31,
0,
26,
1,
"IMAGE"
]
],
"groups": [
{
"id": 1,
"title": "Group",
"bounding": [
20360.296397399088,
-693.5307633964666,
423.05741857026806,
762.725858797222
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
},
{
"id": 2,
"title": "Group",
"bounding": [
20800.419992988376,
-696.2643982036351,
707.3995033604333,
775.6013422163574
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
},
{
"id": 3,
"title": "Group",
"bounding": [
20343.9040135276,
150.3523702537851,
1179.091696966756,
432.6842093303753
],
"color": "#3f789e",
"font_size": 24,
"flags": {}
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.5960521358998819,
"offset": [
-19117.54563399457,
694.7222119219581
]
},
"frontendVersion": "1.41.21",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+76
View File
@@ -0,0 +1,76 @@
from dataclasses import dataclass, field
from pathlib import Path
from .src.infer_runtime.infer_config import InferConfig
# def _resolve_root() -> Path:
# here = Path(__file__).resolve().parent
# if (here / "transformer").exists() and (here / "vae").exists() and (here / "JoyAI-Image-Und").exists():
# return here
# raise ValueError(
# "Place this config file directly inside the checkpoint root."
# )
# _ROOT = _resolve_root()
@dataclass
class JoyAIImageInferConfig(InferConfig):
dit_arch_config: dict = field(
default_factory=lambda: {
"target": "modules.models.Transformer3DModel",
"params": {
"hidden_size": 4096,
"in_channels": 16,
"heads_num": 32,
"mm_double_blocks_depth": 40,
"out_channels": 16,
"patch_size": [1, 2, 2],
"rope_dim_list": [16, 56, 56],
"text_states_dim": 4096,
"rope_type": "rope",
"dit_modulation_type": "wanx",
"theta": 10000,
"attn_backend": "flash_attn",
},
}
)
# vae_arch_config: dict = field(
# default_factory=lambda: {
# "target": "modules.models.WanxVAE",
# "params": {
# "pretrained": str(_ROOT / "vae" / "Wan2.1_VAE.pth"),
# },
# }
# )
# text_encoder_arch_config: dict = field(
# default_factory=lambda: {
# "target": "modules.models.load_text_encoder",
# "params": {
# "text_encoder_ckpt": str(_ROOT / "JoyAI-Image-Und"),
# },
# }
# )
scheduler_arch_config: dict = field(
default_factory=lambda: {
"target": "modules.models.FlowMatchDiscreteScheduler",
"params": {
"num_train_timesteps": 1000,
"shift": 4.0,
},
}
)
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precision: str = "bf16"
text_token_max_length: int = 2048
# Keep these fields visible in the active config because they control multi-GPU inference.
hsdp_shard_dim: int = 1
reshard_after_forward: bool = False
use_fsdp_inference: bool = False
cpu_offload: bool = False
pin_cpu_memory: bool = False
+303
View File
@@ -0,0 +1,303 @@
"""Local inference entrypoint for the clean JoyAI-Image release."""
from __future__ import annotations
import os
import time
import numpy as np
import torch
from PIL import Image
from diffusers.utils.torch_utils import randn_tensor
from einops import rearrange
from .src.infer_runtime.model import InferenceParams, build_model
from .src.infer_runtime.settings import InferSettings
from .src.modules.models.mmdit.vae import WanxVAE
from .model_loader_utils import map_0_1_to_neg1_1,map_neg1_1_to_0_1
cur_path = os.path.dirname(os.path.abspath(__file__))
def load_vae(vae_path,device,dtype ):
vae=WanxVAE(vae_path,dtype,device)
return vae
joy_ai_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
]
joy_ai_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.9160
]
def vae_decode(vae,latents):
latents=latents["samples"] if isinstance(latents, dict) else latents
if isinstance(vae, WanxVAE):
with torch.autocast(device_type="cuda", dtype=vae.dtype, enabled=True):
image = vae.decode(latents, return_dict=False)[0]
#print(f" {image.shape}, dtype: {image.dtype}, device: {image.device}")
image = rearrange(image, "(b n) c f h w -> b n c f h w", b=1)
image = (image / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
image = image.cpu().float().permute(0, 1, 3, 2, 4, 5)
# image_tensor = (image[0, -1, 0] * 255).to(torch.uint8).cpu()
# img= Image.fromarray(image_tensor.permute(1, 2, 0).numpy())
# img.save(os.path.join(cur_path, 'decoded_image_12.png'))
image = image[0, -1, 0] # (c, f, h, w)
#print(image.shape) # torch.Size([1, 3, 1024, 1024])
image = image.cpu().float().unsqueeze(0).permute(0, 2, 3, 1)
#print(image.shape) # torch.Size([1, 1024, 1024, 3])
else:
mean = torch.tensor(joy_ai_mean, dtype=latents.dtype, device=latents.device)
std = torch.tensor(joy_ai_std, dtype=latents.dtype, device=latents.device)
scale = [mean, 1.0 / std]
latents = latents / scale[1].view(1, 16, 1, 1, 1) + scale[0].view(1, 16, 1, 1, 1)
#latents=map_neg1_1_to_0_1(latents)
image=vae.decode(latents) ##Decoded image shape: torch.Size([2, 1, 1024, 1024, 3]), dtype: torch.float32, device: cpu
image= rearrange(image, "f b h w c -> (f b) h w c")
#print(f"Decoded image shape: {image.shape}, dtype: {image.dtype}, device: {image.device}")
return image
def prepare_conditions( latents, image=None, last_image=None,vae=None):
"""
Prepare conditional inputs for video generation.
Args:
latents: Generated latent tensor with shape (B, N, C, T, latent_H, latent_W)
image: First frame condition, shape (B, N, 3, 1, H, W)
last_image: Last frame condition, shape (B, N, 3, 1, H, W)
Returns:
Combined condition tensor with shape (B, N, C+1, T, H, W)
"""
device, dtype = latents.device, latents.dtype
batch_size, num_items, latent_channels, latent_frames, latent_h, latent_w = latents.shape
# If no conditions provided, return zero condition
if image is None and last_image is None:
return torch.zeros(
batch_size, num_items, latent_channels + 1, latent_frames, latent_h, latent_w,
device=device, dtype=dtype
)
num_frame = (latent_frames - 1) * 4 + 1
height = latent_h * 8
width = latent_w * 8
# Initialize mask
mask = torch.zeros(batch_size, num_items, 1, latent_frames,
latent_h, latent_w, device=device, dtype=dtype)
# Build video condition
if image is not None and last_image is not None:
# Both first and last frame conditions
image = image.to(device=device, dtype=dtype)
last_image = last_image.to(device=device, dtype=dtype)
middle_frames = torch.zeros(
batch_size, num_items, image.shape[2], num_frame -
2, height, width,
device=device, dtype=dtype
)
video_condition = torch.cat(
[image, middle_frames, last_image], dim=3)
mask[:, :, :, 0] = 1 # Mark first frame as conditional
mask[:, :, :, -1] = 1 # Mark last frame as conditional
elif image is not None:
# Only first frame condition
image = image.to(device=device, dtype=dtype)
remaining_frames = torch.zeros(
batch_size, num_items, image.shape[2], num_frame -
1, height, width,
device=device, dtype=dtype
)
video_condition = torch.cat([image, remaining_frames], dim=3)
mask[:, :, :, 0] = 1 # Mark first frame as conditional
else:
raise NotImplementedError
# VAE encode the video condition
video_condition = rearrange(
video_condition, "b n c t h w -> (b n) c t h w")
latent_condition = vae.encode(
video_condition).latent_dist.sample()
# Normalize
normalize_latents=lambda x: x * (2.0 / x.shape[-1]) # TODO is's not right,just for test
latent_condition = normalize_latents(latent_condition)
# Reshape back to (B, N, C, T, H, W)
latent_condition = rearrange(
latent_condition, "(b n) c t h w -> b n c t h w", b=batch_size)
# Concat
return torch.cat([latent_condition, mask], dim=2)
def prepare_latents_(
batch_size,
num_items,
num_channels_latents,
height,
width,
video_length,
dtype,
device,
generator,
latents=None,
reference_images=None,
image=None,
last_image=None,
vae= None,
image_tensor=None
):
shape = (
batch_size,
num_items,
num_channels_latents,
(video_length - 1) // 4 + 1,
int(height) // 8,
int(width) // 8,
)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
if latents is None:
if reference_images is not None:
ref_img = [torch.from_numpy(
np.array(x.convert("RGB"))) for x in reference_images]
ref_img = torch.stack(ref_img).to(device=device, dtype=dtype)
ref_img = ref_img / 127.5 - 1.0
ref_img = rearrange(ref_img, "x h w c -> x c 1 h w")
if isinstance(vae, WanxVAE):
ref_vae = vae.encode(ref_img) #(torch.Size([1, 16, 1, 128, 128]), torch.float32, True)
else:
mean = torch.tensor(joy_ai_mean, dtype=dtype, device=device)
std = torch.tensor(joy_ai_std, dtype=dtype, device=device)
scale = [mean, 1.0 / std]
ref_vae = vae.encode(image_tensor).to(device=device, dtype=dtype) # (torch.Size([1, 16, 1, 128, 128]), torch.bfloat16, True)
ref_vae=map_0_1_to_neg1_1(ref_vae) # comfyUI 0.1 to -1.1
ref_vae = (ref_vae - scale[0].view(1, 16, 1, 1, 1)) * scale[1].view(1, 16, 1, 1, 1)
#print(f"Reference VAE shape: {ref_vae.shape,ref_vae.dtype,ref_vae.is_cuda}")
ref_vae = rearrange(
ref_vae, "(b n) c 1 h w -> b n c 1 h w", n=(num_items - 1))
#print(f"Reference VAE reshaped: {ref_vae.shape}") # torch.Size([1, 1, 16, 1, 128, 128])
noise = randn_tensor(
(shape[0], 1, *shape[2:]),
generator=generator, device=device, dtype=dtype
)
latents = torch.cat([ref_vae, noise], dim=1)
else:
latents = randn_tensor(
shape, generator=generator, device=device, dtype=dtype
)
else:
latents = latents.to(device)
enable_multi_task=False
if not enable_multi_task:
return latents, None
# image: (b, n, c, 1, h, w), last_image: (b, n, c, 1, h, w)
condition = prepare_conditions(latents, image, last_image, vae)
return latents, condition
def get_latents(vae, images, height, width, device,seed,image_tensor, dtype):
num_items = 1 if images is None or len(
images) == 0 else 1 + len(images)
num_channels_latents =16
num_frames = 1
generator = torch.Generator(device='cuda').manual_seed(int(seed))
latents, condition = prepare_latents_(
1,
num_items,
num_channels_latents,
height,
width,
num_frames,
dtype,
device,
generator,
reference_images=images,
vae=vae,
image_tensor=image_tensor
)
return latents, condition
def load_input_image(image_path: str | None) -> Image.Image | None:
if not image_path:
return None
return Image.open(image_path).convert('RGB')
def is_rank0() -> bool:
return int(os.environ.get('RANK', '0')) == 0
def resolve_device() -> torch.device:
if not torch.cuda.is_available():
return torch.device('cpu')
local_rank = int(os.environ.get('LOCAL_RANK', '0'))
torch.cuda.set_device(local_rank)
return torch.device(f'cuda:{local_rank}')
def load_mmdit(dit_path,gguf_path,offload):
settings = InferSettings(
config_path=os.path.join(cur_path, 'infer_config.py') ,
ckpt_path=dit_path or gguf_path,
rewrite_model=None ,#'gpt-5'
openai_api_key=os.environ.get('OPENAI_API_KEY', None),
openai_base_url=os.environ.get('OPENAI_BASE_URL', None),
default_seed=42,
repo_path=os.path.join(cur_path, 'JoyAI-Image-Und'),
)
device = resolve_device() if not offload else torch.device('cpu')
model = build_model(
settings,
device=device,
hsdp_shard_dim_override=False,
)
return model
def infer_joyai(model,lat,positive,negative, steps, guidance_scale,offload,offload_block_num):
start_time = time.time()
output_image = model.infer(
images=lat.get("images"),
height=lat["height"],
width=lat["width"],
steps=steps,
guidance_scale=guidance_scale,
prompt_embeds=positive[0][0] ,
prompt_embeds_mask=positive[0][1]['prompt_attention_mask'] ,
negative_prompt_embeds=negative[0][0] ,
negative_prompt_embeds_mask=negative[0][1]['prompt_attention_mask'] if negative else None,
offload=offload,
offload_block_num=offload_block_num,
lat=lat["samples"]
)
elapsed = time.time() - start_time
print(f'Time taken: {elapsed:.2f} seconds')
return output_image
+484
View File
@@ -0,0 +1,484 @@
"""Local inference entrypoint for the image understanding capability of JoyAI-Image."""
from __future__ import annotations
import argparse
from .src.modules.utils import _dynamic_resize_from_bucket
import time
import warnings
from pathlib import Path
from transformers import Qwen3VLForConditionalGeneration, AutoProcessor,Qwen3VLConfig
import torch
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
from transformers import AutoTokenizer
import os
from contextlib import nullcontext
from accelerate import init_empty_weights
from diffusers.utils import is_accelerate_available
from safetensors.torch import load_file
import gc
cur_path = os.path.dirname(os.path.abspath(__file__))
# ROOT_DIR = Path(__file__).resolve().parent
# SRC_DIR = ROOT_DIR / "src"
# if str(SRC_DIR) not in sys.path:
# sys.path.insert(0, str(SRC_DIR))
from PIL import Image
warnings.filterwarnings("ignore")
def load_images(image_arg: str) -> list[Image.Image]:
paths = [p.strip() for p in image_arg.split(",")]
images = []
for p in paths:
if not Path(p).is_file():
raise FileNotFoundError(f"Image not found: {p}")
images.append(Image.open(p).convert("RGB"))
return images
def resolve_text_encoder_path(ckpt_root: str) -> Path:
root = Path(ckpt_root).expanduser().resolve()
text_encoder_dir = root / "JoyAI-Image-Und"
if not text_encoder_dir.is_dir():
raise FileNotFoundError(
f"Expected text_encoder/ directory inside checkpoint root: {root}"
)
return text_encoder_dir
def build_conversation(
images: list[Image.Image],
prompt: str | None,
) -> list[dict]:
SYS_PROMPT = "You are a helpful assistant."
default_prompt = "Describe this image in detail."
user_text = prompt if prompt is not None else default_prompt
image_content = [{"type": "image", "image": img} for img in images]
user_content = image_content + [{"type": "text", "text": user_text}]
messages = [
{"role": "system", "content": SYS_PROMPT},
{"role": "user", "content": user_content},
]
return messages
joy_prompt_template_encode = {
'image': "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n",
'multiple_images': "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n{}<|im_start|>assistant\n",
'video': "<|im_start|>system\n \\nDescribe the video by detailing the following aspects:\n1. The main content and theme of the video.\n2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects.\n3. Actions, events, behaviors temporal relationships, physical movement changes of the objects.\n4. background environment, light, style and atmosphere.\n5. camera angles, movements, and transitions used in the video:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
}
joy_prompt_template_encode_start_idx = {
'image': 34,
'multiple_images': 34,
'video': 91,
}
def extract_masked_hidden( hidden_states: torch.Tensor, mask: torch.Tensor):
bool_mask = mask.bool()
valid_lengths = bool_mask.sum(dim=1)
selected = hidden_states[bool_mask]
split_result = torch.split(selected, valid_lengths.tolist(), dim=0)
return split_result
def get_qwen_prompt_embeds(text_encoder,tokenizer,
prompt: Union[str, List[str]] = None,
template_type: str = 'image',
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
template = joy_prompt_template_encode[template_type]
drop_idx = joy_prompt_template_encode_start_idx[template_type]
txt = [template.format(e) for e in prompt]
txt_tokens = tokenizer(
txt, max_length=2048 + drop_idx, padding=True, truncation=True, return_tensors="pt"
).to(device)
encoder_hidden_states = text_encoder(
input_ids=txt_tokens.input_ids,
attention_mask=txt_tokens.attention_mask,
output_hidden_states=True,
)
hidden_states = encoder_hidden_states.hidden_states[-1]
split_hidden_states = extract_masked_hidden(
hidden_states, txt_tokens.attention_mask)
split_hidden_states = [e[drop_idx:] for e in split_hidden_states]
attn_mask_list = [torch.ones(
e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states]
max_seq_len = min([
2048,
max([u.size(0) for u in split_hidden_states]),
max([u.size(0) for u in attn_mask_list])])
prompt_embeds = torch.stack(
[torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))])
for u in split_hidden_states]
)
encoder_attention_mask = torch.stack(
[torch.cat([u, u.new_zeros(max_seq_len - u.size(0))])
for u in attn_mask_list]
)
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
return prompt_embeds, encoder_attention_mask
def encode_prompt_multiple_images(
text_encoder,qwen_processor,
prompt: Union[str, List[str]],
device: Optional[torch.device] = None,
images: Optional[torch.Tensor] = None,
template_type: Optional[str] = 'multiple_images',
max_sequence_length: Optional[int] = None,
drop_vit_feature: Optional[float] = False,
):
assert template_type == 'multiple_images', "template_type must be 'multiple_images'"
device = device
template = joy_prompt_template_encode[template_type]
drop_idx = joy_prompt_template_encode_start_idx[template_type]
prompt = [p.replace('<image>\n', '<|vision_start|><|image_pad|><|vision_end|>') for p in prompt]
prompt = [template.format(p) for p in prompt]
inputs = qwen_processor(
text=prompt,
images=images,
padding=True,
return_tensors="pt",
).to(device)
encoder_hidden_states = text_encoder(
**inputs,
output_hidden_states=True,
)
last_hidden_states = encoder_hidden_states.hidden_states[-1]
if drop_vit_feature:
input_ids = inputs['input_ids']
vlm_image_end_idx = torch.where(input_ids[0] == 151653)[0][-1]
drop_idx = vlm_image_end_idx + 1
prompt_embeds = last_hidden_states[:, drop_idx:]
prompt_embeds_mask = inputs['attention_mask'][:, drop_idx:]
if max_sequence_length is not None and prompt_embeds.shape[1] > max_sequence_length:
prompt_embeds = prompt_embeds[:, -max_sequence_length:, :]
prompt_embeds_mask = prompt_embeds_mask[:, -max_sequence_length:]
return prompt_embeds, prompt_embeds_mask
def encode_prompt(
text_encoder,
prompt: Union[str, List[str]],
images: Optional[torch.Tensor] = None,
device: Optional[torch.device] = None,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_embeds_mask: Optional[torch.Tensor] = None,
max_sequence_length: int = 4096,
template_type: str = 'image',
drop_vit_feature: bool = False,
):
r"""
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
device: (`torch.device`):
torch device
num_videos_per_prompt (`int`):
number of videos that should be generated per prompt
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
"""
qwen_processor = AutoProcessor.from_pretrained(os.path.join(cur_path, 'JoyAI-Image-Und'))
tokenizer = AutoTokenizer.from_pretrained(
os.path.join(cur_path, 'JoyAI-Image-Und'),
local_files_only=True,
trust_remote_code=True,
)
if images is not None:
##################################################
# from PIL import Image
# images = [img.resize((512, 512), Image.LANCZOS) for img in images]
##################################################
return encode_prompt_multiple_images(text_encoder,qwen_processor,
prompt=prompt,
images=images,
device=device,
max_sequence_length=max_sequence_length,
drop_vit_feature=drop_vit_feature,
)
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(
prompt) if prompt_embeds is None else prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds, prompt_embeds_mask = get_qwen_prompt_embeds(text_encoder,tokenizer,
prompt, template_type, device)
prompt_embeds = prompt_embeds[:, :max_sequence_length]
prompt_embeds_mask = prompt_embeds_mask[:, :max_sequence_length]
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
batch_size * num_videos_per_prompt, seq_len, -1)
prompt_embeds_mask = prompt_embeds_mask.repeat(
1, num_videos_per_prompt, 1)
prompt_embeds_mask = prompt_embeds_mask.view(
batch_size * num_videos_per_prompt, seq_len)
return prompt_embeds, prompt_embeds_mask
def load_qwen3vl_model(safetensors_path,gguf_path ) -> torch.nn.Module:
#text_encoder_path = resolve_text_encoder_path(text_encoder_path)
#device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
ctx = init_empty_weights if is_accelerate_available() else nullcontext
device = torch.device('cpu')
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
print(f"Loading MLLM from: {safetensors_path}")
print(f"Device: {device}, dtype: {dtype}")
configs=Qwen3VLConfig.from_pretrained(os.path.join(cur_path, 'JoyAI-Image-Und'),local_files_only=True,model_type=dtype,trust_remote_code=True,low_cpu_mem_usage=False)
with ctx():
model = Qwen3VLForConditionalGeneration(configs)
if safetensors_path is not None :
model_dict=load_file(safetensors_path)
match_state_dict(model, model_dict,show_num=20)
x,y=model.load_state_dict(model_dict,strict=False,assign=True)
print(x,"########_missing \n")
print(y,"########_unused \n")
del model_dict
gc.collect()
model.eval().to(dtype)
elif gguf_path is not None:
g_dict=load_gguf_checkpoint(gguf_path)
match_state_dict(model, g_dict,show_num=20)
set_gguf2meta_model(model,g_dict,dtype,torch.device("cpu"))
del g_dict
gc.collect()
else:
raise ValueError(
"Please provide either a safetensors_path or a gguf_path."
)
return model
def encoder_input(model,prompt,images,max_new_tokens,top_p,top_k,temperature,infer_device):
device = torch.device(infer_device)
processor = AutoProcessor.from_pretrained(
os.path.join(cur_path, 'JoyAI-Image-Und'),
local_files_only=True,
trust_remote_code=True,
)
#images = load_images(args.image)
messages = build_conversation(images, prompt)
text_input = processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True,
)
inputs = processor(
text=[text_input],
images=images,
padding=True,
return_tensors="pt",
).to(device)
print(f"Input tokens: {inputs['input_ids'].shape[1]}")
print("Generating...")
start_time = time.time()
generate_kwargs = dict(
max_new_tokens=max_new_tokens,
)
if temperature == 0:
generate_kwargs["do_sample"] = False
else:
generate_kwargs["do_sample"] = True
generate_kwargs["temperature"] = temperature
generate_kwargs["top_p"] = top_p
generate_kwargs["top_k"] = top_k
if infer_device=="cuda":
model.to("cuda")
with torch.no_grad():
output_ids = model.generate(**inputs, **generate_kwargs)
if infer_device=="cuda":
model.to("cpu")
# Strip the input tokens from the output
generated_ids = output_ids[:, inputs["input_ids"].shape[1]:]
response = processor.batch_decode(
generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False,
)[0]
elapsed = time.time() - start_time
num_output_tokens = generated_ids.shape[1]
print(f"\n{'=' * 60}")
print(f"Response:\n{response}")
print(f"{'=' * 60}")
print(f"Output tokens: {num_output_tokens}")
print(f"Time: {elapsed:.2f}s ({num_output_tokens / elapsed:.1f} tok/s)")
# if save_response:
# output_path = Path(output)
# output_path.parent.mkdir(parents=True, exist_ok=True)
# output_path.write_text(response, encoding="utf-8")
# print(f"Saved response to: {output_path}")
return response
def get_conditioning(clip,prompt, images,infer_device):
device = torch.device(infer_device)
num_items = 1 if images is None or len(
images) == 0 else 1 + len(images)
default_negative_prompt = ""
if images is None:
prompts = [f"<|im_start|>user\n{prompt}<|im_end|>\n"]
negative_prompt = [f"<|im_start|>user\n{default_negative_prompt}<|im_end|>\n"]
else:
images = _dynamic_resize_from_bucket(images[0], basesize=1024)
images.save("temp_image.png")
width, height = images.size
image_tokens = '<image>\n'
prompts = [f"<|im_start|>user\n{image_tokens}{prompt}<|im_end|>\n"]
negative_prompt = [f"<|im_start|>user\n{image_tokens}{default_negative_prompt}<|im_end|>\n"]
# if num_items <= 1:
# negative_prompt = [
# f"<|im_start|>user\n{default_negative_prompt}<|im_end|>\n"] * 1
# else:
# image_tokens = "<image>\n" * (num_items - 1)
# negative_prompt = [
# f"<|im_start|>user\n{image_tokens}{default_negative_prompt}<|im_end|>\n"] * 1
if infer_device=="cuda":
clip.to("cuda")
prompt_embeds, prompt_embeds_mask=encode_prompt(clip,prompts,images,device)
#print(prompt_embeds.shape, prompt_embeds_mask.shape) #torch.Size([1, 1037, 4096]) torch.Size([1, 1037])
n_prompt_embeds, n_prompt_embeds_mask=encode_prompt(clip,negative_prompt,images,device)
if infer_device=="cuda":
clip.to("cpu")
torch.cuda.empty_cache()
positive=[[prompt_embeds,{"prompt_attention_mask": prompt_embeds_mask}]]
negative=[[n_prompt_embeds,{"prompt_attention_mask": n_prompt_embeds_mask}]]
return positive,negative
def load_gguf_checkpoint(gguf_checkpoint_path):
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
from diffusers.utils import is_gguf_available, is_torch_available
if is_gguf_available() and is_torch_available():
import gguf
from gguf import GGUFReader
from diffusers.quantizers.gguf.utils import SUPPORTED_GGUF_QUANT_TYPES, GGUFParameter
else:
logger.error(
"Loading a GGUF checkpoint in PyTorch, requires both PyTorch and GGUF>=0.10.0 to be installed. Please see "
"https://pytorch.org/ and https://github.com/ggerganov/llama.cpp/tree/master/gguf-py for installation instructions."
)
raise ImportError("Please install torch and gguf>=0.10.0 to load a GGUF checkpoint in PyTorch.")
reader = GGUFReader(gguf_checkpoint_path)
parsed_parameters = {}
for i, tensor in enumerate(reader.tensors):
name = tensor.name
quant_type = tensor.tensor_type
is_gguf_quant = quant_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
if is_gguf_quant and quant_type not in SUPPORTED_GGUF_QUANT_TYPES:
_supported_quants_str = "\n".join([str(type) for type in SUPPORTED_GGUF_QUANT_TYPES])
raise ValueError(
(
f"{name} has a quantization type: {str(quant_type)} which is unsupported."
"\n\nCurrently the following quantization types are supported: \n\n"
f"{_supported_quants_str}"
"\n\nTo request support for this quantization type please open an issue here: https://github.com/huggingface/diffusers"
)
)
weights = torch.from_numpy(tensor.data) #tensor.data.copy()
parsed_parameters[name.replace("model.", "")] = GGUFParameter(weights, quant_type=quant_type) if is_gguf_quant else weights
del tensor,weights
if i > 0 and i % 1000 == 0: # 每1000个tensor执行一次gc
logger.info(f"Processed {i}tensors...")
gc.collect()
del reader
gc.collect()
return parsed_parameters
def set_gguf2meta_model(meta_model,model_state_dict,dtype,device):
from diffusers import GGUFQuantizationConfig
from diffusers.quantizers.gguf import GGUFQuantizer
g_config = GGUFQuantizationConfig(compute_dtype=dtype or torch.bfloat16)
hf_quantizer = GGUFQuantizer(quantization_config=g_config)
hf_quantizer.pre_quantized = True
hf_quantizer._process_model_before_weight_loading(
meta_model,
device_map={"": device} if device else None,
state_dict=model_state_dict
)
from diffusers.models.model_loading_utils import load_model_dict_into_meta
x,y=load_model_dict_into_meta(
meta_model,
model_state_dict,
hf_quantizer=hf_quantizer,
device_map={"": device} if device else None,
dtype=dtype
)
print(x,"offload_index")
print(y,"state_dict_index")
hf_quantizer._process_model_after_weight_loading(meta_model)
del model_state_dict
gc.collect()
return meta_model.to(dtype=dtype)
def match_state_dict(meta_model, sd,show_num=10):
meta_model_keys = set(meta_model.state_dict().keys())
state_dict_keys = set(sd.keys())
# 打印匹配的键的数量
matching_keys = meta_model_keys.intersection(state_dict_keys)
print(f"Matching keys count: {len(matching_keys)}")
# 打印不在 meta_model 中但在 state_dict 中的键(多余键)
extra_keys = state_dict_keys - meta_model_keys
if extra_keys:
print(f"Extra keys in state_dict (not in meta_model): {len(extra_keys)}")
for key in list(extra_keys)[:show_num]: # 只显示前10个
print(f" - {key}")
# 打印不在 state_dict 中但在 meta_model 中的键(缺失键)
missing_keys = meta_model_keys - state_dict_keys
if missing_keys:
print(f"Missing keys in state_dict (not in state_dict): {len(missing_keys)}")
for key in list(missing_keys)[:show_num]: # 只显示前10个
print(f" - {key}")
# 如果需要,也可以打印部分匹配的键
print(f"Sample matching keys: {list(matching_keys)[:5]}")
+165
View File
@@ -0,0 +1,165 @@
# !/usr/bin/env python
# -*- coding: UTF-8 -*-
import os
import torch
import gc
import comfy.model_management as mm
from PIL import Image
import numpy as np
from comfy.utils import common_upscale
import folder_paths
import time
cur_path = os.path.dirname(os.path.abspath(__file__))
def clear_comfyui_cache():
cf_models=mm.loaded_models()
try:
for pipe in cf_models:
pipe.unpatch_model(device_to=torch.device("cpu"))
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 gc_cleanup():
gc.collect()
torch.cuda.empty_cache()
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 read_lat_emb(prefix, device):
if prefix =="embeds":
if not os.path.exists(os.path.join(folder_paths.get_output_directory(),"raw_embeds_JOY_sm.pt")):
raise Exception("No backup prompt embeddings found. Please run JOY_SM_ENCODER node first.")
else:
prompt_embeds=torch.load(os.path.join(folder_paths.get_output_directory(),"raw_embeds_JOY_sm.pt"),weights_only=False)
if os.path.exists(os.path.join(folder_paths.get_output_directory(),"n_raw_embeds_JOY_sm.pt")):
negative_prompt_embeds=torch.load(os.path.join(folder_paths.get_output_directory(),"n_raw_embeds_JOY_sm.pt"),weights_only=False)
else:
negative_prompt_embeds=[[torch.zeros_like(prompt_embeds[0][0]),prompt_embeds[0][1]]]
#print("Loaded backup prompt embeddings",prompt_embeds[0][0].shape) # Loaded backup prompt embeddings torch.Size([1, 640, 3584])
positive=[[prompt_embeds[0][0].to(device,torch.bfloat16),prompt_embeds[0][1]]]
negative=[[negative_prompt_embeds[0][0].to(device,torch.bfloat16),negative_prompt_embeds[0][1]]]
return positive,negative
elif prefix =="latents":
if not os.path.exists(os.path.join(folder_paths.get_output_directory(),"raw_latents_JOY_sm.pt")) or not os.path.exists(os.path.join(folder_paths.get_output_directory(),"raw_audio_latents_JOY_sm.pt")):
raise Exception("No backup latents found. Please run JOY_SM_KSampler node first.")
else:
video_latents=torch.load(os.path.join(folder_paths.get_output_directory(),"raw_latents_JOY_sm.pt"),weights_only=False)
video_latents["samples"]=video_latents["samples"].to(device,torch.bfloat16)
print(f"video shape: {video_latents['samples'].shape}")
return video_latents, None
def save_lat_emb(save_prefix,data1,data2,mode=""):
data1_prefix, data2_prefix = ("raw_embeds_JOY", "n_raw_embeds_JOY") if save_prefix == "embeds" else ("raw_latents_JOY", "raw_audio_latents_JOY")
default_data1_path = os.path.join(folder_paths.get_output_directory(),f"{data1_prefix}_sm.pt")
default_data2_path = os.path.join(folder_paths.get_output_directory(),f"{data2_prefix}_sm.pt")
prefix = mode+str(int(time.time()))
if os.path.exists(default_data1_path): # use a different path if the file already exists
default_data1_path=os.path.join(folder_paths.get_output_directory(),f"{data1_prefix}_sm_{prefix}.pt")
torch.save(data1,default_data1_path)
if data2 is not None:
if os.path.exists(default_data2_path):
default_data2_path=os.path.join(folder_paths.get_output_directory(),f"{data2_prefix}_sm_{prefix}.pt")
torch.save(data2,default_data2_path)
def map_0_1_to_neg1_1(t):
"""
接受 torch.Tensor 或可转 torch.Tensor 的输入。
可处理形状: H,W,C 或 B,H,W,C 或 B,T,H,W,C(会按元素处理)。
确保 float,0..255 -> 0..1,再把 0..1 -> -1..1(如果已经在 -1..1 则不变)。
"""
if not torch.is_tensor(t):
t = torch.tensor(t)
t = t.float()
# 处理 0..255 的情况
try:
vmax = float(t.max())
except Exception:
vmax = 1.0
if vmax > 2.0:
t = t / 255.0
# 若当前处于 0..1 范围,则映射到 -1..1
try:
vmin = float(t.min())
vmax = float(t.max())
except Exception:
vmin, vmax = -1.0, 1.0
if vmin >= 0.0 and vmax <= 1.1:
t = t * 2.0 - 1.0
return t
def map_neg1_1_to_0_1(t):
"""
接受 torch.Tensor 或可转 torch.Tensor 的输入。
可处理形状: H,W,C 或 B,H,W,C 或 B,T,H,W,C(会按元素处理)。
返回 float tensor,范围 0..1。
"""
if not torch.is_tensor(t):
t = torch.tensor(t)
t = t.float()
# map -1..1 -> 0..1
t = (t + 1.0) * 0.5
# 限幅到 [0,1]
t = t.clamp(0.0, 1.0)
# 保持在 cpu 端,调用方可决定是否转 device/dtype
return t
+16
View File
@@ -0,0 +1,16 @@
[project]
name = "joyai_image"
description = ""
version = "1.0.0"
license = {file = "LICENSE"}
dependencies = ["accelerate", "diffusers>=0.34.0", "einops", "fastapi", "flash-attn>=2.8.0", "loguru", "openai", "packaging", "pillow", "pydantic>=2", "requests", "safetensors", "sentencepiece", "torch", "torchvision", "transformers>=4.57.0,<4.58.0", "uvicorn"]
[project.urls]
Repository = "https://github.com/smthemex/ComfyUI_JoyAI_Image"
# Used by Comfy Registry https://registry.comfy.org
[tool.comfy]
PublisherId = "smthemex"
DisplayName = "ComfyUI_JoyAI_Image"
Icon = ""
includes = []
+18
View File
@@ -0,0 +1,18 @@
accelerate
diffusers>=0.34.0
einops
fastapi
flash-attn>=2.8.0
loguru
openai
packaging
pillow
pydantic>=2
requests
safetensors
sentencepiece
torch
torchvision
transformers>=4.57.0,<4.58.0
uvicorn
+1
View File
@@ -0,0 +1 @@
__all__: list[str] = []
+63
View File
@@ -0,0 +1,63 @@
from __future__ import annotations
from dataclasses import dataclass
import json
from pathlib import Path
@dataclass(frozen=True)
class CheckpointLayout:
root: Path
transformer_ckpt: Path
vae_ckpt: Path
text_encoder_ckpt: Path
def _must_exist(path: Path, kind: str) -> Path:
if not (path.exists() or path.is_symlink()):
raise FileNotFoundError(f"Missing {kind}: {path}")
return path
def _find_single_entry(directory: Path, kind: str, *, expect_dir: bool) -> Path:
_must_exist(directory, f"{kind} directory")
entries = sorted(p for p in directory.iterdir() if not p.name.startswith("."))
if len(entries) != 1:
raise FileNotFoundError(f"Expected exactly one entry in {directory} for {kind}, found {len(entries)}")
entry = entries[0]
if expect_dir:
if not (entry.is_dir() or entry.is_symlink()):
raise FileNotFoundError(f"Expected directory-like entry for {kind}: {entry}")
else:
if not (entry.is_file() or entry.is_symlink()):
raise FileNotFoundError(f"Expected file-like entry for {kind}: {entry}")
return entry
def resolve_checkpoint_layout(root: str | Path) -> CheckpointLayout:
root_path = Path(root).expanduser().resolve()
transformer_ckpt = str(root_path / "transformer" / "transformer.pth")
vae_ckpt = _find_single_entry(root_path / "vae", "vae checkpoint", expect_dir=False)
text_encoder_ckpt = _must_exist(root_path / "JoyAI-Image-Und", "text encoder checkpoint directory")
if not text_encoder_ckpt.is_dir():
raise FileNotFoundError(f"Expected text encoder checkpoint directory: {text_encoder_ckpt}")
return CheckpointLayout(
root=root_path,
transformer_ckpt=transformer_ckpt,
vae_ckpt=vae_ckpt,
text_encoder_ckpt=text_encoder_ckpt,
)
def build_manifest(layout: CheckpointLayout) -> dict[str, str]:
return {
"root": str(layout.root),
"transformer_ckpt": str(layout.transformer_ckpt),
"vae_ckpt": str(layout.vae_ckpt),
"text_encoder_ckpt": str(layout.text_encoder_ckpt),
}
def write_manifest(layout: CheckpointLayout, output_path: str | Path) -> Path:
output = Path(output_path)
output.write_text(json.dumps(build_manifest(layout), indent=2) + "\n", encoding="utf-8")
return output
+53
View File
@@ -0,0 +1,53 @@
from __future__ import annotations
from dataclasses import dataclass
import importlib.util
import inspect
from pathlib import Path
from typing import Any, Type
@dataclass
class InferConfig:
dit_ckpt: str | None = None
dit_ckpt_type: str = "pt"
dit_arch_config: dict[str, Any] | None = None
dit_precision: str = "bf16"
vae_arch_config: dict[str, Any] | None = None
vae_precision: str = "bf16"
text_encoder_arch_config: dict[str, Any] | None = None
text_encoder_precision: str = "bf16"
text_token_max_length: int = 2048
scheduler_arch_config: dict[str, Any] | None = None
training_mode: bool = False
hsdp_shard_dim: int = 1
reshard_after_forward: bool = False
use_fsdp_inference: bool = False
cpu_offload: bool = False
pin_cpu_memory: bool = False
def load_infer_config_class_from_pyfile(file_path: str) -> Type[Any]:
path = Path(file_path)
if not path.is_file():
raise FileNotFoundError(f"Configuration file not found: {file_path}")
module_name = path.stem
spec = importlib.util.spec_from_file_location(module_name, file_path)
if spec is None or spec.loader is None:
raise ImportError(f"Could not create module spec for '{file_path}'.")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
for _, obj in inspect.getmembers(module, inspect.isclass):
if obj is InferConfig:
continue
if issubclass(obj, InferConfig):
return obj
raise ValueError(f"No class inheriting from 'InferConfig' was found in '{file_path}'.")
+202
View File
@@ -0,0 +1,202 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional
import os
from PIL import Image
import torch
from .infer_config import InferConfig, load_infer_config_class_from_pyfile
from .prompt_rewrite import rewrite_prompt
from .settings import InferSettings
from ..modules.models import load_dit, load_pipeline
from ..modules.utils import _dynamic_resize_from_bucket, seed_everything
from ..modules.models import Transformer3DModel,FlowMatchDiscreteScheduler
from dataclasses import dataclass, field
@dataclass
class JoyAIImageInferConfig:
dit_arch_config: dict = field(
default_factory=lambda: {
"target": Transformer3DModel,
"params": {
"hidden_size": 4096,
"in_channels": 16,
"heads_num": 32,
"mm_double_blocks_depth": 40,
"out_channels": 16,
"patch_size": [1, 2, 2],
"rope_dim_list": [16, 56, 56],
"text_states_dim": 4096,
"rope_type": "rope",
"dit_modulation_type": "wanx",
"theta": 10000,
"attn_backend": "flash_attn",
},
}
)
vae_arch_config: dict = field(
default_factory=lambda: {
"target": "modules.models.WanxVAE",
"params": {
"pretrained": "vae/Wan2.1_VAE.pth", # 相对于checkpoint根目录的路径
},
}
)
text_encoder_arch_config: dict = field(
default_factory=lambda: {
"target": "modules.models.load_text_encoder",
"params": {
"text_encoder_ckpt": "JoyAI-Image-Und", # 相对于checkpoint根目录的路径
},
}
)
scheduler_arch_config: dict = field(
default_factory=lambda: {
"target": FlowMatchDiscreteScheduler,
"params": {
"num_train_timesteps": 1000,
"shift": 4.0,
},
}
)
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precision: str = "bf16"
text_token_max_length: int = 2048
hsdp_shard_dim: int = 1
reshard_after_forward: bool = False
use_fsdp_inference: bool = False
cpu_offload: bool = False
pin_cpu_memory: bool = False
dit_ckpt: str = ""
training_mode: bool = False
@dataclass
class InferenceParams:
prompt: str
image: Optional[Image.Image]
height: int
width: int
steps: int
guidance_scale: float
seed: int
neg_prompt: str
basesize: int
latents: Optional[torch.Tensor] = None
prompt_embeds: Optional[torch.Tensor] = None
negative_prompt_embeds_mask: Optional[torch.Tensor] = None
neg_prompt_embeds: Optional[torch.Tensor] = None
prompt_embeds_mask: Optional[torch.Tensor] = None
negative_prompt_embeds_mask: Optional[torch.Tensor] = None
offload: bool = False
offload_block_num: int = 1
class EditModel:
def __init__(
self,
settings: InferSettings,
device: torch.device,
hsdp_shard_dim_override: int | None = None,
):
self.settings = settings
self.device = device
self._rewrite_cache: dict[str, str] = {}
# config_class = load_infer_config_class_from_pyfile(settings.config_path)
# self.cfg: InferConfig = config_class()
self.cfg = JoyAIImageInferConfig()
self.cfg.dit_ckpt = settings.ckpt_path
self.cfg.training_mode = False
if hsdp_shard_dim_override is not None:
self.cfg.hsdp_shard_dim = hsdp_shard_dim_override
if int(os.environ.get('WORLD_SIZE', '1')) > 1 and self.cfg.hsdp_shard_dim > 1:
self.cfg.use_fsdp_inference = True
self.dit = load_dit(self.cfg, device=self.device)
self.dit.requires_grad_(False)
self.dit.eval()
self.pipeline = load_pipeline(self.cfg, self.dit, settings.repo_path)
def maybe_rewrite_prompt(self, prompt: str, image: Optional[Image.Image], enabled: bool) -> str:
if not enabled:
return str(prompt or '')
cache_key = f"prompt={prompt.strip()}"
if image is not None:
cache_key += f"|image={image.size[0]}x{image.size[1]}"
if cache_key not in self._rewrite_cache:
self._rewrite_cache[cache_key] = rewrite_prompt(
prompt,
image,
model=self.settings.rewrite_model,
api_key=self.settings.openai_api_key,
base_url=self.settings.openai_base_url,
)
return self._rewrite_cache[cache_key]
@torch.no_grad()
def infer(self, images,height,width,steps,guidance_scale,prompt_embeds,prompt_embeds_mask,negative_prompt_embeds,negative_prompt_embeds_mask,offload,offload_block_num,lat):
# if params.image is None:
# prompts = [params.prompt]
# negative_prompt = [params.neg_prompt]
# images = None
# height = params.height
# width = params.width
# else:
# processed = _dynamic_resize_from_bucket(params.image, basesize=params.basesize)
# width, height = processed.size
# image_tokens = '<image>\n'
# prompts = [f"<|im_start|>user\n{image_tokens}{params.prompt}<|im_end|>\n"]
# negative_prompt = [f"<|im_start|>user\n{image_tokens}{params.neg_prompt}<|im_end|>\n"]
# images = [processed]
# generator_device = 'cuda' if self.device.type == 'cuda' else 'cpu'
# generator = torch.Generator(device=generator_device).manual_seed(int(params.seed))
output = self.pipeline(
prompt=None,
negative_prompt=None,
images=images,
height=height,
width=width,
num_frames=1,
num_inference_steps=steps,
guidance_scale=guidance_scale,
generator=None,
num_videos_per_prompt=1,
output_type='latent',
return_dict=False,
prompt_embeds= prompt_embeds ,
prompt_embeds_mask = prompt_embeds_mask,
negative_prompt_embeds= negative_prompt_embeds ,
negative_prompt_embeds_mask= negative_prompt_embeds_mask,
offload=offload,
offload_block_num=offload_block_num,
lat=lat,
)
return output
#image_tensor = (output[0, -1, 0] * 255).to(torch.uint8).cpu()
#return Image.fromarray(image_tensor.permute(1, 2, 0).numpy())
def build_model(
settings: InferSettings,
device: torch.device | None = None,
hsdp_shard_dim_override: int | None = None,
) -> EditModel:
seed_everything(settings.default_seed)
if device is None:
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
return EditModel(
settings=settings,
device=device,
hsdp_shard_dim_override=hsdp_shard_dim_override,
)
+107
View File
@@ -0,0 +1,107 @@
from __future__ import annotations
import base64
import io
import json
import time
from typing import Optional
from PIL import Image
SYSTEM_PROMPT = r"""
# Edit Prompt Enhancer
You are a professional edit prompt enhancer. Your task is to generate a direct and specific edit prompt based on the user-provided instruction and the image input conditions.
Please strictly follow the enhancing rules below:
## 1. General Principles
- Keep the enhanced prompt direct and specific.
- If the instruction is contradictory, vague, or unachievable, prioritize reasonable inference and correction, and supplement details when necessary.
- Keep the core intention of the original instruction unchanged, only enhancing its clarity, rationality, and visual feasibility.
- All added objects or modifications must align with the logic and style of the edited input image's overall scene.
## 2. Task-Type Handling Rules
### 1. Add, Delete, Replace Tasks
- If the instruction is clear, preserve the original intent and only refine the grammar.
- If the description is vague, supplement with minimal but sufficient details.
### 2. Text Editing Tasks
- All text content must be enclosed in English double quotes.
### 3. Human (ID) Editing Tasks
- Emphasize maintaining the person's core visual consistency.
- For expression changes or beauty changes, they must be natural and subtle.
### 4. Style Conversion or Enhancement Tasks
- Colorization tasks must use: "Restore and colorize the photo."
### 5. Content Filling Tasks
- Inpainting tasks must use: "Perform inpainting on this image. The original caption is: "
- Outpainting tasks must use: "Extend the image beyond its boundaries using outpainting. The original caption is: "
Output a JSON object: {"Rewritten": "..."}
"""
def encode_image_base64_png(image: Image.Image) -> str:
buffer = io.BytesIO()
image.save(buffer, format="PNG")
return base64.b64encode(buffer.getvalue()).decode("utf-8")
def extract_rewritten(content: str) -> str:
text = (content or "").strip().replace("```json", "").replace("```", "").strip()
payload = json.loads(text)
return (payload.get("Rewritten") or "").strip().replace("\n", " ")
def rewrite_prompt(
prompt: str,
image: Optional[Image.Image],
*,
model: str,
api_key: str | None,
base_url: str | None,
max_retries: int = 3,
) -> str:
prompt = str(prompt or "").strip()
if not prompt:
return prompt
if not api_key:
return prompt
from openai import OpenAI
client = OpenAI(api_key=api_key, base_url=base_url) if base_url else OpenAI(api_key=api_key)
user_content: list[dict[str, object]] = []
if image is not None:
user_content.append(
{
"type": "image_url",
"image_url": {"url": f"data:image/png;base64,{encode_image_base64_png(image.convert('RGB'))}"},
}
)
user_content.append({"type": "text", "text": f"User Input: {prompt}\n\nRewritten Prompt:"})
messages = [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_content if image is not None else f"{SYSTEM_PROMPT}\n\nUser Input: {prompt}\n\nRewritten Prompt:"},
]
temperature = 1.0 if "gpt-5" in model.lower() else 0.0
last_error = None
for attempt in range(max_retries):
try:
try:
response = client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
response_format={"type": "json_object"},
)
except Exception:
response = client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
)
rewritten = extract_rewritten(response.choices[0].message.content or "")
return rewritten or prompt
except Exception as exc:
last_error = exc
time.sleep(0.5 * (2 ** attempt))
print(f"[PromptRewrite] failed after {max_retries} retries: {last_error}")
return prompt
+41
View File
@@ -0,0 +1,41 @@
from __future__ import annotations
import os
from dataclasses import dataclass
from .checkpoints import resolve_checkpoint_layout
@dataclass
class InferSettings:
config_path: str
ckpt_path: str
rewrite_model: str
openai_api_key: str | None
openai_base_url: str | None
default_seed: int
repo_path: str | None = None
def load_settings(
*,
ckpt_root: str,
config_path: str | None = None,
rewrite_model: str | None = None,
default_seed: int = 42,
) -> InferSettings:
layout = resolve_checkpoint_layout(ckpt_root)
default_config = layout.root / 'infer_config.py'
if config_path is None and not default_config.exists():
raise FileNotFoundError(
f"Missing inference config: {default_config}. Pass --config explicitly to choose a config file."
)
return InferSettings(
config_path=config_path or str(default_config),
ckpt_path=str(layout.transformer_ckpt),
rewrite_model=rewrite_model or 'gpt-5',
openai_api_key=os.environ.get('OPENAI_API_KEY'),
openai_base_url=os.environ.get('OPENAI_BASE_URL'),
default_seed=default_seed,
)
View File
+253
View File
@@ -0,0 +1,253 @@
import os
import glob
import torch
#import torch.distributed as dist
import gc
from .bucket import BucketGroup
from .mmdit.dit import Transformer3DModel
from .mmdit.text_encoder import load_text_encoder
from .mmdit.vae import WanxVAE
from .pipeline import Pipeline
from .scheduler import FlowMatchDiscreteScheduler
from ..utils.fsdp_load import maybe_load_fsdp_model, pt_weights_iterator, safetensors_weights_iterator
from ..utils.logging import get_logger
from ..utils.constants import PRECISION_TO_TYPE
from ..utils.utils import build_from_config
from transformers import AutoTokenizer
from safetensors.torch import load_file
def load_pipeline(cfg, dit,repo,):
# vae
#factory_kwargs = {
# 'torch_dtype': PRECISION_TO_TYPE[cfg.vae_precision], "device": device}
# vae = build_from_config(cfg.vae_arch_config, **factory_kwargs)
# if getattr(cfg.vae_arch_config, "enable_feature_caching", False):
# vae.enable_feature_caching()
# text_encoder
# factory_kwargs = {
# 'torch_dtype': PRECISION_TO_TYPE[cfg.text_encoder_precision], "device": device}
# tokenizer, text_encoder = build_from_config(
# cfg.text_encoder_arch_config, **factory_kwargs)
tokenizer = AutoTokenizer.from_pretrained(
repo,
local_files_only=True,
trust_remote_code=True,
)
# scheduler
#scheduler = build_from_config(cfg.scheduler_arch_config)
scheduler = FlowMatchDiscreteScheduler(**cfg.scheduler_arch_config["params"])
cfg.repo = repo
pipeline = Pipeline(
vae=None,
tokenizer=tokenizer,
text_encoder=None,
transformer=dit,
scheduler=scheduler,
args=cfg,
)
#pipeline = pipeline.to(device)
return pipeline
def load_dit(cfg, device: torch.device) -> torch.nn.Module:
"""Load DiT model with FSDP support."""
logger = get_logger()
dtype = PRECISION_TO_TYPE[cfg.dit_precision]
model_kwargs = {'dtype': dtype, 'device': device, 'args': cfg}
#model = build_from_config(cfg.dit_arch_config, **model_kwargs)
with torch.device('meta'):
model=Transformer3DModel(**model_kwargs,**cfg.dit_arch_config["params"])
state_dict = None
use_gguf=False
if cfg.dit_ckpt is not None:
logger.info(f"Loading model from: {cfg.dit_ckpt}")
if cfg.dit_ckpt.endswith(".safetensors"):
# Find all safetensors files
# safetensors_files = glob.glob(
# os.path.join(str(cfg.dit_ckpt), "*.safetensors"))
# if not safetensors_files:
# raise ValueError(
# f"No safetensors files found in {cfg.dit_ckpt}")
# state_dict = dict(
# safetensors_weights_iterator(cfg.dit_ckpt))
state_dict=load_file(cfg.dit_ckpt)
elif cfg.dit_ckpt.endswith(".pth"):
# pt_files = [cfg.dit_ckpt]
# state_dict = dict(pt_weights_iterator(pt_files))
state_dict=torch.load(cfg.dit_ckpt, map_location="cpu", weights_only=True)
if "model" in state_dict:
state_dict = state_dict["model"]
elif cfg.dit_ckpt.endswith(".gguf"):
state_dict=load_gguf_checkpoint(cfg.dit_ckpt)
use_gguf=True
else:
raise ValueError(
f"Unknown checkpoint format: {cfg.dit_ckpt}, must be 'safetensor' or 'pth' or 'gguf'")
# if not dist.is_initialized() or dist.get_world_size() == 1:
# # Debug mode
# model.to(device=device)
match_state_dict(model,state_dict)
if state_dict is not None:
# filter unused params
if not use_gguf:
load_state_dict = {}
for k, v in state_dict.items():
if k == "img_in.weight" and model.img_in.weight.shape != v.shape:
logger.info(
f"Inflate {k} from {v.shape} to {model.img_in.weight.shape}")
v_new = v.new_zeros(model.img_in.weight.shape)
v_new[:, :v.shape[1], :, :, :] = v
v = v_new
load_state_dict[k] = v
model.load_state_dict(load_state_dict, strict=False,assign=True)
else:
model = set_gguf2meta_model(model,state_dict,dtype,device)
# model = maybe_load_fsdp_model(
# model=model,
# hsdp_shard_dim=cfg.hsdp_shard_dim,
# reshard_after_forward=cfg.reshard_after_forward,
# param_dtype=dtype,
# reduce_dtype=torch.float32,
# output_dtype=None,
# cpu_offload=cfg.cpu_offload,
# fsdp_inference=cfg.use_fsdp_inference,
# training_mode=cfg.training_mode,
# pin_cpu_memory=cfg.pin_cpu_memory,
# )
# Log model info
total_params = sum(p.numel() for p in model.parameters())
logger.info(f"Instantiate model with {total_params / 1e9:.2f}B parameters")
# Ensure consistent dtype
param_dtypes = {param.dtype for param in model.parameters()}
if len(param_dtypes) > 1:
logger.warning(
f"Model has mixed dtypes: {param_dtypes}. Converting to {dtype}")
model = model.to(dtype)
return model.eval()
def load_gguf_checkpoint(gguf_checkpoint_path):
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
from diffusers.utils import is_gguf_available, is_torch_available
if is_gguf_available() and is_torch_available():
import gguf
from gguf import GGUFReader
from diffusers.quantizers.gguf.utils import SUPPORTED_GGUF_QUANT_TYPES, GGUFParameter
else:
logger.error(
"Loading a GGUF checkpoint in PyTorch, requires both PyTorch and GGUF>=0.10.0 to be installed. Please see "
"https://pytorch.org/ and https://github.com/ggerganov/llama.cpp/tree/master/gguf-py for installation instructions."
)
raise ImportError("Please install torch and gguf>=0.10.0 to load a GGUF checkpoint in PyTorch.")
reader = GGUFReader(gguf_checkpoint_path)
parsed_parameters = {}
for i, tensor in enumerate(reader.tensors):
name = tensor.name
quant_type = tensor.tensor_type
is_gguf_quant = quant_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
if is_gguf_quant and quant_type not in SUPPORTED_GGUF_QUANT_TYPES:
_supported_quants_str = "\n".join([str(type) for type in SUPPORTED_GGUF_QUANT_TYPES])
raise ValueError(
(
f"{name} has a quantization type: {str(quant_type)} which is unsupported."
"\n\nCurrently the following quantization types are supported: \n\n"
f"{_supported_quants_str}"
"\n\nTo request support for this quantization type please open an issue here: https://github.com/huggingface/diffusers"
)
)
weights = torch.from_numpy(tensor.data) #tensor.data.copy()
parsed_parameters[name.replace("model.", "")] = GGUFParameter(weights, quant_type=quant_type) if is_gguf_quant else weights
del tensor,weights
if i > 0 and i % 1000 == 0: # 每1000个tensor执行一次gc
logger.info(f"Processed {i}tensors...")
gc.collect()
del reader
gc.collect()
return parsed_parameters
def set_gguf2meta_model(meta_model,model_state_dict,dtype,device):
from diffusers import GGUFQuantizationConfig
from diffusers.quantizers.gguf import GGUFQuantizer
g_config = GGUFQuantizationConfig(compute_dtype=dtype or torch.bfloat16)
hf_quantizer = GGUFQuantizer(quantization_config=g_config)
hf_quantizer.pre_quantized = True
hf_quantizer._process_model_before_weight_loading(
meta_model,
device_map={"": device} if device else None,
state_dict=model_state_dict
)
from diffusers.models.model_loading_utils import load_model_dict_into_meta
x,y=load_model_dict_into_meta(
meta_model,
model_state_dict,
hf_quantizer=hf_quantizer,
device_map={"": device} if device else None,
dtype=dtype
)
print(x,"offload_index")
print(y,"state_dict_index")
hf_quantizer._process_model_after_weight_loading(meta_model)
del model_state_dict
gc.collect()
return meta_model.to(dtype=dtype)
def match_state_dict(meta_model, sd,show_num=10):
meta_model_keys = set(meta_model.state_dict().keys())
state_dict_keys = set(sd.keys())
# 打印匹配的键的数量
matching_keys = meta_model_keys.intersection(state_dict_keys)
print(f"Matching keys count: {len(matching_keys)}")
# 打印不在 meta_model 中但在 state_dict 中的键(多余键)
extra_keys = state_dict_keys - meta_model_keys
if extra_keys:
print(f"Extra keys in state_dict (not in meta_model): {len(extra_keys)}")
for key in list(extra_keys)[:show_num]: # 只显示前10个
print(f" - {key}")
# 打印不在 state_dict 中但在 meta_model 中的键(缺失键)
missing_keys = meta_model_keys - state_dict_keys
if missing_keys:
print(f"Missing keys in state_dict (not in state_dict): {len(missing_keys)}")
for key in list(missing_keys)[:show_num]: # 只显示前10个
print(f" - {key}")
# 如果需要,也可以打印部分匹配的键
print(f"Sample matching keys: {list(matching_keys)[:5]}")
__all__ = [
"BucketGroup",
"FlowMatchDiscreteScheduler",
"Pipeline",
"Transformer3DModel",
"WanxVAE",
"load_pipeline",
"load_text_encoder",
]
+119
View File
@@ -0,0 +1,119 @@
# Adapted from https://github.com/hao-ai-lab/FastVideo/tree/main/fastvideo/attention
import os
import sys
import torch
from einops import rearrange
_FLASH_ATTN_IMPORT_ERROR = None
try:
# Check for Flash Attention 3 installation path
flash_attn3_path = os.getenv("FLASH_ATTN3_PATH")
if flash_attn3_path:
print(f"Using Flash Attention 3 from: {flash_attn3_path}")
sys.path.insert(0, flash_attn3_path)
from flash_attn_interface import flash_attn_varlen_func
else:
from flash_attn.flash_attn_interface import flash_attn_varlen_func
except ImportError as exc:
flash_attn_varlen_func = None
_FLASH_ATTN_IMPORT_ERROR = exc
def is_flash_attn_available() -> bool:
return flash_attn_varlen_func is not None
def get_preferred_attention_backend() -> str:
return "flash_attn" if is_flash_attn_available() else "torch_spda"
def describe_attention_backend() -> str:
backend = get_preferred_attention_backend()
if backend == "flash_attn":
return "flash_attn"
if _FLASH_ATTN_IMPORT_ERROR is None:
return "torch_spda"
return f"torch_spda (flash_attn unavailable: {_FLASH_ATTN_IMPORT_ERROR})"
def get_cu_seqlens(text_mask, img_len):
"""Calculate cu_seqlens_q, cu_seqlens_kv using text_mask and img_len
Args:
text_mask (torch.Tensor): the mask of text
img_len (int): the length of image
Returns:
torch.Tensor: the calculated cu_seqlens for flash attention
"""
batch_size = text_mask.shape[0]
text_len = text_mask.sum(dim=1)
max_len = text_mask.shape[1] + img_len
cu_seqlens = torch.zeros([2 * batch_size + 1],
dtype=torch.int32, device="cuda")
for i in range(batch_size):
s = text_len[i] + img_len
s1 = i * max_len + s
s2 = (i + 1) * max_len
cu_seqlens[2 * i + 1] = s1
cu_seqlens[2 * i + 2] = s2
return cu_seqlens
def attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
backend: str = "flash_attn",
*,
causal: bool = False,
softmax_scale: float = None,
attn_kwargs: dict = None,
):
"""
Args:
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads
"""
if backend == "auto":
backend = get_preferred_attention_backend()
# Fall back to torch_spda when flash_attn was requested but unavailable
if backend == "flash_attn" and flash_attn_varlen_func is None:
backend = "torch_spda"
assert backend in [
"torch_spda", "flash_attn"], f"Unsupported attention backend: {backend}"
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Input tensors must be 4D"
batch_size = q.shape[0]
if backend == "torch_spda":
q = rearrange(q, "b l h c -> b h l c")
k = rearrange(k, "b l h c -> b h l c")
v = rearrange(v, "b l h c -> b h l c")
output = torch.nn.functional.scaled_dot_product_attention(
q, k, v, is_causal=causal, scale=softmax_scale)
output = rearrange(output, "b h l c -> b l h c")
elif backend == "flash_attn":
cu_seqlens_q = attn_kwargs['cu_seqlens_q']
cu_seqlens_kv = attn_kwargs['cu_seqlens_kv']
max_seqlen_q = attn_kwargs['max_seqlen_q']
max_seqlen_kv = attn_kwargs['max_seqlen_kv']
x = flash_attn_varlen_func(
q.view(q.shape[0] * q.shape[1], *q.shape[2:]),
k.view(k.shape[0] * k.shape[1], *k.shape[2:]),
v.view(v.shape[0] * v.shape[1], *v.shape[2:]),
cu_seqlens_q,
cu_seqlens_kv,
max_seqlen_q,
max_seqlen_kv,
)
output = x.view(
batch_size, max_seqlen_q, x.shape[-2], x.shape[-1]
)
return output
+108
View File
@@ -0,0 +1,108 @@
class BucketGroup:
"""Manages dynamic batch grouping buckets for image inference."""
def __init__(
self,
bucket_configs: list[tuple[int, int, int, int, int]],
prioritize_frame_matching: bool = True,
):
"""
Initialize bucket group with predefined configurations.
Args:
bucket_configs: List of (batch_size, num_items, num_frames, height, width) tuples
prioritize_frame_matching: Unused, kept for API compatibility.
"""
self.bucket_configs = [tuple(b) for b in bucket_configs]
def find_best_bucket(self, media_shape: tuple[int, int, int, int]) -> tuple[int, int, int, int, int]:
"""
Find the best matching bucket for given media dimensions.
Args:
media_shape: (num_items, num_frames, height, width) of input media
Returns:
Best matching bucket as (batch_size, num_items, num_frames, height, width)
"""
num_items, num_frames, height, width = media_shape
target_aspect_ratio = height / width
if num_frames != 1:
raise ValueError(
f"Only image inference (num_frames=1) is supported, got num_frames={num_frames}")
valid_buckets = [
b for b in self.bucket_configs
if b[1] == num_items and b[2] == 1
]
if not valid_buckets:
raise ValueError(
f"No image buckets found for shape {media_shape}")
return min(
valid_buckets,
key=lambda bucket: abs(
(bucket[3] / bucket[4]) - target_aspect_ratio)
)
def __repr__(self) -> str:
return (
f"BucketGroup("
f"total_buckets={len(self.bucket_configs)}, "
f"configs={self.bucket_configs})"
)
def _generate_hw_buckets(base_height=256, base_width=256, step_width=16, step_height=16, max_ratio=4.0) -> list[tuple[int, int, int, int, int]]:
"""Generate dimension buckets based on aspect ratios."""
buckets = []
target_pixels = base_height * base_width
height = target_pixels // step_width
width = step_width
while height >= step_height:
if max(height, width) / min(height, width) <= max_ratio:
buckets.append((1, 1, 1, height, width))
if height * (width + step_width) <= target_pixels:
width += step_width
else:
height -= step_height
return buckets
def generate_video_image_bucket(basesize=256, min_temporal=65, max_temporal=129, bs_img=8, bs_vid=1, bs_mimg=4, min_items=1, max_items=1):
"""Generate bucket configs for image inference.
Returns:
List of (batch_size, num_items, num_frames, height, width) tuples.
"""
assert basesize in [
256, 512, 768, 1024], f"[generate_video_image_bucket] wrong basesize {basesize}"
bucket_list = []
base_bucket_list = _generate_hw_buckets()
# image
for _bucket in base_bucket_list:
bucket = list(_bucket)
bucket[0] = bs_img
bucket_list.append(bucket)
# multiple images
for num_items in range(min_items, max_items + 1):
for _bucket in base_bucket_list:
bucket = list(_bucket)
bucket[0] = bs_mimg
bucket[1] = num_items
bucket_list.append(bucket)
# spatial resize
if basesize > 256:
ratio = basesize // 256
def resize(bucket, r):
bucket[-2] *= r
bucket[-1] *= r
return bucket
bucket_list = [resize(bucket, ratio) for bucket in bucket_list]
return bucket_list
+4
View File
@@ -0,0 +1,4 @@
from .models import Transformer3DModel
__all__ = ["Transformer3DModel"]
+660
View File
@@ -0,0 +1,660 @@
from typing import Any, List, Tuple, Optional, Union, Dict
from einops import rearrange
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.attention import FeedForward
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps
import gc
import copy
from ....models.attention import attention, get_cu_seqlens
from .posemb_layers import apply_rotary_emb, get_nd_rotary_pos_embed
from .modulate_layers import load_modulation, modulate, apply_gate
class BlockGPUManager:
def __init__(self, device="cuda", block_group_size=1):
self.device = torch.device(device)
self.managed_modules = []
self.submodule = []
self.block_group_size = block_group_size # 每次加载的连续层数
self._original_model_ref = None
self._original_block_ref = None
self._num_groups = 0 # 总批次数
self._group_loaded: list[bool] = []
def setup_for_inference(self, transformer_model):
self._collect_managed_modules(transformer_model)
self._initialize_submodule()
return self
def _collect_managed_modules(self, transformer_model):
self.submodule = []
self._original_model_ref = transformer_model
self._original_block_ref = transformer_model.double_blocks
self._num_groups = (len(self._original_block_ref) + self.block_group_size - 1) // self.block_group_size
for attr in ['img_in', 'condition_embedder', 'norm_out', 'proj_out', ]:
if hasattr(transformer_model, attr):
self.submodule.append(getattr(transformer_model, attr))
self.managed_modules = [None] * self._num_groups
self._group_loaded = [False] * self._num_groups
def _load_group(self, group_index):
"""加载指定组的数据块"""
if self._group_loaded[group_index]:
return
start_idx = group_index * self.block_group_size
end_idx = min(start_idx + self.block_group_size, len(self._original_block_ref))
group = nn.ModuleList()
for layer in self._original_block_ref[start_idx:end_idx]:
# 深拷贝当前层
cpu_layer = copy.deepcopy(layer)
# 移动到目标设备
cpu_layer.to(self.device)
group.append(cpu_layer)
self.managed_modules[group_index] = group
self._group_loaded[group_index] = True
def _unload_group(self, group_index):
"""卸载指定组的数据块"""
if not self._group_loaded[group_index]:
return
group = self.managed_modules[group_index]
self.managed_modules[group_index] = None
self._group_loaded[group_index] = False
# 显式删除引用
group = None
del group
# 清理GPU缓存
torch.cuda.empty_cache()
def _get_layer(self, layer_index):
"""按需获取层,实现按组加载"""
group_index = layer_index // self.block_group_size
local_idx = layer_index % self.block_group_size
# 如果组未加载,则加载该组
if not self._group_loaded[group_index]:
self._load_group(group_index)
# 返回组中的指定层
return self.managed_modules[group_index][local_idx]
def _unload_unused_groups(self, keep: set[int]):
"""卸载不需要的组"""
for index in range(self._num_groups):
if self._group_loaded[index] and index not in keep:
self._unload_group(index)
def _initialize_submodule(self):
for module in self.submodule:
if hasattr(module, 'to'):
module.to(self.device)
return self
def unload_all_blocks_to_cpu(self):
# 卸载所有组
for group_index in range(self._num_groups):
self._unload_group(group_index)
# 将embedder和output模块移到CPU
for module in self.submodule:
if hasattr(module, 'to'):
module.to('cpu', non_blocking=True)
torch.cuda.empty_cache()
return self
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal dit block with seperate modulation for
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
def __init__(
self,
hidden_size: int,
heads_num: int,
mlp_width_ratio: float,
mlp_act_type: str = "gelu_tanh",
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
dit_modulation_type: Optional[str] = "wanx",
attn_backend: str = 'flash_attn',
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.attn_backend = attn_backend
self.dit_modulation_type = dit_modulation_type
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.img_mod = load_modulation(
modulate_type=self.dit_modulation_type,
hidden_size=hidden_size,
factor=6,
**factory_kwargs,
)
self.img_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=True, **factory_kwargs
)
self.img_attn_q_norm = RMSNorm(head_dim, elementwise_affine=True,
eps=1e-6, **factory_kwargs)
self.img_attn_k_norm = RMSNorm(head_dim, elementwise_affine=True,
eps=1e-6, **factory_kwargs)
self.img_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=True, **factory_kwargs
)
self.img_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
# There is no dtype fpr FeedForward, because FSDP2 casts the dtype for all parameters.
# You may need to give the dtype when no autocast and fsdp !!!
self.img_mlp = FeedForward(hidden_size, inner_dim=mlp_hidden_dim,
activation_fn="gelu-approximate")
self.txt_mod = load_modulation(
modulate_type=self.dit_modulation_type,
hidden_size=hidden_size,
factor=6,
**factory_kwargs,
)
self.txt_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=True, **factory_kwargs
)
self.txt_attn_q_norm = RMSNorm(head_dim, elementwise_affine=True,
eps=1e-6, **factory_kwargs)
self.txt_attn_k_norm = RMSNorm(head_dim, elementwise_affine=True,
eps=1e-6, **factory_kwargs)
self.txt_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=True, **factory_kwargs
)
self.txt_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_mlp = FeedForward(hidden_size, inner_dim=mlp_hidden_dim,
activation_fn="gelu-approximate")
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
vis_freqs_cis: tuple = None,
txt_freqs_cis: tuple = None,
attn_kwargs: Optional[dict] = {},
) -> Tuple[torch.Tensor, torch.Tensor]:
tt, th, tw = attn_kwargs['thw']
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(vec)
(
txt_mod1_shift,
txt_mod1_scale,
txt_mod1_gate,
txt_mod2_shift,
txt_mod2_scale,
txt_mod2_gate,
) = self.txt_mod(vec)
# Prepare image for attention.
img_modulated = self.img_norm1(img)
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply RoPE if needed.
if vis_freqs_cis is not None:
img_qq, img_kk = apply_rotary_emb(
img_q, img_k, vis_freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
if txt_freqs_cis is not None:
raise NotImplementedError("RoPE text is not supported for inference")
txt_qq, txt_kk = apply_rotary_emb(
txt_q, txt_k, txt_freqs_cis, head_first=False)
assert (
txt_qq.shape == txt_q.shape and txt_kk.shape == txt_k.shape
), f"txt_kk: {txt_qq.shape}, txt_q: {txt_q.shape}, txt_kk: {txt_kk.shape}, txt_k: {txt_k.shape}"
txt_q, txt_k = txt_qq, txt_kk
# attention computation start
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention(
q, k, v,
backend=self.attn_backend,
attn_kwargs=attn_kwargs,
)
attn = attn.flatten(2, 3)
# attention computation end
img_attn, txt_attn = attn[:,
: img.shape[1]], attn[:, img.shape[1]:]
# Calculate the img bloks.
img = img + apply_gate(self.img_attn_proj(img_attn),
gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn),
gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
),
gate=txt_mod2_gate,
)
return img, txt
class WanTimeTextImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
time_proj_dim: int,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
pos_embed_seq_len: Optional[int] = None,
):
super().__init__()
self.timesteps_proj = Timesteps(
num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = TimestepEmbedding(
in_channels=time_freq_dim, time_embed_dim=dim)
self.act_fn = nn.SiLU()
self.time_proj = nn.Linear(dim, time_proj_dim)
self.text_embedder = PixArtAlphaTextProjection(
text_embed_dim, dim, act_fn="gelu_tanh")
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
):
timestep = self.timesteps_proj(timestep)
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
timestep = timestep.to(time_embedder_dtype)
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
timestep_proj = self.time_proj(self.act_fn(temb))
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
return temb, timestep_proj, encoder_hidden_states
class Transformer3DModel(ModelMixin, ConfigMixin):
_fsdp_shard_conditions: list = [
lambda name, module: isinstance(module, (MMDoubleStreamBlock))]
_supports_gradient_checkpointing = True
@register_to_config
def __init__(
self,
args: Any,
patch_size: list = [1, 2, 2],
in_channels: int = 4, # Should be VAE.config.latent_channels.
out_channels: int = None,
hidden_size: int = 3072,
heads_num: int = 24,
text_states_dim: int = 4096,
mlp_width_ratio: float = 4.0,
mm_double_blocks_depth: int = 20,
rope_dim_list: List[int] = [16, 56, 56],
rope_type: str = 'rope',
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
dit_modulation_type: str = "wanx",
attn_backend: str = 'flash_attn',
theta: int = 256,
):
self.args = args
self.out_channels = out_channels or in_channels
self.patch_size = patch_size
self.hidden_size = hidden_size
self.heads_num = heads_num
self.rope_dim_list = rope_dim_list
self.dit_modulation_type = dit_modulation_type
self.mm_double_blocks_depth = mm_double_blocks_depth
self.attn_backend = attn_backend
self.rope_type = rope_type
self.theta = theta
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.hidden_size = hidden_size
if hidden_size % heads_num != 0:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
)
# image projection
self.img_in = nn.Conv3d(
in_channels, hidden_size, kernel_size=patch_size, stride=patch_size)
# condition embedding
self.condition_embedder = WanTimeTextImageEmbedding(
dim=hidden_size,
time_freq_dim=256,
time_proj_dim=hidden_size * 6,
text_embed_dim=text_states_dim,
)
# double blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
dit_modulation_type=self.dit_modulation_type,
attn_backend=attn_backend,
**factory_kwargs,
)
for _ in range(mm_double_blocks_depth)
]
)
# Output norm & projection
self.norm_out = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6
)
self.proj_out = nn.Linear(
hidden_size, out_channels * math.prod(patch_size),
**factory_kwargs)
def get_rotary_pos_embed(self, vis_rope_size, txt_rope_size=None):
target_ndim = 3
ndim = 5 - 2
if len(vis_rope_size) != target_ndim:
vis_rope_size = [1] * (target_ndim - len(vis_rope_size)
) + vis_rope_size # time axis
head_dim = self.hidden_size // self.heads_num
rope_dim_list = self.rope_dim_list
if rope_dim_list is None:
rope_dim_list = [head_dim //
target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
vis_freqs, txt_freqs = get_nd_rotary_pos_embed(
rope_dim_list,
vis_rope_size,
txt_rope_size=txt_rope_size,
theta=self.theta,
use_real=True,
theta_rescale_factor=1,
)
return vis_freqs, txt_freqs
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor, # Should be in range(0, 1000).
encoder_hidden_states: torch.Tensor = None,
encoder_hidden_states_mask: torch.Tensor = None,
return_dict: bool = True,
gpu_manager: Optional[BlockGPUManager] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
# For Multi-item Input: hidden_states: (b, n, c, t, h, w)
# Permute the items into the temporal dimension
is_multi_item = (len(hidden_states.shape) == 6)
num_items = 0
if is_multi_item:
num_items = hidden_states.shape[1]
if num_items > 1:
assert self.patch_size[0] == 1, "For multi-item input, patch_size[0] must be 1"
# Move the last item to the first position
hidden_states = torch.cat(
[
hidden_states[:, -1:],
hidden_states[:, :-1]
],
dim=1
)
hidden_states = rearrange(
hidden_states, 'b n c t h w -> b c (n t) h w')
out = {}
batch_size, _, ot, oh, ow = hidden_states.shape
tt, th, tw = (
ot // self.patch_size[0],
oh // self.patch_size[1],
ow // self.patch_size[2],
)
# Text Mask
if encoder_hidden_states_mask == None:
encoder_hidden_states_mask = torch.ones(
(encoder_hidden_states.shape[0], encoder_hidden_states.shape[1]), dtype=torch.bool).to(encoder_hidden_states.device)
# Prepare img, txt, vec.
img = self.img_in(hidden_states).flatten(2).transpose(1, 2)
temb, vec, txt = self.condition_embedder(
timestep, encoder_hidden_states)
if vec.shape[-1] > self.hidden_size:
vec = vec.unflatten(1, (6, -1))
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
# rope
vis_freqs_cis, txt_freqs_cis = self.get_rotary_pos_embed(vis_rope_size=(
tt, th, tw), txt_rope_size=txt_seq_len if self.rope_type == 'mrope' else None)
# Compute attn_kwargs
attn_kwargs = {'thw': [tt, th, tw], 'txt_len': txt_seq_len}
if self.attn_backend == 'flash_attn':
cu_seqlens_q = get_cu_seqlens(
encoder_hidden_states_mask, img_seq_len)
cu_seqlens_kv = cu_seqlens_q
max_seqlen_q = img_seq_len + txt_seq_len
max_seqlen_kv = max_seqlen_q
attn_kwargs.update({
'cu_seqlens_q': cu_seqlens_q,
'cu_seqlens_kv': cu_seqlens_kv,
'max_seqlen_q': max_seqlen_q,
'max_seqlen_kv': max_seqlen_kv,
})
# --------------------- Pass through DiT blocks ------------------------
#for _, block in enumerate(self.double_blocks):
for layer_index in range(len(self.double_blocks)):
if gpu_manager is not None:
block = gpu_manager._get_layer(layer_index)
else:
block = self.double_blocks[layer_index]
double_block_args = [
img,
txt,
vec,
vis_freqs_cis,
txt_freqs_cis,
attn_kwargs
]
img, txt = block(*double_block_args)
if gpu_manager is not None:
current_group = layer_index // gpu_manager.block_group_size
next_group = current_group + 1
gpu_manager._unload_unused_groups(keep={current_group, next_group})
img_len = img.shape[1]
x = torch.cat((img, txt), 1)
img = x[:, :img_len, ...]
# ---------------------------- Final layer ------------------------------
img = self.proj_out(self.norm_out(img))
img = self.unpatchify(img, tt, th, tw)
# Reshape back to multiple items
if is_multi_item:
img = rearrange(
img, 'b c (n t) h w -> b n c t h w', n=num_items)
if num_items > 1:
# Move the first item back to the last position
img = torch.cat(
[
img[:, 1:],
img[:, :1]
],
dim=1
)
return (img, txt)
def unpatchify(self, x, t, h, w):
"""
x: (N, T, patch_size**2 * C)
imgs: (N, H, W, C)
"""
c = self.out_channels
pt, ph, pw = self.patch_size
assert t * h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], t, h, w, pt, ph, pw, c))
x = torch.einsum("nthwopqc->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
@@ -0,0 +1,80 @@
import torch
import torch.nn as nn
def load_modulation(
modulate_type: str,
hidden_size: int,
factor: int,
act_layer=nn.SiLU,
dtype=None,
device=None):
factory_kwargs = {"dtype": dtype, "device": device}
if modulate_type == 'wanx':
return ModulateWan(hidden_size, factor, **factory_kwargs)
raise ValueError(
f"Unknown modulation type: {modulate_type}. Only 'wanx' is supported.")
class ModulateWan(nn.Module):
"""Modulation layer for WanX."""
def __init__(
self,
hidden_size: int,
factor: int,
dtype=None,
device=None,
):
super().__init__()
self.factor = factor
self.modulate_table = nn.Parameter(
torch.zeros(1, factor, hidden_size,
dtype=dtype, device=device) / hidden_size**0.5,
requires_grad=True
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if len(x.shape) != 3:
x = x.unsqueeze(1)
return [o.squeeze(1) for o in (self.modulate_table + x).chunk(self.factor, dim=1)]
def modulate(x, shift=None, scale=None):
"""modulate by shift and scale
Args:
x (torch.Tensor): input tensor.
shift (torch.Tensor, optional): shift tensor. Defaults to None.
scale (torch.Tensor, optional): scale tensor. Defaults to None.
Returns:
torch.Tensor: the output tensor after modulate.
"""
if scale is None and shift is None:
return x
elif shift is None:
return x * (1 + scale.unsqueeze(1))
elif scale is None:
return x + shift.unsqueeze(1)
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def apply_gate(x, gate=None, tanh=False):
"""Apply gating to tensor.
Args:
x (torch.Tensor): input tensor.
gate (torch.Tensor, optional): gate tensor. Defaults to None.
tanh (bool, optional): whether to use tanh function. Defaults to False.
Returns:
torch.Tensor: the output tensor after apply gate.
"""
if gate is None:
return x
if tanh:
return x * gate.unsqueeze(1).tanh()
else:
return x * gate.unsqueeze(1)
@@ -0,0 +1,320 @@
import torch
from typing import Union, Tuple, List
def _to_tuple(x, dim=2):
if isinstance(x, int):
return (x,) * dim
elif len(x) == dim:
return x
else:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start, *args, dim=2):
"""
Get n-D meshgrid with start, stop and num.
Args:
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
n-tuples.
*args: See above.
dim (int): Dimension of the meshgrid. Defaults to 2.
Returns:
grid (np.ndarray): [dim, ...]
"""
if len(args) == 0:
# start is grid_size
num = _to_tuple(start, dim=dim)
start = (0,) * dim
stop = num
elif len(args) == 1:
# start is start, args[0] is stop, step is 1
start = _to_tuple(start, dim=dim)
stop = _to_tuple(args[0], dim=dim)
num = [stop[i] - start[i] for i in range(dim)]
elif len(args) == 2:
# start is start, args[0] is stop, args[1] is num
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
else:
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
axis_grid = []
for i in range(dim):
a, b, n = start[i], stop[i], num[i]
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
axis_grid.append(g)
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
return grid
#################################################################################
# Rotary Positional Embedding Functions #
#################################################################################
# https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L80
def reshape_for_broadcast(
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
x: torch.Tensor,
head_first=False,
):
"""
Reshape frequency tensor for broadcasting it with another tensor.
This function reshapes the frequency tensor to have the same shape as the target tensor 'x'
for the purpose of broadcasting the frequency tensor during element-wise operations.
Notes:
When using FlashMHAModified, head_first should be False.
When using Attention, head_first should be True.
Args:
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.
x (torch.Tensor): Target tensor for broadcasting compatibility.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
torch.Tensor: Reshaped frequency tensor.
Raises:
AssertionError: If the frequency tensor doesn't match the expected shape.
AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.
"""
ndim = x.ndim
assert 0 <= 1 < ndim
if isinstance(freqs_cis, tuple):
# freqs_cis: (cos, sin) in real space
if head_first:
assert freqs_cis[0].shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis[0].shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim -
1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
else:
# freqs_cis: values in complex space
if head_first:
assert freqs_cis.shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis.shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim -
1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis.view(*shape)
def rotate_half(x):
x_real, x_imag = (
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
) # [B, S, H, D//2]
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
def apply_rotary_emb(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
head_first: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor.
This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided
frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor
is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are
returned as real tensors.
Args:
xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]
xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D]
freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
xk_out = None
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
cos, sin = cos.to(xq.device), sin.to(xq.device)
# real * cos - imag * sin
# imag * cos + real * sin
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
return xq_out, xk_out
def get_nd_rotary_pos_embed(
rope_dim_list,
start,
*args,
theta=10000.0,
use_real=False,
txt_rope_size = None,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
):
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
Args:
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
sum(rope_dim_list) should equal to head_dim of attention layer.
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
*args: See above.
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.
Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real
part and an imaginary part separately.
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
Returns:
pos_embed (torch.Tensor): [HW, D/2]
"""
grid = get_meshgrid_nd(
start, *args, dim=len(rope_dim_list)
) # [3, W, H, D] / [2, W, H]
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list
), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(
rope_dim_list[i],
grid[i].reshape(-1),
theta,
use_real=use_real,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
if use_real:
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
vis_emb = (cos, sin)
else:
vis_emb = torch.cat(embs, dim=1) # (WHD, D/2)
# add text rope
if txt_rope_size is not None:
embs_txt = []
vis_max_ids = grid.view(-1).max().item()
grid_txt = torch.arange(txt_rope_size) + vis_max_ids + 1
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(
rope_dim_list[i],
grid_txt,
theta,
use_real=use_real,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
)
embs_txt.append(emb)
if use_real:
cos = torch.cat([emb[0] for emb in embs_txt], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs_txt], dim=1) # (WHD, D/2)
txt_emb = (cos, sin)
else:
txt_emb = torch.cat(embs_txt, dim=1) # (WHD, D/2)
else:
txt_emb = None
return vis_emb, txt_emb
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
theta: float = 10000.0,
use_real: bool = False,
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
and the end index 'end'. The 'theta' parameter scales the frequencies.
The returned tensor contains complex values in complex64 data type.
Args:
dim (int): Dimension of the frequency tensor.
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool, optional): If True, return real part and imaginary part separately.
Otherwise, return complex numbers.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
Returns:
freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor ** (dim / (dim - 2))
freqs = 1.0 / (
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
) # [D/2]
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
if use_real:
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(
torch.ones_like(freqs), freqs
) # complex64 # [S, D/2]
return freqs_cis
@@ -0,0 +1,23 @@
import torch
from transformers import Qwen3VLForConditionalGeneration, AutoTokenizer
def load_text_encoder(
text_encoder_ckpt: str,
device: torch.device = torch.device("cpu"),
torch_dtype: torch.dtype = torch.bfloat16,
):
loader = Qwen3VLForConditionalGeneration #or AutoModelForVision2Seq
model = loader.from_pretrained(
text_encoder_ckpt,
torch_dtype=torch_dtype,
local_files_only=True,
trust_remote_code=True,
).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(
text_encoder_ckpt,
local_files_only=True,
trust_remote_code=True,
)
return tokenizer, model
+3
View File
@@ -0,0 +1,3 @@
from .wanvae import WanxVAE
__all__ = ["WanxVAE"]
+697
View File
@@ -0,0 +1,697 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import logging
import torch
import torch.cuda.amp as amp
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
__all__ = [
'WanVAE',
]
CACHE_T = 2
class CausalConv3d(nn.Conv3d):
"""
Causal 3d convolusion.
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
return super().forward(x)
class RMS_norm(nn.Module):
def __init__(self, dim, channel_first=True, images=True, bias=False):
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
def forward(self, x):
return F.normalize(
x, dim=(1 if self.channel_first else
-1)) * self.scale * self.gamma + self.bias
class Upsample(nn.Upsample):
def forward(self, x):
"""
Fix bfloat16 support for nearest neighbor interpolation.
"""
return super().forward(x.float()).type_as(x)
class Resample(nn.Module):
def __init__(self, dim, mode):
assert mode in ('none', 'upsample2d', 'upsample3d', 'downsample2d',
'downsample3d')
super().__init__()
self.dim = dim
self.mode = mode
# layers
if mode == 'upsample2d':
self.resample = nn.Sequential(
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
nn.Conv2d(dim, dim // 2, 3, padding=1))
elif mode == 'upsample3d':
self.resample = nn.Sequential(
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
nn.Conv2d(dim, dim // 2, 3, padding=1))
self.time_conv = CausalConv3d(
dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
elif mode == 'downsample2d':
self.resample = nn.Sequential(
nn.ZeroPad2d((0, 1, 0, 1)),
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
elif mode == 'downsample3d':
self.resample = nn.Sequential(
nn.ZeroPad2d((0, 1, 0, 1)),
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
self.time_conv = CausalConv3d(
dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
else:
self.resample = nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]):
b, c, t, h, w = x.size()
if self.mode == 'upsample3d':
if feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = 'Rep'
feat_idx[0] += 1
else:
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[
idx] is not None and feat_cache[idx] != 'Rep':
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
if cache_x.shape[2] < 2 and feat_cache[
idx] is not None and feat_cache[idx] == 'Rep':
cache_x = torch.cat([
torch.zeros_like(cache_x).to(cache_x.device),
cache_x
],
dim=2)
if feat_cache[idx] == 'Rep':
x = self.time_conv(x)
else:
x = self.time_conv(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
3)
x = x.reshape(b, c, t * 2, h, w)
t = x.shape[2]
x = rearrange(x, 'b c t h w -> (b t) c h w')
x = self.resample(x)
x = rearrange(x, '(b t) c h w -> b c t h w', t=t)
if self.mode == 'downsample3d':
if feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = x.clone()
feat_idx[0] += 1
else:
cache_x = x[:, :, -1:, :, :].clone()
# if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx]!='Rep':
# # cache last frame of last two chunk
# cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = self.time_conv(
torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
feat_cache[idx] = cache_x
feat_idx[0] += 1
return x
def init_weight(self, conv):
conv_weight = conv.weight
nn.init.zeros_(conv_weight)
c1, c2, t, h, w = conv_weight.size()
one_matrix = torch.eye(c1, c2)
init_matrix = one_matrix
nn.init.zeros_(conv_weight)
#conv_weight.data[:,:,-1,1,1] = init_matrix * 0.5
conv_weight.data[:, :, 1, 0, 0] = init_matrix #* 0.5
conv.weight.data.copy_(conv_weight)
nn.init.zeros_(conv.bias.data)
def init_weight2(self, conv):
conv_weight = conv.weight.data
nn.init.zeros_(conv_weight)
c1, c2, t, h, w = conv_weight.size()
init_matrix = torch.eye(c1 // 2, c2)
#init_matrix = repeat(init_matrix, 'o ... -> (o 2) ...').permute(1,0,2).contiguous().reshape(c1,c2)
conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
conv.weight.data.copy_(conv_weight)
nn.init.zeros_(conv.bias.data)
class ResidualBlock(nn.Module):
def __init__(self, in_dim, out_dim, dropout=0.0):
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
# layers
self.residual = nn.Sequential(
RMS_norm(in_dim, images=False), nn.SiLU(),
CausalConv3d(in_dim, out_dim, 3, padding=1),
RMS_norm(out_dim, images=False), nn.SiLU(), nn.Dropout(dropout),
CausalConv3d(out_dim, out_dim, 3, padding=1))
self.shortcut = CausalConv3d(in_dim, out_dim, 1) \
if in_dim != out_dim else nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x)
for layer in self.residual:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x + h
class AttentionBlock(nn.Module):
"""
Causal self-attention with a single head.
"""
def __init__(self, dim):
super().__init__()
self.dim = dim
# layers
self.norm = RMS_norm(dim)
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
self.proj = nn.Conv2d(dim, dim, 1)
# zero out the last layer params
nn.init.zeros_(self.proj.weight)
def forward(self, x):
identity = x
b, c, t, h, w = x.size()
x = rearrange(x, 'b c t h w -> (b t) c h w')
x = self.norm(x)
# compute query, key, value
q, k, v = self.to_qkv(x).reshape(b * t, 1, c * 3,
-1).permute(0, 1, 3,
2).contiguous().chunk(
3, dim=-1)
# apply attention
x = F.scaled_dot_product_attention(
q,
k,
v,
)
x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
# output
x = self.proj(x)
x = rearrange(x, '(b t) c h w-> b c t h w', t=t)
return x + identity
class Encoder3d(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
# dimensions
dims = [dim * u for u in [1] + dim_mult]
scale = 1.0
# init block
self.conv1 = CausalConv3d(3, dims[0], 3, padding=1)
# downsample blocks
downsamples = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
for _ in range(num_res_blocks):
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
downsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
# downsample block
if i != len(dim_mult) - 1:
mode = 'downsample3d' if temperal_downsample[
i] else 'downsample2d'
downsamples.append(Resample(out_dim, mode=mode))
scale /= 2.0
self.downsamples = nn.Sequential(*downsamples)
# middle blocks
self.middle = nn.Sequential(
ResidualBlock(out_dim, out_dim, dropout), AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout))
# output blocks
self.head = nn.Sequential(
RMS_norm(out_dim, images=False), nn.SiLU(),
CausalConv3d(out_dim, z_dim, 3, padding=1))
def forward(self, x, feat_cache=None, feat_idx=[0]):
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
## downsamples
for layer in self.downsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## middle
for layer in self.middle:
if isinstance(layer, ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## head
for layer in self.head:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x
class Decoder3d(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_upsample=[False, True, True],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_upsample = temperal_upsample
# dimensions
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
scale = 1.0 / 2**(len(dim_mult) - 2)
# init block
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks
self.middle = nn.Sequential(
ResidualBlock(dims[0], dims[0], dropout), AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout))
# upsample blocks
upsamples = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
if i == 1 or i == 2 or i == 3:
in_dim = in_dim // 2
for _ in range(num_res_blocks + 1):
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
upsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
# upsample block
if i != len(dim_mult) - 1:
mode = 'upsample3d' if temperal_upsample[i] else 'upsample2d'
upsamples.append(Resample(out_dim, mode=mode))
scale *= 2.0
self.upsamples = nn.Sequential(*upsamples)
# output blocks
self.head = nn.Sequential(
RMS_norm(out_dim, images=False), nn.SiLU(),
CausalConv3d(out_dim, 3, 3, padding=1))
def forward(self, x, feat_cache=None, feat_idx=[0]):
## conv1
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
## middle
for layer in self.middle:
if isinstance(layer, ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## upsamples
for layer in self.upsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## head
for layer in self.head:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x
def count_conv3d(model):
count = 0
for m in model.modules():
if isinstance(m, CausalConv3d):
count += 1
return count
class WanVAE_(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
self.temperal_upsample = temperal_downsample[::-1]
# modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout)
def forward(self, x):
mu, log_var = self.encode(x)
z = self.reparameterize(mu, log_var)
x_recon = self.decode(z)
return x_recon, mu, log_var
def encode(self, x, scale=None, return_posterior=False):
self.clear_cache()
## cache
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
## 对encode输入的x,按时间拆分为1、4、4、4....
for i in range(iter_):
self._enc_conv_idx = [0]
if i == 0:
out = self.encoder(
x[:, :, :1, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
else:
out_ = self.encoder(
x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = torch.cat([out, out_], 2)
mu, log_var = self.conv1(out).chunk(2, dim=1)
if scale is None or return_posterior:
return mu, log_var
mu = self.reparameterize(mu, log_var)
if isinstance(scale[0], torch.Tensor):
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
mu = (mu - scale[0]) * scale[1]
self.clear_cache()
return mu
def decode(self, z, scale=None):
self.clear_cache()
# z: [b,c,t,h,w]
if scale is not None:
if isinstance(scale[0], torch.Tensor):
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
z = z / scale[1] + scale[0]
iter_ = z.shape[2]
x = self.conv2(z)
for i in range(iter_):
self._conv_idx = [0]
if i == 0:
out = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
else:
out_ = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2)
self.clear_cache()
return out
def reparameterize(self, mu, log_var):
std = torch.exp(0.5 * log_var)
eps = torch.randn_like(std)
return eps * std + mu
def sample(self, imgs, deterministic=False, scale=None):
mu, log_var = self.encode(imgs)
if deterministic:
return mu
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
mu = mu + std * torch.randn_like(std)
if isinstance(scale[0], torch.Tensor):
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
mu = (mu - scale[0]) * scale[1]
self.clear_cache()
return mu
def clear_cache(self):
self._conv_num = count_conv3d(self.decoder)
self._conv_idx = [0]
self._feat_map = [None] * self._conv_num
#cache encode
self._enc_conv_num = count_conv3d(self.encoder)
self._enc_conv_idx = [0]
self._enc_feat_map = [None] * self._enc_conv_num
def _video_vae(pretrained_path=None, z_dim=None, device='cpu', **kwargs):
"""
Autoencoder3d adapted from Stable Diffusion 1.x, 2.x and XL.
"""
# params
cfg = dict(
dim=96,
z_dim=z_dim,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0)
cfg.update(**kwargs)
# init model
with torch.device('meta'):
model = WanVAE_(**cfg)
# load checkpoint
logging.info(f'loading {pretrained_path}')
if pretrained_path.endswith('.safetensors'):
from safetensors.torch import load_file
pretrained_state_dict = load_file(pretrained_path, device='cpu')
else:
pretrained_state_dict = torch.load(pretrained_path, map_location='cpu')
model.load_state_dict(pretrained_state_dict, assign=True)
return model
class WanxVAE(nn.Module):
# @register_to_config
def __init__(self,
pretrained='',
torch_dtype=torch.float32,
device='cuda'
):
super().__init__()
self.dtype = torch_dtype
self.device = device
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
]
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.9160
]
self.mean = torch.tensor(mean, dtype=self.dtype, device=device)
self.std = torch.tensor(std, dtype=self.dtype, device=device)
self.scale = [self.mean, 1.0 / self.std]
self.config = lambda: None
self.config.latents_mean = self.mean
self.config.latents_std = self.std
self.ffactor_spatial = 8
self.ffactor_temporal = 4
self.config.latent_channels = 16
# init model
self.model = _video_vae(
pretrained_path=pretrained,
z_dim=16,
).eval().requires_grad_(False)
self.model = self.model.to(device=device, dtype=torch_dtype)
def encode(self, videos, return_posterior=False, **kwargs):
"""
videos: A list of videos each with shape [C, T, H, W].
"""
with amp.autocast(dtype=torch.float):
if return_posterior:
mus, log_vars = self.model.encode(
videos, scale=self.scale, return_posterior=True)
return mus, log_vars
else:
latents = self.model.encode(videos, scale=self.scale)
return latents
def decode(self, zs, **kwargs):
with amp.autocast(dtype=torch.float):
videos = [
self.model.decode(u.unsqueeze(0), scale=self.scale).clamp_(-1, 1).squeeze(0)
for u in zs
]
videos = torch.stack(videos, dim=0)
return (videos, )
+994
View File
@@ -0,0 +1,994 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# 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.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
import inspect
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
import torch
import torch.distributed as dist
import numpy as np
from dataclasses import dataclass
from packaging import version
from einops import rearrange
from typing import Any
from transformers import AutoProcessor
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.image_processor import VaeImageProcessor
from diffusers.models import AutoencoderKL
from diffusers.schedulers import KarrasDiffusionSchedulers
from diffusers.utils import (
logging,
replace_example_docstring,
)
from diffusers.utils.torch_utils import randn_tensor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput
from .mmdit.dit import Transformer3DModel
from .mmdit.dit.models import BlockGPUManager
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """"""
PRECISION_TO_TYPE = {
'fp32': torch.float32,
'fp16': torch.float16,
'bf16': torch.bfloat16,
}
def retrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] = None,
device: Optional[Union[str, torch.device]] = None,
timesteps: Optional[List[int]] = None,
sigmas: Optional[List[float]] = None,
**kwargs,
):
"""
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
Args:
scheduler (`SchedulerMixin`):
The scheduler to get timesteps from.
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
must be `None`.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
timesteps (`List[int]`, *optional*):
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
`num_inference_steps` and `sigmas` must be `None`.
sigmas (`List[float]`, *optional*):
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
`num_inference_steps` and `timesteps` must be `None`.
Returns:
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError(
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
)
if timesteps is not None:
accepts_timesteps = "timesteps" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys()
)
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(
inspect.signature(scheduler.set_timesteps).parameters.keys()
)
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
@dataclass
class PipelineOutput(BaseOutput):
videos: Union[torch.Tensor, np.ndarray]
class Pipeline(DiffusionPipeline):
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = ["latents", "prompt_embeds"]
def __init__(
self,
vae: AutoencoderKL,
text_encoder: Any,
tokenizer: Any,
transformer: Transformer3DModel,
scheduler: KarrasDiffusionSchedulers,
args=None,
):
super().__init__()
self.args = args
self.register_modules(
#vae=vae,
#text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer,
scheduler=scheduler,
)
self.enable_multi_task = getattr(
self.args, "enable_multi_task_training", False)
self.vae_scale_factor = 8
self.vae_scale_factor_temporal = 4
# if hasattr(self.vae, "ffactor_spatial"):
# self.vae_scale_factor = self.vae.ffactor_spatial
# self.vae_scale_factor_temporal = self.vae.ffactor_temporal
# else:
# self.vae_scale_factor = 2 ** (
# len(self.vae.config.block_out_channels) - 1)
# self.vae_scale_factor_temporal = 4 # hard code for HunyuanVideoVAE
self.image_processor = VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor)
# text_encoder_ckpt = dict(args.text_encoder_arch_config.get("params", {}))[
# 'text_encoder_ckpt']
self.qwen_processor = AutoProcessor.from_pretrained(self.args.repo)
self.text_token_max_length = self.args.text_token_max_length
self.prompt_template_encode = {
'image': "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n",
'multiple_images': "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n{}<|im_start|>assistant\n",
'video': "<|im_start|>system\n \\nDescribe the video by detailing the following aspects:\n1. The main content and theme of the video.\n2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects.\n3. Actions, events, behaviors temporal relationships, physical movement changes of the objects.\n4. background environment, light, style and atmosphere.\n5. camera angles, movements, and transitions used in the video:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
}
# [36:-4]
# [36:-4]
# [93:-4]
self.prompt_template_encode_start_idx = {
'image': 34,
'multiple_images': 34,
'video': 91,
}
def _extract_masked_hidden(self, hidden_states: torch.Tensor, mask: torch.Tensor):
bool_mask = mask.bool()
valid_lengths = bool_mask.sum(dim=1)
selected = hidden_states[bool_mask]
split_result = torch.split(selected, valid_lengths.tolist(), dim=0)
return split_result
def _get_qwen_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
template_type: str = 'image',
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
template = self.prompt_template_encode[template_type]
drop_idx = self.prompt_template_encode_start_idx[template_type]
txt = [template.format(e) for e in prompt]
txt_tokens = self.tokenizer(
txt, max_length=self.text_token_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt"
).to(device)
encoder_hidden_states = self.text_encoder(
input_ids=txt_tokens.input_ids,
attention_mask=txt_tokens.attention_mask,
output_hidden_states=True,
)
hidden_states = encoder_hidden_states.hidden_states[-1]
split_hidden_states = self._extract_masked_hidden(
hidden_states, txt_tokens.attention_mask)
split_hidden_states = [e[drop_idx:] for e in split_hidden_states]
attn_mask_list = [torch.ones(
e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states]
max_seq_len = min([
self.text_token_max_length,
max([u.size(0) for u in split_hidden_states]),
max([u.size(0) for u in attn_mask_list])])
prompt_embeds = torch.stack(
[torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))])
for u in split_hidden_states]
)
encoder_attention_mask = torch.stack(
[torch.cat([u, u.new_zeros(max_seq_len - u.size(0))])
for u in attn_mask_list]
)
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
return prompt_embeds, encoder_attention_mask
def encode_prompt_multiple_images(
self,
prompt: Union[str, List[str]],
device: Optional[torch.device] = None,
images: Optional[torch.Tensor] = None,
template_type: Optional[str] = 'multiple_images',
max_sequence_length: Optional[int] = None,
drop_vit_feature: Optional[float] = False,
):
assert template_type == 'multiple_images', "template_type must be 'multiple_images'"
device = device or self._execution_device
template = self.prompt_template_encode[template_type]
drop_idx = self.prompt_template_encode_start_idx[template_type]
prompt = [p.replace(
'<image>\n', '<|vision_start|><|image_pad|><|vision_end|>') for p in prompt]
prompt = [template.format(p) for p in prompt]
inputs = self.qwen_processor(
text=prompt,
images=images,
padding=True,
return_tensors="pt",
).to(device)
encoder_hidden_states = self.text_encoder(
**inputs,
output_hidden_states=True,
)
last_hidden_states = encoder_hidden_states.hidden_states[-1]
if drop_vit_feature:
input_ids = inputs['input_ids']
vlm_image_end_idx = torch.where(input_ids[0] == 151653)[0][-1]
drop_idx = vlm_image_end_idx + 1
prompt_embeds = last_hidden_states[:, drop_idx:]
prompt_embeds_mask = inputs['attention_mask'][:, drop_idx:]
if max_sequence_length is not None and prompt_embeds.shape[1] > max_sequence_length:
prompt_embeds = prompt_embeds[:, -max_sequence_length:, :]
prompt_embeds_mask = prompt_embeds_mask[:, -max_sequence_length:]
return prompt_embeds, prompt_embeds_mask
def encode_prompt_images(
self,
prompt: Union[str, List[str]],
device: Optional[torch.device] = None,
images: Optional[torch.Tensor] = None,
max_sequence_length: int = 1024,
):
r"""
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
device: (`torch.device`):
torch device
"""
device = device or self._execution_device
bs = images.shape[0]
if images.shape[2] == 1:
template = self.prompt_template_encode['image']
drop_idx = self.prompt_template_encode_start_idx['image']
prompt = template.replace(
'{}', '<|vision_start|><|image_pad|><|vision_end|>')
prompt = [prompt] * bs
inputs = self.qwen_processor(
text=prompt,
images=images.squeeze(2),
padding=True,
return_tensors="pt",
).to(device)
output_tensor = self.text_encoder(
**inputs,
output_hidden_states=True,
)
last_hidden_states = output_tensor.hidden_states[-1]
vis_hidden_states = last_hidden_states[:, drop_idx+3:-6]
patchify_size = self.qwen_processor.image_processor.merge_size
image_grid_thw = inputs['image_grid_thw']
vis_hidden_states = vis_hidden_states.view(
bs, image_grid_thw[0, 1]//patchify_size, image_grid_thw[0, 2]//patchify_size, -1)
return vis_hidden_states
def encode_prompt(
self,
prompt: Union[str, List[str]],
images: Optional[torch.Tensor] = None,
device: Optional[torch.device] = None,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_embeds_mask: Optional[torch.Tensor] = None,
max_sequence_length: int = 1024,
template_type: str = 'image',
drop_vit_feature: bool = False,
):
r"""
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
device: (`torch.device`):
torch device
num_videos_per_prompt (`int`):
number of videos that should be generated per prompt
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
"""
if images is not None:
##################################################
# from PIL import Image
# images = [img.resize((512, 512), Image.LANCZOS) for img in images]
##################################################
return self.encode_prompt_multiple_images(
prompt=prompt,
images=images,
device=device,
max_sequence_length=max_sequence_length,
drop_vit_feature=drop_vit_feature,
)
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(
prompt) if prompt_embeds is None else prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds, prompt_embeds_mask = self._get_qwen_prompt_embeds(
prompt, template_type, device)
prompt_embeds = prompt_embeds[:, :max_sequence_length]
prompt_embeds_mask = prompt_embeds_mask[:, :max_sequence_length]
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
batch_size * num_videos_per_prompt, seq_len, -1)
prompt_embeds_mask = prompt_embeds_mask.repeat(
1, num_videos_per_prompt, 1)
prompt_embeds_mask = prompt_embeds_mask.view(
batch_size * num_videos_per_prompt, seq_len)
return prompt_embeds, prompt_embeds_mask
def check_inputs(
self,
prompt,
height,
width,
images=None,
negative_prompt=None,
prompt_embeds=None,
negative_prompt_embeds=None,
prompt_embeds_mask=None,
negative_prompt_embeds_mask=None,
callback_on_step_end_tensor_inputs=None,
):
# if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0:
# logger.warning(
# f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly"
# )
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if prompt_embeds is not None and prompt_embeds_mask is None:
raise ValueError(
"If `prompt_embeds` are provided, `prompt_embeds_mask` also have to be passed. Make sure to generate `prompt_embeds_mask` from the same text encoder that was used to generate `prompt_embeds`."
)
if negative_prompt_embeds is not None and negative_prompt_embeds_mask is None:
raise ValueError(
"If `negative_prompt_embeds` are provided, `negative_prompt_embeds_mask` also have to be passed. Make sure to generate `negative_prompt_embeds_mask` from the same text encoder that was used to generate `negative_prompt_embeds`."
)
def prepare_conditions(self, latents: torch.Tensor, image=None, last_image=None):
"""
Prepare conditional inputs for video generation.
Args:
latents: Generated latent tensor with shape (B, N, C, T, latent_H, latent_W)
image: First frame condition, shape (B, N, 3, 1, H, W)
last_image: Last frame condition, shape (B, N, 3, 1, H, W)
Returns:
Combined condition tensor with shape (B, N, C+1, T, H, W)
"""
device, dtype = latents.device, latents.dtype
batch_size, num_items, latent_channels, latent_frames, latent_h, latent_w = latents.shape
# If no conditions provided, return zero condition
if image is None and last_image is None:
return torch.zeros(
batch_size, num_items, latent_channels + 1, latent_frames, latent_h, latent_w,
device=device, dtype=dtype
)
num_frame = (latent_frames - 1) * self.vae_scale_factor_temporal + 1
height = latent_h * self.vae_scale_factor
width = latent_w * self.vae_scale_factor
# Initialize mask
mask = torch.zeros(batch_size, num_items, 1, latent_frames,
latent_h, latent_w, device=device, dtype=dtype)
# Build video condition
if image is not None and last_image is not None:
# Both first and last frame conditions
image = image.to(device=device, dtype=dtype)
last_image = last_image.to(device=device, dtype=dtype)
middle_frames = torch.zeros(
batch_size, num_items, image.shape[2], num_frame -
2, height, width,
device=device, dtype=dtype
)
video_condition = torch.cat(
[image, middle_frames, last_image], dim=3)
mask[:, :, :, 0] = 1 # Mark first frame as conditional
mask[:, :, :, -1] = 1 # Mark last frame as conditional
elif image is not None:
# Only first frame condition
image = image.to(device=device, dtype=dtype)
remaining_frames = torch.zeros(
batch_size, num_items, image.shape[2], num_frame -
1, height, width,
device=device, dtype=dtype
)
video_condition = torch.cat([image, remaining_frames], dim=3)
mask[:, :, :, 0] = 1 # Mark first frame as conditional
else:
raise NotImplementedError
# VAE encode the video condition
video_condition = rearrange(
video_condition, "b n c t h w -> (b n) c t h w")
latent_condition = self.vae.encode(
video_condition).latent_dist.sample()
# Normalize
latent_condition = self.normalize_latents(latent_condition)
# Reshape back to (B, N, C, T, H, W)
latent_condition = rearrange(
latent_condition, "(b n) c t h w -> b n c t h w", b=batch_size)
# Concat
return torch.cat([latent_condition, mask], dim=2)
def prepare_latents(
self,
batch_size,
num_items,
num_channels_latents,
height,
width,
video_length,
dtype,
device,
generator,
latents=None,
reference_images=None,
image=None,
last_image=None,
):
shape = (
batch_size,
num_items,
num_channels_latents,
(video_length - 1) // self.vae_scale_factor_temporal + 1,
int(height) // self.vae_scale_factor,
int(width) // self.vae_scale_factor,
)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
if latents is None:
if reference_images is not None:
ref_img = [torch.from_numpy(
np.array(x.convert("RGB"))) for x in reference_images]
ref_img = torch.stack(ref_img).to(device=device, dtype=dtype)
ref_img = ref_img / 127.5 - 1.0
ref_img = rearrange(ref_img, "x h w c -> x c 1 h w")
ref_vae = self.vae.encode(ref_img)
ref_vae = rearrange(
ref_vae, "(b n) c 1 h w -> b n c 1 h w", n=(num_items - 1))
noise = randn_tensor(
(shape[0], 1, *shape[2:]),
generator=generator, device=device, dtype=dtype
)
latents = torch.cat([ref_vae, noise], dim=1)
else:
latents = randn_tensor(
shape, generator=generator, device=device, dtype=dtype
)
else:
latents = latents.to(device)
if not self.enable_multi_task:
return latents, None
# image: (b, n, c, 1, h, w), last_image: (b, n, c, 1, h, w)
condition = self.prepare_conditions(latents, image, last_image)
return latents, condition
@property
def guidance_scale(self):
return self._guidance_scale
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
# corresponds to doing no classifier free guidance.
@property
def do_classifier_free_guidance(self):
# return self._guidance_scale > 1 and self.transformer.config.time_cond_proj_dim is None
return self._guidance_scale > 1
@property
def num_timesteps(self):
return self._num_timesteps
@property
def interrupt(self):
return self._interrupt
def pad_sequence(self, x: torch.Tensor, target_length: int):
current_length = x.shape[1]
if current_length >= target_length:
return x[:, -target_length:]
padding_length = target_length - current_length
if x.ndim >= 3:
padding = torch.zeros(
(x.shape[0], padding_length, *x.shape[2:]),
dtype=x.dtype, device=x.device
)
else:
padding = torch.zeros(
(x.shape[0], padding_length),
dtype=x.dtype, device=x.device
)
return torch.cat([x, padding], dim=1)
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]],
height: int,
width: int,
num_frames: int,
images: Optional[torch.Tensor] = None,
image_condition: torch.Tensor | None = None,
last_image_condition: torch.Tensor | None = None,
data_type: str = "video",
num_inference_steps: int = 50,
timesteps: List[int] = None,
sigmas: List[float] = None,
guidance_scale: float = 7.5,
negative_prompt: Optional[Union[str, List[str]]] = None,
num_videos_per_prompt: Optional[int] = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator,
List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
prompt_embeds_mask: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds_mask: Optional[torch.Tensor] = None,
output_type: Optional[str] = "pil",
return_dict: bool = True,
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
guidance_rescale: float = 0.0,
clip_skip: Optional[int] = None,
callback_on_step_end: Optional[
Union[
Callable[[int, int, Dict], None],
PipelineCallback,
MultiPipelineCallbacks,
]
] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
enable_tiling: bool = False,
max_sequence_length: int = 4096,
drop_vit_feature: bool = False,
offload=False,
offload_block_num: int = 1,
lat=None,
**kwargs,
):
r"""
The call function to the pipeline for generation.
Args:
prompt (`str` or `List[str]`):
The prompt or prompts to guide image generation. If not defined, you need to pass `prompt_embeds`.
height (`int`):
The height in pixels of the generated image.
width (`int`):
The width in pixels of the generated image.
video_length (`int`):
The number of frames in the generated video.
num_inference_steps (`int`, *optional*, defaults to 50):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
timesteps (`List[int]`, *optional*):
Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument
in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is
passed will be used. Must be in descending order.
sigmas (`List[float]`, *optional*):
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
will be used.
guidance_scale (`float`, *optional*, defaults to 7.5):
A higher guidance scale value encourages the model to generate images closely linked to the text
`prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`.
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide what to not include in image generation. If not defined, you need to
pass `negative_prompt_embeds` instead. Ignored when not using guidance (`guidance_scale < 1`).
num_videos_per_prompt (`int`, *optional*, defaults to 1):
The number of images to generate per prompt.
eta (`float`, *optional*, defaults to 0.0):
Corresponds to parameter eta (η) from the [DDIM](https://arxiv.org/abs/2010.02502) paper. Only applies
to the [`~schedulers.DDIMScheduler`], and is ignored in other schedulers.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
generation deterministic.
latents (`torch.Tensor`, *optional*):
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor is generated by sampling using the supplied random `generator`.
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
provided, text embeddings are generated from the `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs (prompt weighting). If
not provided, `negative_prompt_embeds` are generated from the `negative_prompt` input argument.
output_type (`str`, *optional*, defaults to `"pil"`):
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`HunyuanVideoPipelineOutput`] instead of a
plain tuple.
cross_attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the [`AttentionProcessor`] as defined in
[`self.processor`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
guidance_rescale (`float`, *optional*, defaults to 0.0):
Guidance rescale factor from [Common Diffusion Noise Schedules and Sample Steps are
Flawed](https://arxiv.org/pdf/2305.08891.pdf). Guidance rescale factor should fix overexposure when
using zero terminal SNR.
clip_skip (`int`, *optional*):
Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that
the output of the pre-final layer will be used for computing the prompt embeddings.
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
callback_on_step_end_tensor_inputs (`List`, *optional*):
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
`._callback_tensor_inputs` attribute of your pipeline class.
Examples:
Returns:
[`~HunyuanVideoPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`HunyuanVideoPipelineOutput`] is returned,
otherwise a `tuple` is returned where the first element is a list with the generated images and the
second element is a list of `bool`s indicating whether the corresponding generated image contains
"not-safe-for-work" (nsfw) content.
"""
# 1. Check inputs. Raise error if not correct
# self.check_inputs(
# prompt,
# height,
# width,
# images=images,
# negative_prompt=negative_prompt,
# prompt_embeds=prompt_embeds,
# negative_prompt_embeds=negative_prompt_embeds,
# prompt_embeds_mask=prompt_embeds_mask,
# negative_prompt_embeds_mask=negative_prompt_embeds_mask,
# callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
# )
self._guidance_scale = guidance_scale
self._interrupt = False
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
print(f"check: {prompt_embeds.shape,prompt_embeds.device,prompt_embeds.dtype}") #check: (torch.Size([1, 1037, 4096]), device(type='cuda', index=0), torch.bfloat16)
device = prompt_embeds.device
# 3. Encode input prompt
template_type = 'image' if num_frames == 1 else "video"
num_items = 1 if images is None or len(
images) == 0 else 1 + len(images)
if prompt is not None:
prompt_embeds, prompt_embeds_mask = self.encode_prompt(
prompt=prompt,
prompt_embeds=prompt_embeds,
prompt_embeds_mask=prompt_embeds_mask,
images=images,
device=device,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
template_type=template_type,
drop_vit_feature=drop_vit_feature,
)
# For classifier free guidance, we need to do two forward passes.
# Here we concatenate the unconditional and text embeddings into a single batch
# to avoid doing two forward passes
if self.do_classifier_free_guidance:
if negative_prompt is None and negative_prompt_embeds is None:
# default_negative_prompt = 'low quality, jpeg artifacts, ugly, duplicate, morbid, mutilated, extra fingers, mutated hands, poorly drawn hands, poorly drawn face, mutation, deformed, blurry, dehydrated, bad anatomy, bad proportions, extra limbs, cloned face, disfigured, gross proportions, malformed limbs, missing arms, missing legs, extra arms, extra legs, fused fingers, too many fingers.' # noqa
default_negative_prompt = ""
if num_items <= 1:
negative_prompt = [
f"<|im_start|>user\n{default_negative_prompt}<|im_end|>\n"] * batch_size
else:
image_tokens = "<image>\n" * (num_items - 1)
negative_prompt = [
f"<|im_start|>user\n{image_tokens}{default_negative_prompt}<|im_end|>\n"] * batch_size
if negative_prompt is not None:
negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt(
prompt=negative_prompt,
prompt_embeds=negative_prompt_embeds,
prompt_embeds_mask=negative_prompt_embeds_mask,
images=images,
device=device,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
template_type=template_type,
)
max_seq_len = max(
prompt_embeds.shape[1], negative_prompt_embeds.shape[1])
prompt_embeds = torch.cat([
self.pad_sequence(negative_prompt_embeds, max_seq_len),
self.pad_sequence(prompt_embeds, max_seq_len)])
if prompt_embeds_mask is not None:
prompt_embeds_mask = torch.cat([
self.pad_sequence(
negative_prompt_embeds_mask, max_seq_len),
self.pad_sequence(prompt_embeds_mask, max_seq_len)])
# 4. Prepare timesteps
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
device,
timesteps,
sigmas,
)
print(prompt_embeds.shape,prompt_embeds_mask.shape,num_items) #torch.Size([2, 1037, 4096]) torch.Size([2, 1037]) 2
# 5. Prepare latent variables
num_channels_latents =16
#num_channels_latents = self.vae.config.latent_channels
if lat is None:
latents, condition = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_items,
num_channels_latents,
height,
width,
num_frames,
prompt_embeds.dtype,
device,
generator,
latents,
reference_images=images,
image=image_condition,
last_image=last_image_condition,
)
else:
latents = lat
condition = None
target_dtype = PRECISION_TO_TYPE[self.args.dit_precision]
autocast_enabled = (
target_dtype != torch.float32
)
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32
)
# 6. Denoising loop
num_warmup_steps = len(timesteps) - \
num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
if num_items > 1:
ref_latents = latents[:, :(num_items - 1)].clone()
if offload:
gpu_manager=BlockGPUManager()
gpu_manager.setup_for_inference(self.transformer)
else:
gpu_manager=None
# if is_progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
# copy reference latents
if num_items > 1:
latents[:, :(num_items - 1)] = ref_latents.clone()
# concat condition if enable multi-task
if condition is not None:
latents_ = torch.cat([latents, condition], dim=2)
else:
latents_ = latents
# expand the latents if we are doing classifier free guidance
latent_model_input = (
torch.cat([latents_] * 2)
if self.do_classifier_free_guidance
else latents_
)
t_expand = t.repeat(latent_model_input.shape[0])
# predict the noise residual
with torch.autocast(
device_type="cuda", dtype=target_dtype, enabled=autocast_enabled
):
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
# [2, 16, 33, 24, 42]
hidden_states=latent_model_input,
timestep=t_expand, # [2]
encoder_hidden_states=prompt_embeds, # [2, 256, 4096]
# [2, 256]
encoder_hidden_states_mask=prompt_embeds_mask,
return_dict=False,
gpu_manager=gpu_manager,
)[0]
if (noise_pred.isnan()).any() or (noise_pred.isinf()).any():
print("handle with nan/inf data")
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond
)
cond_norm = torch.norm(
noise_pred_text, dim=2, keepdim=True)
noise_norm = torch.norm(noise_pred, dim=2, keepdim=True)
noise_pred = noise_pred * (cond_norm / noise_norm)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(
noise_pred, t, latents, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(
self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop(
"prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop(
"negative_prompt_embeds", negative_prompt_embeds
)
# call the callback, if provided
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
if progress_bar is not None:
progress_bar.update()
if gpu_manager is not None:
gpu_manager.unload_all_blocks_to_cpu()
if not output_type == "latent":
if not (len(latents.shape) == 5 or len(latents.shape) == 6):
raise ValueError(
f"Only support latents with shape (b, n, c, h, w) or (b, n, c, f, h, w), but got {latents.shape}."
)
if (latents.isnan()).any() or (latents.isinf()).any():
print("handle with nan/inf data")
# Decode latents (VAE handles denormalization internally via scale)
latents = rearrange(
latents, "b n c f h w -> (b n) c f h w")
with torch.autocast(
device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled
):
image = self.vae.decode(
latents, return_dict=False
)[0]
image = rearrange(
image, "(b n) c f h w -> b n c f h w", b=batch_size)
# Replace the first frame with the image condition
if image_condition is not None:
image[:, :, :, :1] = image_condition
# Replace the last frame with the last image condition
if last_image_condition is not None:
image[:, :, :, -1:] = last_image_condition
else:
#print("return latents without decoding",latents.shape) #torch.Size([1, 1, 16, 1, 128, 128])
image = rearrange(latents, "b n c f h w -> (b n) c f h w")
#print(image.shape) #torch.Size([1, 16, 1, 128, 128])
return image
image = (image / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
image = image.cpu().float().permute(0, 1, 3, 2, 4, 5)
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return image
return PipelineOutput(videos=image)
+265
View File
@@ -0,0 +1,265 @@
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
#
# 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.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.schedulers.scheduling_utils import SchedulerMixin
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
"""
Euler scheduler.
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.
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.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
@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 _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def set_timesteps(
self,
num_inference_steps: int,
device: Union[str, torch.device] = None,
n_tokens: int = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
if not self.config.reverse:
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device
)
# Reset step index
self._step_index = None
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()
def _init_step_index(self, timestep):
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 scale_model_input(
self, sample: torch.Tensor, timestep: Optional[int] = None
) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
# print("sd3:self.config.shift",self.config.shift)
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def flux_time_shift(self, t: torch.Tensor):
# print("flux2:self.config.shift",self.config.shift)
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
pass
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return (prev_sample,)
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
+65
View File
@@ -0,0 +1,65 @@
import os
import random
import re
from pathlib import Path
import numpy as np
from PIL import Image
import torch
import torch.distributed as dist
def seed_everything(seed: int | None = None) -> None:
if seed is not None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
# ---------------------------------------------------------------------------
# Distributed helpers (replaces modules.distributed.parallel_states)
# ---------------------------------------------------------------------------
def maybe_init_distributed() -> bool:
"""Initialize torch distributed if WORLD_SIZE > 1. Returns True if initialized."""
world_size = int(os.environ.get('WORLD_SIZE', '1'))
if world_size <= 1:
return False
rank = int(os.environ.get('RANK', '0'))
dist.init_process_group(backend='nccl', world_size=world_size, rank=rank)
return True
def clean_dist_env() -> None:
"""Destroy the distributed process group if it was initialized."""
if dist.is_initialized():
dist.destroy_process_group()
def _dynamic_resize_from_bucket(image: Image, basesize: int = 512):
from ..models.bucket import BucketGroup, generate_video_image_bucket
from typing import Tuple
import math
import torchvision.transforms.functional as TF
def resize_center_crop(img: Image.Image, target_size: Tuple[int, int]) -> Image.Image:
"""等比缩放到 >= 目标尺寸,再中心裁剪到目标尺寸。(PIL输入/输出)"""
w, h = img.size # PIL: (width, height)
bh, bw = target_size
scale = max(bh / h, bw / w)
resize_h, resize_w = math.ceil(h * scale), math.ceil(w * scale)
img = TF.resize(img, (resize_h, resize_w),
interpolation=TF.InterpolationMode.BILINEAR, antialias=True)
img = TF.center_crop(img, target_size)
return img
bucket_config = generate_video_image_bucket(
basesize=basesize, min_temporal=56, max_temporal=56, bs_img=4, bs_vid=4, bs_mimg=8, min_items=2, max_items=2
)
bucket_group = BucketGroup(bucket_config)
img_w, img_h = image.size
bucket = bucket_group.find_best_bucket((1, 1, img_h, img_w))
target_height, target_width = bucket[-2], bucket[-1] # (height, width)
img_proc = resize_center_crop(image, (target_height, target_width))
return img_proc
+7
View File
@@ -0,0 +1,7 @@
import torch
PRECISION_TO_TYPE = {
"fp32": torch.float32,
"fp16": torch.float16,
"bf16": torch.bfloat16,
}
+208
View File
@@ -0,0 +1,208 @@
# Adapted from https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/models/loader/
from typing import Generator
import os
import contextlib
from collections.abc import Generator, Callable
from tqdm import tqdm
import torch
from torch import nn
from torch.distributed import init_device_mesh, DeviceMesh
from torch.distributed.checkpoint.state_dict import set_model_state_dict, get_model_state_dict, StateDictOptions
from torch.distributed.fsdp import CPUOffloadPolicy, MixedPrecisionPolicy, fully_shard
from safetensors.torch import safe_open
from .logging import get_logger
# TODO(PY): move this to utils elsewhere
@contextlib.contextmanager
def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
"""
Context manager to set torch's default dtype.
Args:
dtype (torch.dtype): The desired default dtype inside the context manager.
Returns:
ContextManager: context manager for setting default dtype.
Example:
>>> with set_default_dtype(torch.bfloat16):
>>> x = torch.tensor([1, 2, 3])
>>> x.dtype
torch.bfloat16
"""
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(dtype)
try:
yield
finally:
torch.set_default_dtype(old_dtype)
# explicitly use pure text format, with a newline at the end
# this makes it impossible to see the animation in the progress bar
# but will avoid messing up with ray or multiprocessing, which wraps
# each line of output with some prefix.
_BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]\n" # noqa: E501
def safetensors_weights_iterator(hf_weights_files: list[str]) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model safetensor files."""
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0
device = "cpu"
for st_file in tqdm(
hf_weights_files,
desc="Loading safetensors checkpoint shards",
disable=not enable_tqdm,
bar_format=_BAR_FORMAT,
):
with safe_open(st_file, framework="pt", device=device) as f:
for name in f.keys(): # noqa: SIM118
param = f.get_tensor(name)
yield name, param
def pt_weights_iterator(hf_weights_files: list[str]) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model bin/pt files."""
device = "cpu"
enable_tqdm = not torch.distributed.is_initialized(
) or torch.distributed.get_rank() == 0
for bin_file in tqdm(
hf_weights_files,
desc="Loading pt checkpoint shards",
disable=not enable_tqdm,
bar_format=_BAR_FORMAT,
):
state = torch.load(bin_file, map_location=device, weights_only=True)
yield from state.items()
del state
def maybe_load_fsdp_model(
model: nn.Module,
hsdp_shard_dim: int,
reshard_after_forward: bool,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
fsdp_inference: bool = False,
output_dtype: torch.dtype | None = None,
training_mode: bool = True,
pin_cpu_memory: bool = True,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
"""
logger = get_logger()
mp_policy = MixedPrecisionPolicy(param_dtype,
reduce_dtype,
output_dtype,
cast_forward_inputs=False)
# Check if we should use FSDP
world_size = int(os.getenv("WORLD_SIZE", "1"))
assert world_size % hsdp_shard_dim == 0, f"world_size {world_size} must be divisible by hsdp_shard_dim {hsdp_shard_dim}"
hsdp_replicate_dim = world_size // hsdp_shard_dim
use_fsdp = training_mode or fsdp_inference
if hsdp_shard_dim * hsdp_replicate_dim <= 1:
use_fsdp = False
logger.warning(
f"hsdp_replicate_dim * hsdp_shard_dim = {hsdp_replicate_dim}x{hsdp_shard_dim} <= 1, not using FSDP.")
if use_fsdp:
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=reshard_after_forward,
mp_policy=mp_policy,
mesh=device_mesh,
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
return model
def shard_model(
model,
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: MixedPrecisionPolicy | None = MixedPrecisionPolicy(), # noqa
mesh: DeviceMesh | None = None,
fsdp_shard_conditions: list[Callable[[str, nn.Module], bool]] = [], # noqa
pin_cpu_memory: bool = True,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
This method will over the model's named modules from the bottom-up and apply shard modules
based on whether they meet any of the criteria from shard_conditions.
Args:
model (TransformerDecoder): Model to shard with FSDP.
shard_conditions (List[Callable[[str, nn.Module], bool]]): A list of functions to determine
which modules to shard with FSDP. Each function should take module name (relative to root)
and the module itself, returning True if FSDP should shard the module and False otherwise.
If any of shard_conditions return True for a given module, it will be sharded by FSDP.
cpu_offload (bool): If set to True, FSDP will offload parameters, gradients, and optimizer
states to CPU.
reshard_after_forward (bool): Whether to reshard parameters and buffers after
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
Default to None.
fsdp_shard_conditions (List[Callable[[str, nn.Module], bool]]): A list of functions to determine
which modules to shard with FSDP.
pin_cpu_memory (bool): If set to True, FSDP will pin the CPU memory of the offloaded parameters.
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
"""
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
logger = get_logger()
logger.warning(
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
type(model).__name__)
return
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
pin_memory=pin_cpu_memory)
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
for shard_condition in fsdp_shard_conditions
]):
fully_shard(m, **fsdp_kwargs)
num_layers_sharded += 1
if num_layers_sharded == 0:
raise ValueError(
"No layer modules were sharded. Please check if shard conditions are working as expected."
)
# Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs)
+38
View File
@@ -0,0 +1,38 @@
import os
import loguru
_logger = loguru.logger
class NullLogger:
def __getattr__(self, name):
return lambda *args, **kwargs: None
def bind(self, **kwargs):
return self
def setup_logger(exp_dir: str):
global _logger
if int(os.getenv("RANK", 0)) <= 0:
_logger.add(
os.path.join(exp_dir, "train.log"),
level="DEBUG",
colorize=False,
backtrace=True,
diagnose=True,
encoding="utf-8",
)
else:
_logger = NullLogger()
_logger.info(f"Experiment directory created at: {exp_dir}")
return _logger
def get_logger():
return _logger
__all__ = ["setup_logger", "get_logger"]
+27
View File
@@ -0,0 +1,27 @@
import importlib
def get_obj_from_str(string, reload=False):
module, cls = string.rsplit(".", 1)
if reload:
module_imp = importlib.import_module(module)
importlib.reload(module_imp)
return getattr(importlib.import_module(module, package=None), cls)
def build_from_config(config, **kwargs):
if "target" not in config:
if config in ("__is_first_stage__", "__is_unconditional__"):
return None
raise KeyError("Expected key `target` to instantiate.")
cls = get_obj_from_str(config["target"])
params = dict(config.get("params", {}))
params.update(kwargs)
pretrained_path = config.get("pretrained", None)
if pretrained_path is not None and hasattr(cls, "from_pretrained"):
return cls.from_pretrained(pretrained_path, **params)
obj = cls(**params)
return obj
Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 300 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 506 KiB