@@ -0,0 +1,7 @@
|
||||
__pycache__/
|
||||
*.pyc
|
||||
|
||||
.pytest_cache/
|
||||
|
||||
test_outputs/
|
||||
outputs/
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 %}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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()
|
||||
@@ -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
|
||||
----
|
||||
|
||||

|
||||

|
||||

|
||||
|
||||
5.Citation
|
||||
----
|
||||
|
||||
```
|
||||
@article{JoyAI2023,}
|
||||
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
|
||||
from .JoyAI_Image_node import *
|
||||
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 681 KiB |
|
After Width: | Height: | Size: 4.2 MiB |
|
After Width: | Height: | Size: 537 KiB |
|
After Width: | Height: | Size: 9.6 MiB |
|
After Width: | Height: | Size: 598 KiB |
|
After Width: | Height: | Size: 795 KiB |
|
After Width: | Height: | Size: 342 KiB |
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]}")
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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 = []
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
__all__: list[str] = []
|
||||
@@ -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
|
||||
@@ -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}'.")
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,4 @@
|
||||
from .models import Transformer3DModel
|
||||
|
||||
|
||||
__all__ = ["Transformer3DModel"]
|
||||
@@ -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
|
||||
@@ -0,0 +1,3 @@
|
||||
from .wanvae import WanxVAE
|
||||
|
||||
__all__ = ["WanxVAE"]
|
||||
@@ -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, )
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,7 @@
|
||||
import torch
|
||||
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 300 KiB |
|
After Width: | Height: | Size: 506 KiB |