Spaces:
Runtime error
Runtime error
Migrated files batch 30
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- diffsynth/models/__init__.py +1 -0
- diffsynth/models/attention.py +89 -0
- diffsynth/models/cog_dit.py +408 -0
- diffsynth/models/cog_vae.py +518 -0
- diffsynth/models/downloader.py +111 -0
- diffsynth/models/flux_controlnet.py +331 -0
- diffsynth/models/flux_dit.py +749 -0
- diffsynth/models/flux_infiniteyou.py +129 -0
- diffsynth/models/flux_ipadapter.py +94 -0
- diffsynth/models/flux_lora_encoder.py +111 -0
- diffsynth/models/flux_text_encoder.py +32 -0
- diffsynth/models/flux_vae.py +303 -0
- diffsynth/models/flux_value_control.py +60 -0
- diffsynth/models/hunyuan_dit.py +451 -0
- diffsynth/models/hunyuan_dit_text_encoder.py +163 -0
- diffsynth/models/hunyuan_video_dit.py +920 -0
- diffsynth/models/hunyuan_video_text_encoder.py +68 -0
- diffsynth/models/hunyuan_video_vae_decoder.py +507 -0
- diffsynth/models/hunyuan_video_vae_encoder.py +307 -0
- diffsynth/models/kolors_text_encoder.py +1551 -0
- diffsynth/models/lora.py +387 -0
- diffsynth/models/model_manager.py +467 -0
- diffsynth/models/nexus_gen.py +161 -0
- diffsynth/models/nexus_gen_ar_model.py +1143 -0
- diffsynth/models/nexus_gen_projector.py +417 -0
- diffsynth/models/omnigen.py +803 -0
- diffsynth/models/qwen_image_dit.py +357 -0
- diffsynth/models/qwen_image_text_encoder.py +255 -0
- diffsynth/models/qwen_image_vae.py +736 -0
- diffsynth/models/qwenvl.py +168 -0
- diffsynth/models/sd3_dit.py +567 -0
- diffsynth/models/sd3_text_encoder.py +0 -0
- diffsynth/models/sd3_vae_decoder.py +81 -0
- diffsynth/models/sd3_vae_encoder.py +95 -0
- diffsynth/models/sd_controlnet.py +589 -0
- diffsynth/models/sd_ipadapter.py +57 -0
- diffsynth/models/sd_motion.py +199 -0
- diffsynth/models/sd_text_encoder.py +321 -0
- diffsynth/models/sd_unet.py +0 -0
- diffsynth/models/sd_vae_decoder.py +336 -0
- diffsynth/models/sd_vae_encoder.py +282 -0
- diffsynth/models/sdxl_controlnet.py +318 -0
- diffsynth/models/sdxl_ipadapter.py +122 -0
- diffsynth/models/sdxl_motion.py +104 -0
- diffsynth/models/sdxl_text_encoder.py +759 -0
- diffsynth/models/sdxl_unet.py +0 -0
- diffsynth/models/sdxl_vae_decoder.py +24 -0
- diffsynth/models/sdxl_vae_encoder.py +24 -0
- diffsynth/models/step1x_connector.py +683 -0
- diffsynth/models/stepvideo_dit.py +940 -0
diffsynth/models/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
from .model_manager import *
|
diffsynth/models/attention.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from einops import rearrange
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def low_version_attention(query, key, value, attn_bias=None):
|
| 6 |
+
scale = 1 / query.shape[-1] ** 0.5
|
| 7 |
+
query = query * scale
|
| 8 |
+
attn = torch.matmul(query, key.transpose(-2, -1))
|
| 9 |
+
if attn_bias is not None:
|
| 10 |
+
attn = attn + attn_bias
|
| 11 |
+
attn = attn.softmax(-1)
|
| 12 |
+
return attn @ value
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class Attention(torch.nn.Module):
|
| 16 |
+
|
| 17 |
+
def __init__(self, q_dim, num_heads, head_dim, kv_dim=None, bias_q=False, bias_kv=False, bias_out=False):
|
| 18 |
+
super().__init__()
|
| 19 |
+
dim_inner = head_dim * num_heads
|
| 20 |
+
kv_dim = kv_dim if kv_dim is not None else q_dim
|
| 21 |
+
self.num_heads = num_heads
|
| 22 |
+
self.head_dim = head_dim
|
| 23 |
+
|
| 24 |
+
self.to_q = torch.nn.Linear(q_dim, dim_inner, bias=bias_q)
|
| 25 |
+
self.to_k = torch.nn.Linear(kv_dim, dim_inner, bias=bias_kv)
|
| 26 |
+
self.to_v = torch.nn.Linear(kv_dim, dim_inner, bias=bias_kv)
|
| 27 |
+
self.to_out = torch.nn.Linear(dim_inner, q_dim, bias=bias_out)
|
| 28 |
+
|
| 29 |
+
def interact_with_ipadapter(self, hidden_states, q, ip_k, ip_v, scale=1.0):
|
| 30 |
+
batch_size = q.shape[0]
|
| 31 |
+
ip_k = ip_k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 32 |
+
ip_v = ip_v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 33 |
+
ip_hidden_states = torch.nn.functional.scaled_dot_product_attention(q, ip_k, ip_v)
|
| 34 |
+
hidden_states = hidden_states + scale * ip_hidden_states
|
| 35 |
+
return hidden_states
|
| 36 |
+
|
| 37 |
+
def torch_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None):
|
| 38 |
+
if encoder_hidden_states is None:
|
| 39 |
+
encoder_hidden_states = hidden_states
|
| 40 |
+
|
| 41 |
+
batch_size = encoder_hidden_states.shape[0]
|
| 42 |
+
|
| 43 |
+
q = self.to_q(hidden_states)
|
| 44 |
+
k = self.to_k(encoder_hidden_states)
|
| 45 |
+
v = self.to_v(encoder_hidden_states)
|
| 46 |
+
|
| 47 |
+
q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 48 |
+
k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 49 |
+
v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 50 |
+
|
| 51 |
+
if qkv_preprocessor is not None:
|
| 52 |
+
q, k, v = qkv_preprocessor(q, k, v)
|
| 53 |
+
|
| 54 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 55 |
+
if ipadapter_kwargs is not None:
|
| 56 |
+
hidden_states = self.interact_with_ipadapter(hidden_states, q, **ipadapter_kwargs)
|
| 57 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 58 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 59 |
+
|
| 60 |
+
hidden_states = self.to_out(hidden_states)
|
| 61 |
+
|
| 62 |
+
return hidden_states
|
| 63 |
+
|
| 64 |
+
def xformers_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None):
|
| 65 |
+
if encoder_hidden_states is None:
|
| 66 |
+
encoder_hidden_states = hidden_states
|
| 67 |
+
|
| 68 |
+
q = self.to_q(hidden_states)
|
| 69 |
+
k = self.to_k(encoder_hidden_states)
|
| 70 |
+
v = self.to_v(encoder_hidden_states)
|
| 71 |
+
|
| 72 |
+
q = rearrange(q, "b f (n d) -> (b n) f d", n=self.num_heads)
|
| 73 |
+
k = rearrange(k, "b f (n d) -> (b n) f d", n=self.num_heads)
|
| 74 |
+
v = rearrange(v, "b f (n d) -> (b n) f d", n=self.num_heads)
|
| 75 |
+
|
| 76 |
+
if attn_mask is not None:
|
| 77 |
+
hidden_states = low_version_attention(q, k, v, attn_bias=attn_mask)
|
| 78 |
+
else:
|
| 79 |
+
import xformers.ops as xops
|
| 80 |
+
hidden_states = xops.memory_efficient_attention(q, k, v)
|
| 81 |
+
hidden_states = rearrange(hidden_states, "(b n) f d -> b f (n d)", n=self.num_heads)
|
| 82 |
+
|
| 83 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 84 |
+
hidden_states = self.to_out(hidden_states)
|
| 85 |
+
|
| 86 |
+
return hidden_states
|
| 87 |
+
|
| 88 |
+
def forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None):
|
| 89 |
+
return self.torch_forward(hidden_states, encoder_hidden_states=encoder_hidden_states, attn_mask=attn_mask, ipadapter_kwargs=ipadapter_kwargs, qkv_preprocessor=qkv_preprocessor)
|
diffsynth/models/cog_dit.py
ADDED
|
@@ -0,0 +1,408 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from einops import rearrange, repeat
|
| 3 |
+
from .sd3_dit import TimestepEmbeddings
|
| 4 |
+
from .attention import Attention
|
| 5 |
+
from .utils import load_state_dict_from_folder
|
| 6 |
+
from .tiler import TileWorker2Dto3D
|
| 7 |
+
import numpy as np
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class CogPatchify(torch.nn.Module):
|
| 12 |
+
def __init__(self, dim_in, dim_out, patch_size) -> None:
|
| 13 |
+
super().__init__()
|
| 14 |
+
self.proj = torch.nn.Conv3d(dim_in, dim_out, kernel_size=(1, patch_size, patch_size), stride=(1, patch_size, patch_size))
|
| 15 |
+
|
| 16 |
+
def forward(self, hidden_states):
|
| 17 |
+
hidden_states = self.proj(hidden_states)
|
| 18 |
+
hidden_states = rearrange(hidden_states, "B C T H W -> B (T H W) C")
|
| 19 |
+
return hidden_states
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class CogAdaLayerNorm(torch.nn.Module):
|
| 24 |
+
def __init__(self, dim, dim_cond, single=False):
|
| 25 |
+
super().__init__()
|
| 26 |
+
self.single = single
|
| 27 |
+
self.linear = torch.nn.Linear(dim_cond, dim * (2 if single else 6))
|
| 28 |
+
self.norm = torch.nn.LayerNorm(dim, elementwise_affine=True, eps=1e-5)
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def forward(self, hidden_states, prompt_emb, emb):
|
| 32 |
+
emb = self.linear(torch.nn.functional.silu(emb))
|
| 33 |
+
if self.single:
|
| 34 |
+
shift, scale = emb.unsqueeze(1).chunk(2, dim=2)
|
| 35 |
+
hidden_states = self.norm(hidden_states) * (1 + scale) + shift
|
| 36 |
+
return hidden_states
|
| 37 |
+
else:
|
| 38 |
+
shift_a, scale_a, gate_a, shift_b, scale_b, gate_b = emb.unsqueeze(1).chunk(6, dim=2)
|
| 39 |
+
hidden_states = self.norm(hidden_states) * (1 + scale_a) + shift_a
|
| 40 |
+
prompt_emb = self.norm(prompt_emb) * (1 + scale_b) + shift_b
|
| 41 |
+
return hidden_states, prompt_emb, gate_a, gate_b
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
class CogDiTBlock(torch.nn.Module):
|
| 46 |
+
def __init__(self, dim, dim_cond, num_heads):
|
| 47 |
+
super().__init__()
|
| 48 |
+
self.norm1 = CogAdaLayerNorm(dim, dim_cond)
|
| 49 |
+
self.attn1 = Attention(q_dim=dim, num_heads=48, head_dim=dim//num_heads, bias_q=True, bias_kv=True, bias_out=True)
|
| 50 |
+
self.norm_q = torch.nn.LayerNorm((dim//num_heads,), eps=1e-06, elementwise_affine=True)
|
| 51 |
+
self.norm_k = torch.nn.LayerNorm((dim//num_heads,), eps=1e-06, elementwise_affine=True)
|
| 52 |
+
|
| 53 |
+
self.norm2 = CogAdaLayerNorm(dim, dim_cond)
|
| 54 |
+
self.ff = torch.nn.Sequential(
|
| 55 |
+
torch.nn.Linear(dim, dim*4),
|
| 56 |
+
torch.nn.GELU(approximate="tanh"),
|
| 57 |
+
torch.nn.Linear(dim*4, dim)
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def apply_rotary_emb(self, x, freqs_cis):
|
| 62 |
+
cos, sin = freqs_cis # [S, D]
|
| 63 |
+
cos = cos[None, None]
|
| 64 |
+
sin = sin[None, None]
|
| 65 |
+
cos, sin = cos.to(x.device), sin.to(x.device)
|
| 66 |
+
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
|
| 67 |
+
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
| 68 |
+
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
|
| 69 |
+
return out
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def process_qkv(self, q, k, v, image_rotary_emb, text_seq_length):
|
| 73 |
+
q = self.norm_q(q)
|
| 74 |
+
k = self.norm_k(k)
|
| 75 |
+
q[:, :, text_seq_length:] = self.apply_rotary_emb(q[:, :, text_seq_length:], image_rotary_emb)
|
| 76 |
+
k[:, :, text_seq_length:] = self.apply_rotary_emb(k[:, :, text_seq_length:], image_rotary_emb)
|
| 77 |
+
return q, k, v
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def forward(self, hidden_states, prompt_emb, time_emb, image_rotary_emb):
|
| 81 |
+
# Attention
|
| 82 |
+
norm_hidden_states, norm_encoder_hidden_states, gate_a, gate_b = self.norm1(
|
| 83 |
+
hidden_states, prompt_emb, time_emb
|
| 84 |
+
)
|
| 85 |
+
attention_io = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
|
| 86 |
+
attention_io = self.attn1(
|
| 87 |
+
attention_io,
|
| 88 |
+
qkv_preprocessor=lambda q, k, v: self.process_qkv(q, k, v, image_rotary_emb, prompt_emb.shape[1])
|
| 89 |
+
)
|
| 90 |
+
|
| 91 |
+
hidden_states = hidden_states + gate_a * attention_io[:, prompt_emb.shape[1]:]
|
| 92 |
+
prompt_emb = prompt_emb + gate_b * attention_io[:, :prompt_emb.shape[1]]
|
| 93 |
+
|
| 94 |
+
# Feed forward
|
| 95 |
+
norm_hidden_states, norm_encoder_hidden_states, gate_a, gate_b = self.norm2(
|
| 96 |
+
hidden_states, prompt_emb, time_emb
|
| 97 |
+
)
|
| 98 |
+
ff_io = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
|
| 99 |
+
ff_io = self.ff(ff_io)
|
| 100 |
+
|
| 101 |
+
hidden_states = hidden_states + gate_a * ff_io[:, prompt_emb.shape[1]:]
|
| 102 |
+
prompt_emb = prompt_emb + gate_b * ff_io[:, :prompt_emb.shape[1]]
|
| 103 |
+
|
| 104 |
+
return hidden_states, prompt_emb
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class CogDiT(torch.nn.Module):
|
| 109 |
+
def __init__(self):
|
| 110 |
+
super().__init__()
|
| 111 |
+
self.patchify = CogPatchify(16, 3072, 2)
|
| 112 |
+
self.time_embedder = TimestepEmbeddings(3072, 512)
|
| 113 |
+
self.context_embedder = torch.nn.Linear(4096, 3072)
|
| 114 |
+
self.blocks = torch.nn.ModuleList([CogDiTBlock(3072, 512, 48) for _ in range(42)])
|
| 115 |
+
self.norm_final = torch.nn.LayerNorm((3072,), eps=1e-05, elementwise_affine=True)
|
| 116 |
+
self.norm_out = CogAdaLayerNorm(3072, 512, single=True)
|
| 117 |
+
self.proj_out = torch.nn.Linear(3072, 64, bias=True)
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def get_resize_crop_region_for_grid(self, src, tgt_width, tgt_height):
|
| 121 |
+
tw = tgt_width
|
| 122 |
+
th = tgt_height
|
| 123 |
+
h, w = src
|
| 124 |
+
r = h / w
|
| 125 |
+
if r > (th / tw):
|
| 126 |
+
resize_height = th
|
| 127 |
+
resize_width = int(round(th / h * w))
|
| 128 |
+
else:
|
| 129 |
+
resize_width = tw
|
| 130 |
+
resize_height = int(round(tw / w * h))
|
| 131 |
+
|
| 132 |
+
crop_top = int(round((th - resize_height) / 2.0))
|
| 133 |
+
crop_left = int(round((tw - resize_width) / 2.0))
|
| 134 |
+
|
| 135 |
+
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def get_3d_rotary_pos_embed(
|
| 139 |
+
self, embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
|
| 140 |
+
):
|
| 141 |
+
start, stop = crops_coords
|
| 142 |
+
grid_h = np.linspace(start[0], stop[0], grid_size[0], endpoint=False, dtype=np.float32)
|
| 143 |
+
grid_w = np.linspace(start[1], stop[1], grid_size[1], endpoint=False, dtype=np.float32)
|
| 144 |
+
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
|
| 145 |
+
|
| 146 |
+
# Compute dimensions for each axis
|
| 147 |
+
dim_t = embed_dim // 4
|
| 148 |
+
dim_h = embed_dim // 8 * 3
|
| 149 |
+
dim_w = embed_dim // 8 * 3
|
| 150 |
+
|
| 151 |
+
# Temporal frequencies
|
| 152 |
+
freqs_t = 1.0 / (theta ** (torch.arange(0, dim_t, 2).float() / dim_t))
|
| 153 |
+
grid_t = torch.from_numpy(grid_t).float()
|
| 154 |
+
freqs_t = torch.einsum("n , f -> n f", grid_t, freqs_t)
|
| 155 |
+
freqs_t = freqs_t.repeat_interleave(2, dim=-1)
|
| 156 |
+
|
| 157 |
+
# Spatial frequencies for height and width
|
| 158 |
+
freqs_h = 1.0 / (theta ** (torch.arange(0, dim_h, 2).float() / dim_h))
|
| 159 |
+
freqs_w = 1.0 / (theta ** (torch.arange(0, dim_w, 2).float() / dim_w))
|
| 160 |
+
grid_h = torch.from_numpy(grid_h).float()
|
| 161 |
+
grid_w = torch.from_numpy(grid_w).float()
|
| 162 |
+
freqs_h = torch.einsum("n , f -> n f", grid_h, freqs_h)
|
| 163 |
+
freqs_w = torch.einsum("n , f -> n f", grid_w, freqs_w)
|
| 164 |
+
freqs_h = freqs_h.repeat_interleave(2, dim=-1)
|
| 165 |
+
freqs_w = freqs_w.repeat_interleave(2, dim=-1)
|
| 166 |
+
|
| 167 |
+
# Broadcast and concatenate tensors along specified dimension
|
| 168 |
+
def broadcast(tensors, dim=-1):
|
| 169 |
+
num_tensors = len(tensors)
|
| 170 |
+
shape_lens = {len(t.shape) for t in tensors}
|
| 171 |
+
assert len(shape_lens) == 1, "tensors must all have the same number of dimensions"
|
| 172 |
+
shape_len = list(shape_lens)[0]
|
| 173 |
+
dim = (dim + shape_len) if dim < 0 else dim
|
| 174 |
+
dims = list(zip(*(list(t.shape) for t in tensors)))
|
| 175 |
+
expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
|
| 176 |
+
assert all(
|
| 177 |
+
[*(len(set(t[1])) <= 2 for t in expandable_dims)]
|
| 178 |
+
), "invalid dimensions for broadcastable concatenation"
|
| 179 |
+
max_dims = [(t[0], max(t[1])) for t in expandable_dims]
|
| 180 |
+
expanded_dims = [(t[0], (t[1],) * num_tensors) for t in max_dims]
|
| 181 |
+
expanded_dims.insert(dim, (dim, dims[dim]))
|
| 182 |
+
expandable_shapes = list(zip(*(t[1] for t in expanded_dims)))
|
| 183 |
+
tensors = [t[0].expand(*t[1]) for t in zip(tensors, expandable_shapes)]
|
| 184 |
+
return torch.cat(tensors, dim=dim)
|
| 185 |
+
|
| 186 |
+
freqs = broadcast((freqs_t[:, None, None, :], freqs_h[None, :, None, :], freqs_w[None, None, :, :]), dim=-1)
|
| 187 |
+
|
| 188 |
+
t, h, w, d = freqs.shape
|
| 189 |
+
freqs = freqs.view(t * h * w, d)
|
| 190 |
+
|
| 191 |
+
# Generate sine and cosine components
|
| 192 |
+
sin = freqs.sin()
|
| 193 |
+
cos = freqs.cos()
|
| 194 |
+
|
| 195 |
+
if use_real:
|
| 196 |
+
return cos, sin
|
| 197 |
+
else:
|
| 198 |
+
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
|
| 199 |
+
return freqs_cis
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
def prepare_rotary_positional_embeddings(
|
| 203 |
+
self,
|
| 204 |
+
height: int,
|
| 205 |
+
width: int,
|
| 206 |
+
num_frames: int,
|
| 207 |
+
device: torch.device,
|
| 208 |
+
):
|
| 209 |
+
grid_height = height // 2
|
| 210 |
+
grid_width = width // 2
|
| 211 |
+
base_size_width = 720 // (8 * 2)
|
| 212 |
+
base_size_height = 480 // (8 * 2)
|
| 213 |
+
|
| 214 |
+
grid_crops_coords = self.get_resize_crop_region_for_grid(
|
| 215 |
+
(grid_height, grid_width), base_size_width, base_size_height
|
| 216 |
+
)
|
| 217 |
+
freqs_cos, freqs_sin = self.get_3d_rotary_pos_embed(
|
| 218 |
+
embed_dim=64,
|
| 219 |
+
crops_coords=grid_crops_coords,
|
| 220 |
+
grid_size=(grid_height, grid_width),
|
| 221 |
+
temporal_size=num_frames,
|
| 222 |
+
use_real=True,
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
freqs_cos = freqs_cos.to(device=device)
|
| 226 |
+
freqs_sin = freqs_sin.to(device=device)
|
| 227 |
+
return freqs_cos, freqs_sin
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
def unpatchify(self, hidden_states, height, width):
|
| 231 |
+
hidden_states = rearrange(hidden_states, "B (T H W) (C P Q) -> B C T (H P) (W Q)", P=2, Q=2, H=height//2, W=width//2)
|
| 232 |
+
return hidden_states
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
def build_mask(self, T, H, W, dtype, device, is_bound):
|
| 236 |
+
t = repeat(torch.arange(T), "T -> T H W", T=T, H=H, W=W)
|
| 237 |
+
h = repeat(torch.arange(H), "H -> T H W", T=T, H=H, W=W)
|
| 238 |
+
w = repeat(torch.arange(W), "W -> T H W", T=T, H=H, W=W)
|
| 239 |
+
border_width = (H + W) // 4
|
| 240 |
+
pad = torch.ones_like(h) * border_width
|
| 241 |
+
mask = torch.stack([
|
| 242 |
+
pad if is_bound[0] else t + 1,
|
| 243 |
+
pad if is_bound[1] else T - t,
|
| 244 |
+
pad if is_bound[2] else h + 1,
|
| 245 |
+
pad if is_bound[3] else H - h,
|
| 246 |
+
pad if is_bound[4] else w + 1,
|
| 247 |
+
pad if is_bound[5] else W - w
|
| 248 |
+
]).min(dim=0).values
|
| 249 |
+
mask = mask.clip(1, border_width)
|
| 250 |
+
mask = (mask / border_width).to(dtype=dtype, device=device)
|
| 251 |
+
mask = rearrange(mask, "T H W -> 1 1 T H W")
|
| 252 |
+
return mask
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
def tiled_forward(self, hidden_states, timestep, prompt_emb, tile_size=(60, 90), tile_stride=(30, 45)):
|
| 256 |
+
B, C, T, H, W = hidden_states.shape
|
| 257 |
+
value = torch.zeros((B, C, T, H, W), dtype=hidden_states.dtype, device=hidden_states.device)
|
| 258 |
+
weight = torch.zeros((B, C, T, H, W), dtype=hidden_states.dtype, device=hidden_states.device)
|
| 259 |
+
|
| 260 |
+
# Split tasks
|
| 261 |
+
tasks = []
|
| 262 |
+
for h in range(0, H, tile_stride):
|
| 263 |
+
for w in range(0, W, tile_stride):
|
| 264 |
+
if (h-tile_stride >= 0 and h-tile_stride+tile_size >= H) or (w-tile_stride >= 0 and w-tile_stride+tile_size >= W):
|
| 265 |
+
continue
|
| 266 |
+
h_, w_ = h + tile_size, w + tile_size
|
| 267 |
+
if h_ > H: h, h_ = max(H - tile_size, 0), H
|
| 268 |
+
if w_ > W: w, w_ = max(W - tile_size, 0), W
|
| 269 |
+
tasks.append((h, h_, w, w_))
|
| 270 |
+
|
| 271 |
+
# Run
|
| 272 |
+
for hl, hr, wl, wr in tasks:
|
| 273 |
+
mask = self.build_mask(
|
| 274 |
+
value.shape[2], (hr-hl), (wr-wl),
|
| 275 |
+
hidden_states.dtype, hidden_states.device,
|
| 276 |
+
is_bound=(True, True, hl==0, hr>=H, wl==0, wr>=W)
|
| 277 |
+
)
|
| 278 |
+
model_output = self.forward(hidden_states[:, :, :, hl:hr, wl:wr], timestep, prompt_emb)
|
| 279 |
+
value[:, :, :, hl:hr, wl:wr] += model_output * mask
|
| 280 |
+
weight[:, :, :, hl:hr, wl:wr] += mask
|
| 281 |
+
value = value / weight
|
| 282 |
+
|
| 283 |
+
return value
|
| 284 |
+
|
| 285 |
+
|
| 286 |
+
def forward(self, hidden_states, timestep, prompt_emb, image_rotary_emb=None, tiled=False, tile_size=90, tile_stride=30, use_gradient_checkpointing=False):
|
| 287 |
+
if tiled:
|
| 288 |
+
return TileWorker2Dto3D().tiled_forward(
|
| 289 |
+
forward_fn=lambda x: self.forward(x, timestep, prompt_emb),
|
| 290 |
+
model_input=hidden_states,
|
| 291 |
+
tile_size=tile_size, tile_stride=tile_stride,
|
| 292 |
+
tile_device=hidden_states.device, tile_dtype=hidden_states.dtype,
|
| 293 |
+
computation_device=self.context_embedder.weight.device, computation_dtype=self.context_embedder.weight.dtype
|
| 294 |
+
)
|
| 295 |
+
num_frames, height, width = hidden_states.shape[-3:]
|
| 296 |
+
if image_rotary_emb is None:
|
| 297 |
+
image_rotary_emb = self.prepare_rotary_positional_embeddings(height, width, num_frames, device=self.context_embedder.weight.device)
|
| 298 |
+
hidden_states = self.patchify(hidden_states)
|
| 299 |
+
time_emb = self.time_embedder(timestep, dtype=hidden_states.dtype)
|
| 300 |
+
prompt_emb = self.context_embedder(prompt_emb)
|
| 301 |
+
|
| 302 |
+
def create_custom_forward(module):
|
| 303 |
+
def custom_forward(*inputs):
|
| 304 |
+
return module(*inputs)
|
| 305 |
+
return custom_forward
|
| 306 |
+
|
| 307 |
+
for block in self.blocks:
|
| 308 |
+
if self.training and use_gradient_checkpointing:
|
| 309 |
+
hidden_states, prompt_emb = torch.utils.checkpoint.checkpoint(
|
| 310 |
+
create_custom_forward(block),
|
| 311 |
+
hidden_states, prompt_emb, time_emb, image_rotary_emb,
|
| 312 |
+
use_reentrant=False,
|
| 313 |
+
)
|
| 314 |
+
else:
|
| 315 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, time_emb, image_rotary_emb)
|
| 316 |
+
|
| 317 |
+
hidden_states = torch.cat([prompt_emb, hidden_states], dim=1)
|
| 318 |
+
hidden_states = self.norm_final(hidden_states)
|
| 319 |
+
hidden_states = hidden_states[:, prompt_emb.shape[1]:]
|
| 320 |
+
hidden_states = self.norm_out(hidden_states, prompt_emb, time_emb)
|
| 321 |
+
hidden_states = self.proj_out(hidden_states)
|
| 322 |
+
hidden_states = self.unpatchify(hidden_states, height, width)
|
| 323 |
+
|
| 324 |
+
return hidden_states
|
| 325 |
+
|
| 326 |
+
|
| 327 |
+
@staticmethod
|
| 328 |
+
def state_dict_converter():
|
| 329 |
+
return CogDiTStateDictConverter()
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
@staticmethod
|
| 333 |
+
def from_pretrained(file_path, torch_dtype=torch.bfloat16):
|
| 334 |
+
model = CogDiT().to(torch_dtype)
|
| 335 |
+
state_dict = load_state_dict_from_folder(file_path, torch_dtype=torch_dtype)
|
| 336 |
+
state_dict = CogDiT.state_dict_converter().from_diffusers(state_dict)
|
| 337 |
+
model.load_state_dict(state_dict)
|
| 338 |
+
return model
|
| 339 |
+
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
class CogDiTStateDictConverter:
|
| 343 |
+
def __init__(self):
|
| 344 |
+
pass
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
def from_diffusers(self, state_dict):
|
| 348 |
+
rename_dict = {
|
| 349 |
+
"patch_embed.proj.weight": "patchify.proj.weight",
|
| 350 |
+
"patch_embed.proj.bias": "patchify.proj.bias",
|
| 351 |
+
"patch_embed.text_proj.weight": "context_embedder.weight",
|
| 352 |
+
"patch_embed.text_proj.bias": "context_embedder.bias",
|
| 353 |
+
"time_embedding.linear_1.weight": "time_embedder.timestep_embedder.0.weight",
|
| 354 |
+
"time_embedding.linear_1.bias": "time_embedder.timestep_embedder.0.bias",
|
| 355 |
+
"time_embedding.linear_2.weight": "time_embedder.timestep_embedder.2.weight",
|
| 356 |
+
"time_embedding.linear_2.bias": "time_embedder.timestep_embedder.2.bias",
|
| 357 |
+
|
| 358 |
+
"norm_final.weight": "norm_final.weight",
|
| 359 |
+
"norm_final.bias": "norm_final.bias",
|
| 360 |
+
"norm_out.linear.weight": "norm_out.linear.weight",
|
| 361 |
+
"norm_out.linear.bias": "norm_out.linear.bias",
|
| 362 |
+
"norm_out.norm.weight": "norm_out.norm.weight",
|
| 363 |
+
"norm_out.norm.bias": "norm_out.norm.bias",
|
| 364 |
+
"proj_out.weight": "proj_out.weight",
|
| 365 |
+
"proj_out.bias": "proj_out.bias",
|
| 366 |
+
}
|
| 367 |
+
suffix_dict = {
|
| 368 |
+
"norm1.linear.weight": "norm1.linear.weight",
|
| 369 |
+
"norm1.linear.bias": "norm1.linear.bias",
|
| 370 |
+
"norm1.norm.weight": "norm1.norm.weight",
|
| 371 |
+
"norm1.norm.bias": "norm1.norm.bias",
|
| 372 |
+
"attn1.norm_q.weight": "norm_q.weight",
|
| 373 |
+
"attn1.norm_q.bias": "norm_q.bias",
|
| 374 |
+
"attn1.norm_k.weight": "norm_k.weight",
|
| 375 |
+
"attn1.norm_k.bias": "norm_k.bias",
|
| 376 |
+
"attn1.to_q.weight": "attn1.to_q.weight",
|
| 377 |
+
"attn1.to_q.bias": "attn1.to_q.bias",
|
| 378 |
+
"attn1.to_k.weight": "attn1.to_k.weight",
|
| 379 |
+
"attn1.to_k.bias": "attn1.to_k.bias",
|
| 380 |
+
"attn1.to_v.weight": "attn1.to_v.weight",
|
| 381 |
+
"attn1.to_v.bias": "attn1.to_v.bias",
|
| 382 |
+
"attn1.to_out.0.weight": "attn1.to_out.weight",
|
| 383 |
+
"attn1.to_out.0.bias": "attn1.to_out.bias",
|
| 384 |
+
"norm2.linear.weight": "norm2.linear.weight",
|
| 385 |
+
"norm2.linear.bias": "norm2.linear.bias",
|
| 386 |
+
"norm2.norm.weight": "norm2.norm.weight",
|
| 387 |
+
"norm2.norm.bias": "norm2.norm.bias",
|
| 388 |
+
"ff.net.0.proj.weight": "ff.0.weight",
|
| 389 |
+
"ff.net.0.proj.bias": "ff.0.bias",
|
| 390 |
+
"ff.net.2.weight": "ff.2.weight",
|
| 391 |
+
"ff.net.2.bias": "ff.2.bias",
|
| 392 |
+
}
|
| 393 |
+
state_dict_ = {}
|
| 394 |
+
for name, param in state_dict.items():
|
| 395 |
+
if name in rename_dict:
|
| 396 |
+
if name == "patch_embed.proj.weight":
|
| 397 |
+
param = param.unsqueeze(2)
|
| 398 |
+
state_dict_[rename_dict[name]] = param
|
| 399 |
+
else:
|
| 400 |
+
names = name.split(".")
|
| 401 |
+
if names[0] == "transformer_blocks":
|
| 402 |
+
suffix = ".".join(names[2:])
|
| 403 |
+
state_dict_[f"blocks.{names[1]}." + suffix_dict[suffix]] = param
|
| 404 |
+
return state_dict_
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
def from_civitai(self, state_dict):
|
| 408 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/cog_vae.py
ADDED
|
@@ -0,0 +1,518 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from einops import rearrange, repeat
|
| 3 |
+
from .tiler import TileWorker2Dto3D
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class Downsample3D(torch.nn.Module):
|
| 8 |
+
def __init__(
|
| 9 |
+
self,
|
| 10 |
+
in_channels: int,
|
| 11 |
+
out_channels: int,
|
| 12 |
+
kernel_size: int = 3,
|
| 13 |
+
stride: int = 2,
|
| 14 |
+
padding: int = 0,
|
| 15 |
+
compress_time: bool = False,
|
| 16 |
+
):
|
| 17 |
+
super().__init__()
|
| 18 |
+
|
| 19 |
+
self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)
|
| 20 |
+
self.compress_time = compress_time
|
| 21 |
+
|
| 22 |
+
def forward(self, x: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:
|
| 23 |
+
if self.compress_time:
|
| 24 |
+
batch_size, channels, frames, height, width = x.shape
|
| 25 |
+
|
| 26 |
+
# (batch_size, channels, frames, height, width) -> (batch_size, height, width, channels, frames) -> (batch_size * height * width, channels, frames)
|
| 27 |
+
x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames)
|
| 28 |
+
|
| 29 |
+
if x.shape[-1] % 2 == 1:
|
| 30 |
+
x_first, x_rest = x[..., 0], x[..., 1:]
|
| 31 |
+
if x_rest.shape[-1] > 0:
|
| 32 |
+
# (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)
|
| 33 |
+
x_rest = torch.nn.functional.avg_pool1d(x_rest, kernel_size=2, stride=2)
|
| 34 |
+
|
| 35 |
+
x = torch.cat([x_first[..., None], x_rest], dim=-1)
|
| 36 |
+
# (batch_size * height * width, channels, (frames // 2) + 1) -> (batch_size, height, width, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, height, width)
|
| 37 |
+
x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)
|
| 38 |
+
else:
|
| 39 |
+
# (batch_size * height * width, channels, frames) -> (batch_size * height * width, channels, frames // 2)
|
| 40 |
+
x = torch.nn.functional.avg_pool1d(x, kernel_size=2, stride=2)
|
| 41 |
+
# (batch_size * height * width, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)
|
| 42 |
+
x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)
|
| 43 |
+
|
| 44 |
+
# Pad the tensor
|
| 45 |
+
pad = (0, 1, 0, 1)
|
| 46 |
+
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
|
| 47 |
+
batch_size, channels, frames, height, width = x.shape
|
| 48 |
+
# (batch_size, channels, frames, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size * frames, channels, height, width)
|
| 49 |
+
x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, height, width)
|
| 50 |
+
x = self.conv(x)
|
| 51 |
+
# (batch_size * frames, channels, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size, channels, frames, height, width)
|
| 52 |
+
x = x.reshape(batch_size, frames, x.shape[1], x.shape[2], x.shape[3]).permute(0, 2, 1, 3, 4)
|
| 53 |
+
return x
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
class Upsample3D(torch.nn.Module):
|
| 58 |
+
def __init__(
|
| 59 |
+
self,
|
| 60 |
+
in_channels: int,
|
| 61 |
+
out_channels: int,
|
| 62 |
+
kernel_size: int = 3,
|
| 63 |
+
stride: int = 1,
|
| 64 |
+
padding: int = 1,
|
| 65 |
+
compress_time: bool = False,
|
| 66 |
+
) -> None:
|
| 67 |
+
super().__init__()
|
| 68 |
+
self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)
|
| 69 |
+
self.compress_time = compress_time
|
| 70 |
+
|
| 71 |
+
def forward(self, inputs: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:
|
| 72 |
+
if self.compress_time:
|
| 73 |
+
if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:
|
| 74 |
+
# split first frame
|
| 75 |
+
x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]
|
| 76 |
+
|
| 77 |
+
x_first = torch.nn.functional.interpolate(x_first, scale_factor=2.0)
|
| 78 |
+
x_rest = torch.nn.functional.interpolate(x_rest, scale_factor=2.0)
|
| 79 |
+
x_first = x_first[:, :, None, :, :]
|
| 80 |
+
inputs = torch.cat([x_first, x_rest], dim=2)
|
| 81 |
+
elif inputs.shape[2] > 1:
|
| 82 |
+
inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)
|
| 83 |
+
else:
|
| 84 |
+
inputs = inputs.squeeze(2)
|
| 85 |
+
inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)
|
| 86 |
+
inputs = inputs[:, :, None, :, :]
|
| 87 |
+
else:
|
| 88 |
+
# only interpolate 2D
|
| 89 |
+
b, c, t, h, w = inputs.shape
|
| 90 |
+
inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
| 91 |
+
inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)
|
| 92 |
+
inputs = inputs.reshape(b, t, c, *inputs.shape[2:]).permute(0, 2, 1, 3, 4)
|
| 93 |
+
|
| 94 |
+
b, c, t, h, w = inputs.shape
|
| 95 |
+
inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
| 96 |
+
inputs = self.conv(inputs)
|
| 97 |
+
inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3, 4)
|
| 98 |
+
|
| 99 |
+
return inputs
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class CogVideoXSpatialNorm3D(torch.nn.Module):
|
| 104 |
+
def __init__(self, f_channels, zq_channels, groups):
|
| 105 |
+
super().__init__()
|
| 106 |
+
self.norm_layer = torch.nn.GroupNorm(num_channels=f_channels, num_groups=groups, eps=1e-6, affine=True)
|
| 107 |
+
self.conv_y = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)
|
| 108 |
+
self.conv_b = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor:
|
| 112 |
+
if f.shape[2] > 1 and f.shape[2] % 2 == 1:
|
| 113 |
+
f_first, f_rest = f[:, :, :1], f[:, :, 1:]
|
| 114 |
+
f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:]
|
| 115 |
+
z_first, z_rest = zq[:, :, :1], zq[:, :, 1:]
|
| 116 |
+
z_first = torch.nn.functional.interpolate(z_first, size=f_first_size)
|
| 117 |
+
z_rest = torch.nn.functional.interpolate(z_rest, size=f_rest_size)
|
| 118 |
+
zq = torch.cat([z_first, z_rest], dim=2)
|
| 119 |
+
else:
|
| 120 |
+
zq = torch.nn.functional.interpolate(zq, size=f.shape[-3:])
|
| 121 |
+
|
| 122 |
+
norm_f = self.norm_layer(f)
|
| 123 |
+
new_f = norm_f * self.conv_y(zq) + self.conv_b(zq)
|
| 124 |
+
return new_f
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
class Resnet3DBlock(torch.nn.Module):
|
| 129 |
+
def __init__(self, in_channels, out_channels, spatial_norm_dim, groups, eps=1e-6, use_conv_shortcut=False):
|
| 130 |
+
super().__init__()
|
| 131 |
+
self.nonlinearity = torch.nn.SiLU()
|
| 132 |
+
if spatial_norm_dim is None:
|
| 133 |
+
self.norm1 = torch.nn.GroupNorm(num_channels=in_channels, num_groups=groups, eps=eps)
|
| 134 |
+
self.norm2 = torch.nn.GroupNorm(num_channels=out_channels, num_groups=groups, eps=eps)
|
| 135 |
+
else:
|
| 136 |
+
self.norm1 = CogVideoXSpatialNorm3D(in_channels, spatial_norm_dim, groups)
|
| 137 |
+
self.norm2 = CogVideoXSpatialNorm3D(out_channels, spatial_norm_dim, groups)
|
| 138 |
+
|
| 139 |
+
self.conv1 = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))
|
| 140 |
+
|
| 141 |
+
self.conv2 = CachedConv3d(out_channels, out_channels, kernel_size=3, padding=(0, 1, 1))
|
| 142 |
+
|
| 143 |
+
if in_channels != out_channels:
|
| 144 |
+
if use_conv_shortcut:
|
| 145 |
+
self.conv_shortcut = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))
|
| 146 |
+
else:
|
| 147 |
+
self.conv_shortcut = torch.nn.Conv3d(in_channels, out_channels, kernel_size=1)
|
| 148 |
+
else:
|
| 149 |
+
self.conv_shortcut = lambda x: x
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
def forward(self, hidden_states, zq):
|
| 153 |
+
residual = hidden_states
|
| 154 |
+
|
| 155 |
+
hidden_states = self.norm1(hidden_states, zq) if isinstance(self.norm1, CogVideoXSpatialNorm3D) else self.norm1(hidden_states)
|
| 156 |
+
hidden_states = self.nonlinearity(hidden_states)
|
| 157 |
+
hidden_states = self.conv1(hidden_states)
|
| 158 |
+
|
| 159 |
+
hidden_states = self.norm2(hidden_states, zq) if isinstance(self.norm2, CogVideoXSpatialNorm3D) else self.norm2(hidden_states)
|
| 160 |
+
hidden_states = self.nonlinearity(hidden_states)
|
| 161 |
+
hidden_states = self.conv2(hidden_states)
|
| 162 |
+
|
| 163 |
+
hidden_states = hidden_states + self.conv_shortcut(residual)
|
| 164 |
+
|
| 165 |
+
return hidden_states
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
class CachedConv3d(torch.nn.Conv3d):
|
| 170 |
+
def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0):
|
| 171 |
+
super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)
|
| 172 |
+
self.cached_tensor = None
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def clear_cache(self):
|
| 176 |
+
self.cached_tensor = None
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def forward(self, input: torch.Tensor, use_cache = True) -> torch.Tensor:
|
| 180 |
+
if use_cache:
|
| 181 |
+
if self.cached_tensor is None:
|
| 182 |
+
self.cached_tensor = torch.concat([input[:, :, :1]] * 2, dim=2)
|
| 183 |
+
input = torch.concat([self.cached_tensor, input], dim=2)
|
| 184 |
+
self.cached_tensor = input[:, :, -2:]
|
| 185 |
+
return super().forward(input)
|
| 186 |
+
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
class CogVAEDecoder(torch.nn.Module):
|
| 190 |
+
def __init__(self):
|
| 191 |
+
super().__init__()
|
| 192 |
+
self.scaling_factor = 0.7
|
| 193 |
+
self.conv_in = CachedConv3d(16, 512, kernel_size=3, stride=1, padding=(0, 1, 1))
|
| 194 |
+
|
| 195 |
+
self.blocks = torch.nn.ModuleList([
|
| 196 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 197 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 198 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 199 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 200 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 201 |
+
Resnet3DBlock(512, 512, 16, 32),
|
| 202 |
+
Upsample3D(512, 512, compress_time=True),
|
| 203 |
+
Resnet3DBlock(512, 256, 16, 32),
|
| 204 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 205 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 206 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 207 |
+
Upsample3D(256, 256, compress_time=True),
|
| 208 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 209 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 210 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 211 |
+
Resnet3DBlock(256, 256, 16, 32),
|
| 212 |
+
Upsample3D(256, 256, compress_time=False),
|
| 213 |
+
Resnet3DBlock(256, 128, 16, 32),
|
| 214 |
+
Resnet3DBlock(128, 128, 16, 32),
|
| 215 |
+
Resnet3DBlock(128, 128, 16, 32),
|
| 216 |
+
Resnet3DBlock(128, 128, 16, 32),
|
| 217 |
+
])
|
| 218 |
+
|
| 219 |
+
self.norm_out = CogVideoXSpatialNorm3D(128, 16, 32)
|
| 220 |
+
self.conv_act = torch.nn.SiLU()
|
| 221 |
+
self.conv_out = CachedConv3d(128, 3, kernel_size=3, stride=1, padding=(0, 1, 1))
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
def forward(self, sample):
|
| 225 |
+
sample = sample / self.scaling_factor
|
| 226 |
+
hidden_states = self.conv_in(sample)
|
| 227 |
+
|
| 228 |
+
for block in self.blocks:
|
| 229 |
+
hidden_states = block(hidden_states, sample)
|
| 230 |
+
|
| 231 |
+
hidden_states = self.norm_out(hidden_states, sample)
|
| 232 |
+
hidden_states = self.conv_act(hidden_states)
|
| 233 |
+
hidden_states = self.conv_out(hidden_states)
|
| 234 |
+
|
| 235 |
+
return hidden_states
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
def decode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):
|
| 239 |
+
if tiled:
|
| 240 |
+
B, C, T, H, W = sample.shape
|
| 241 |
+
return TileWorker2Dto3D().tiled_forward(
|
| 242 |
+
forward_fn=lambda x: self.decode_small_video(x),
|
| 243 |
+
model_input=sample,
|
| 244 |
+
tile_size=tile_size, tile_stride=tile_stride,
|
| 245 |
+
tile_device=sample.device, tile_dtype=sample.dtype,
|
| 246 |
+
computation_device=sample.device, computation_dtype=sample.dtype,
|
| 247 |
+
scales=(3/16, (T//2*8+T%2)/T, 8, 8),
|
| 248 |
+
progress_bar=progress_bar
|
| 249 |
+
)
|
| 250 |
+
else:
|
| 251 |
+
return self.decode_small_video(sample)
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
def decode_small_video(self, sample):
|
| 255 |
+
B, C, T, H, W = sample.shape
|
| 256 |
+
computation_device = self.conv_in.weight.device
|
| 257 |
+
computation_dtype = self.conv_in.weight.dtype
|
| 258 |
+
value = []
|
| 259 |
+
for i in range(T//2):
|
| 260 |
+
tl = i*2 + T%2 - (T%2 and i==0)
|
| 261 |
+
tr = i*2 + 2 + T%2
|
| 262 |
+
model_input = sample[:, :, tl: tr, :, :].to(dtype=computation_dtype, device=computation_device)
|
| 263 |
+
model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)
|
| 264 |
+
value.append(model_output)
|
| 265 |
+
value = torch.concat(value, dim=2)
|
| 266 |
+
for name, module in self.named_modules():
|
| 267 |
+
if isinstance(module, CachedConv3d):
|
| 268 |
+
module.clear_cache()
|
| 269 |
+
return value
|
| 270 |
+
|
| 271 |
+
|
| 272 |
+
@staticmethod
|
| 273 |
+
def state_dict_converter():
|
| 274 |
+
return CogVAEDecoderStateDictConverter()
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
|
| 278 |
+
class CogVAEEncoder(torch.nn.Module):
|
| 279 |
+
def __init__(self):
|
| 280 |
+
super().__init__()
|
| 281 |
+
self.scaling_factor = 0.7
|
| 282 |
+
self.conv_in = CachedConv3d(3, 128, kernel_size=3, stride=1, padding=(0, 1, 1))
|
| 283 |
+
|
| 284 |
+
self.blocks = torch.nn.ModuleList([
|
| 285 |
+
Resnet3DBlock(128, 128, None, 32),
|
| 286 |
+
Resnet3DBlock(128, 128, None, 32),
|
| 287 |
+
Resnet3DBlock(128, 128, None, 32),
|
| 288 |
+
Downsample3D(128, 128, compress_time=True),
|
| 289 |
+
Resnet3DBlock(128, 256, None, 32),
|
| 290 |
+
Resnet3DBlock(256, 256, None, 32),
|
| 291 |
+
Resnet3DBlock(256, 256, None, 32),
|
| 292 |
+
Downsample3D(256, 256, compress_time=True),
|
| 293 |
+
Resnet3DBlock(256, 256, None, 32),
|
| 294 |
+
Resnet3DBlock(256, 256, None, 32),
|
| 295 |
+
Resnet3DBlock(256, 256, None, 32),
|
| 296 |
+
Downsample3D(256, 256, compress_time=False),
|
| 297 |
+
Resnet3DBlock(256, 512, None, 32),
|
| 298 |
+
Resnet3DBlock(512, 512, None, 32),
|
| 299 |
+
Resnet3DBlock(512, 512, None, 32),
|
| 300 |
+
Resnet3DBlock(512, 512, None, 32),
|
| 301 |
+
Resnet3DBlock(512, 512, None, 32),
|
| 302 |
+
])
|
| 303 |
+
|
| 304 |
+
self.norm_out = torch.nn.GroupNorm(32, 512, eps=1e-06, affine=True)
|
| 305 |
+
self.conv_act = torch.nn.SiLU()
|
| 306 |
+
self.conv_out = CachedConv3d(512, 32, kernel_size=3, stride=1, padding=(0, 1, 1))
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
def forward(self, sample):
|
| 310 |
+
hidden_states = self.conv_in(sample)
|
| 311 |
+
|
| 312 |
+
for block in self.blocks:
|
| 313 |
+
hidden_states = block(hidden_states, sample)
|
| 314 |
+
|
| 315 |
+
hidden_states = self.norm_out(hidden_states)
|
| 316 |
+
hidden_states = self.conv_act(hidden_states)
|
| 317 |
+
hidden_states = self.conv_out(hidden_states)[:, :16]
|
| 318 |
+
hidden_states = hidden_states * self.scaling_factor
|
| 319 |
+
|
| 320 |
+
return hidden_states
|
| 321 |
+
|
| 322 |
+
|
| 323 |
+
def encode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):
|
| 324 |
+
if tiled:
|
| 325 |
+
B, C, T, H, W = sample.shape
|
| 326 |
+
return TileWorker2Dto3D().tiled_forward(
|
| 327 |
+
forward_fn=lambda x: self.encode_small_video(x),
|
| 328 |
+
model_input=sample,
|
| 329 |
+
tile_size=(i * 8 for i in tile_size), tile_stride=(i * 8 for i in tile_stride),
|
| 330 |
+
tile_device=sample.device, tile_dtype=sample.dtype,
|
| 331 |
+
computation_device=sample.device, computation_dtype=sample.dtype,
|
| 332 |
+
scales=(16/3, (T//4+T%2)/T, 1/8, 1/8),
|
| 333 |
+
progress_bar=progress_bar
|
| 334 |
+
)
|
| 335 |
+
else:
|
| 336 |
+
return self.encode_small_video(sample)
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
def encode_small_video(self, sample):
|
| 340 |
+
B, C, T, H, W = sample.shape
|
| 341 |
+
computation_device = self.conv_in.weight.device
|
| 342 |
+
computation_dtype = self.conv_in.weight.dtype
|
| 343 |
+
value = []
|
| 344 |
+
for i in range(T//8):
|
| 345 |
+
t = i*8 + T%2 - (T%2 and i==0)
|
| 346 |
+
t_ = i*8 + 8 + T%2
|
| 347 |
+
model_input = sample[:, :, t: t_, :, :].to(dtype=computation_dtype, device=computation_device)
|
| 348 |
+
model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)
|
| 349 |
+
value.append(model_output)
|
| 350 |
+
value = torch.concat(value, dim=2)
|
| 351 |
+
for name, module in self.named_modules():
|
| 352 |
+
if isinstance(module, CachedConv3d):
|
| 353 |
+
module.clear_cache()
|
| 354 |
+
return value
|
| 355 |
+
|
| 356 |
+
|
| 357 |
+
@staticmethod
|
| 358 |
+
def state_dict_converter():
|
| 359 |
+
return CogVAEEncoderStateDictConverter()
|
| 360 |
+
|
| 361 |
+
|
| 362 |
+
|
| 363 |
+
class CogVAEEncoderStateDictConverter:
|
| 364 |
+
def __init__(self):
|
| 365 |
+
pass
|
| 366 |
+
|
| 367 |
+
|
| 368 |
+
def from_diffusers(self, state_dict):
|
| 369 |
+
rename_dict = {
|
| 370 |
+
"encoder.conv_in.conv.weight": "conv_in.weight",
|
| 371 |
+
"encoder.conv_in.conv.bias": "conv_in.bias",
|
| 372 |
+
"encoder.down_blocks.0.downsamplers.0.conv.weight": "blocks.3.conv.weight",
|
| 373 |
+
"encoder.down_blocks.0.downsamplers.0.conv.bias": "blocks.3.conv.bias",
|
| 374 |
+
"encoder.down_blocks.1.downsamplers.0.conv.weight": "blocks.7.conv.weight",
|
| 375 |
+
"encoder.down_blocks.1.downsamplers.0.conv.bias": "blocks.7.conv.bias",
|
| 376 |
+
"encoder.down_blocks.2.downsamplers.0.conv.weight": "blocks.11.conv.weight",
|
| 377 |
+
"encoder.down_blocks.2.downsamplers.0.conv.bias": "blocks.11.conv.bias",
|
| 378 |
+
"encoder.norm_out.weight": "norm_out.weight",
|
| 379 |
+
"encoder.norm_out.bias": "norm_out.bias",
|
| 380 |
+
"encoder.conv_out.conv.weight": "conv_out.weight",
|
| 381 |
+
"encoder.conv_out.conv.bias": "conv_out.bias",
|
| 382 |
+
}
|
| 383 |
+
prefix_dict = {
|
| 384 |
+
"encoder.down_blocks.0.resnets.0.": "blocks.0.",
|
| 385 |
+
"encoder.down_blocks.0.resnets.1.": "blocks.1.",
|
| 386 |
+
"encoder.down_blocks.0.resnets.2.": "blocks.2.",
|
| 387 |
+
"encoder.down_blocks.1.resnets.0.": "blocks.4.",
|
| 388 |
+
"encoder.down_blocks.1.resnets.1.": "blocks.5.",
|
| 389 |
+
"encoder.down_blocks.1.resnets.2.": "blocks.6.",
|
| 390 |
+
"encoder.down_blocks.2.resnets.0.": "blocks.8.",
|
| 391 |
+
"encoder.down_blocks.2.resnets.1.": "blocks.9.",
|
| 392 |
+
"encoder.down_blocks.2.resnets.2.": "blocks.10.",
|
| 393 |
+
"encoder.down_blocks.3.resnets.0.": "blocks.12.",
|
| 394 |
+
"encoder.down_blocks.3.resnets.1.": "blocks.13.",
|
| 395 |
+
"encoder.down_blocks.3.resnets.2.": "blocks.14.",
|
| 396 |
+
"encoder.mid_block.resnets.0.": "blocks.15.",
|
| 397 |
+
"encoder.mid_block.resnets.1.": "blocks.16.",
|
| 398 |
+
}
|
| 399 |
+
suffix_dict = {
|
| 400 |
+
"norm1.norm_layer.weight": "norm1.norm_layer.weight",
|
| 401 |
+
"norm1.norm_layer.bias": "norm1.norm_layer.bias",
|
| 402 |
+
"norm1.conv_y.conv.weight": "norm1.conv_y.weight",
|
| 403 |
+
"norm1.conv_y.conv.bias": "norm1.conv_y.bias",
|
| 404 |
+
"norm1.conv_b.conv.weight": "norm1.conv_b.weight",
|
| 405 |
+
"norm1.conv_b.conv.bias": "norm1.conv_b.bias",
|
| 406 |
+
"norm2.norm_layer.weight": "norm2.norm_layer.weight",
|
| 407 |
+
"norm2.norm_layer.bias": "norm2.norm_layer.bias",
|
| 408 |
+
"norm2.conv_y.conv.weight": "norm2.conv_y.weight",
|
| 409 |
+
"norm2.conv_y.conv.bias": "norm2.conv_y.bias",
|
| 410 |
+
"norm2.conv_b.conv.weight": "norm2.conv_b.weight",
|
| 411 |
+
"norm2.conv_b.conv.bias": "norm2.conv_b.bias",
|
| 412 |
+
"conv1.conv.weight": "conv1.weight",
|
| 413 |
+
"conv1.conv.bias": "conv1.bias",
|
| 414 |
+
"conv2.conv.weight": "conv2.weight",
|
| 415 |
+
"conv2.conv.bias": "conv2.bias",
|
| 416 |
+
"conv_shortcut.weight": "conv_shortcut.weight",
|
| 417 |
+
"conv_shortcut.bias": "conv_shortcut.bias",
|
| 418 |
+
"norm1.weight": "norm1.weight",
|
| 419 |
+
"norm1.bias": "norm1.bias",
|
| 420 |
+
"norm2.weight": "norm2.weight",
|
| 421 |
+
"norm2.bias": "norm2.bias",
|
| 422 |
+
}
|
| 423 |
+
state_dict_ = {}
|
| 424 |
+
for name, param in state_dict.items():
|
| 425 |
+
if name in rename_dict:
|
| 426 |
+
state_dict_[rename_dict[name]] = param
|
| 427 |
+
else:
|
| 428 |
+
for prefix in prefix_dict:
|
| 429 |
+
if name.startswith(prefix):
|
| 430 |
+
suffix = name[len(prefix):]
|
| 431 |
+
state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param
|
| 432 |
+
return state_dict_
|
| 433 |
+
|
| 434 |
+
|
| 435 |
+
def from_civitai(self, state_dict):
|
| 436 |
+
return self.from_diffusers(state_dict)
|
| 437 |
+
|
| 438 |
+
|
| 439 |
+
|
| 440 |
+
class CogVAEDecoderStateDictConverter:
|
| 441 |
+
def __init__(self):
|
| 442 |
+
pass
|
| 443 |
+
|
| 444 |
+
|
| 445 |
+
def from_diffusers(self, state_dict):
|
| 446 |
+
rename_dict = {
|
| 447 |
+
"decoder.conv_in.conv.weight": "conv_in.weight",
|
| 448 |
+
"decoder.conv_in.conv.bias": "conv_in.bias",
|
| 449 |
+
"decoder.up_blocks.0.upsamplers.0.conv.weight": "blocks.6.conv.weight",
|
| 450 |
+
"decoder.up_blocks.0.upsamplers.0.conv.bias": "blocks.6.conv.bias",
|
| 451 |
+
"decoder.up_blocks.1.upsamplers.0.conv.weight": "blocks.11.conv.weight",
|
| 452 |
+
"decoder.up_blocks.1.upsamplers.0.conv.bias": "blocks.11.conv.bias",
|
| 453 |
+
"decoder.up_blocks.2.upsamplers.0.conv.weight": "blocks.16.conv.weight",
|
| 454 |
+
"decoder.up_blocks.2.upsamplers.0.conv.bias": "blocks.16.conv.bias",
|
| 455 |
+
"decoder.norm_out.norm_layer.weight": "norm_out.norm_layer.weight",
|
| 456 |
+
"decoder.norm_out.norm_layer.bias": "norm_out.norm_layer.bias",
|
| 457 |
+
"decoder.norm_out.conv_y.conv.weight": "norm_out.conv_y.weight",
|
| 458 |
+
"decoder.norm_out.conv_y.conv.bias": "norm_out.conv_y.bias",
|
| 459 |
+
"decoder.norm_out.conv_b.conv.weight": "norm_out.conv_b.weight",
|
| 460 |
+
"decoder.norm_out.conv_b.conv.bias": "norm_out.conv_b.bias",
|
| 461 |
+
"decoder.conv_out.conv.weight": "conv_out.weight",
|
| 462 |
+
"decoder.conv_out.conv.bias": "conv_out.bias"
|
| 463 |
+
}
|
| 464 |
+
prefix_dict = {
|
| 465 |
+
"decoder.mid_block.resnets.0.": "blocks.0.",
|
| 466 |
+
"decoder.mid_block.resnets.1.": "blocks.1.",
|
| 467 |
+
"decoder.up_blocks.0.resnets.0.": "blocks.2.",
|
| 468 |
+
"decoder.up_blocks.0.resnets.1.": "blocks.3.",
|
| 469 |
+
"decoder.up_blocks.0.resnets.2.": "blocks.4.",
|
| 470 |
+
"decoder.up_blocks.0.resnets.3.": "blocks.5.",
|
| 471 |
+
"decoder.up_blocks.1.resnets.0.": "blocks.7.",
|
| 472 |
+
"decoder.up_blocks.1.resnets.1.": "blocks.8.",
|
| 473 |
+
"decoder.up_blocks.1.resnets.2.": "blocks.9.",
|
| 474 |
+
"decoder.up_blocks.1.resnets.3.": "blocks.10.",
|
| 475 |
+
"decoder.up_blocks.2.resnets.0.": "blocks.12.",
|
| 476 |
+
"decoder.up_blocks.2.resnets.1.": "blocks.13.",
|
| 477 |
+
"decoder.up_blocks.2.resnets.2.": "blocks.14.",
|
| 478 |
+
"decoder.up_blocks.2.resnets.3.": "blocks.15.",
|
| 479 |
+
"decoder.up_blocks.3.resnets.0.": "blocks.17.",
|
| 480 |
+
"decoder.up_blocks.3.resnets.1.": "blocks.18.",
|
| 481 |
+
"decoder.up_blocks.3.resnets.2.": "blocks.19.",
|
| 482 |
+
"decoder.up_blocks.3.resnets.3.": "blocks.20.",
|
| 483 |
+
}
|
| 484 |
+
suffix_dict = {
|
| 485 |
+
"norm1.norm_layer.weight": "norm1.norm_layer.weight",
|
| 486 |
+
"norm1.norm_layer.bias": "norm1.norm_layer.bias",
|
| 487 |
+
"norm1.conv_y.conv.weight": "norm1.conv_y.weight",
|
| 488 |
+
"norm1.conv_y.conv.bias": "norm1.conv_y.bias",
|
| 489 |
+
"norm1.conv_b.conv.weight": "norm1.conv_b.weight",
|
| 490 |
+
"norm1.conv_b.conv.bias": "norm1.conv_b.bias",
|
| 491 |
+
"norm2.norm_layer.weight": "norm2.norm_layer.weight",
|
| 492 |
+
"norm2.norm_layer.bias": "norm2.norm_layer.bias",
|
| 493 |
+
"norm2.conv_y.conv.weight": "norm2.conv_y.weight",
|
| 494 |
+
"norm2.conv_y.conv.bias": "norm2.conv_y.bias",
|
| 495 |
+
"norm2.conv_b.conv.weight": "norm2.conv_b.weight",
|
| 496 |
+
"norm2.conv_b.conv.bias": "norm2.conv_b.bias",
|
| 497 |
+
"conv1.conv.weight": "conv1.weight",
|
| 498 |
+
"conv1.conv.bias": "conv1.bias",
|
| 499 |
+
"conv2.conv.weight": "conv2.weight",
|
| 500 |
+
"conv2.conv.bias": "conv2.bias",
|
| 501 |
+
"conv_shortcut.weight": "conv_shortcut.weight",
|
| 502 |
+
"conv_shortcut.bias": "conv_shortcut.bias",
|
| 503 |
+
}
|
| 504 |
+
state_dict_ = {}
|
| 505 |
+
for name, param in state_dict.items():
|
| 506 |
+
if name in rename_dict:
|
| 507 |
+
state_dict_[rename_dict[name]] = param
|
| 508 |
+
else:
|
| 509 |
+
for prefix in prefix_dict:
|
| 510 |
+
if name.startswith(prefix):
|
| 511 |
+
suffix = name[len(prefix):]
|
| 512 |
+
state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param
|
| 513 |
+
return state_dict_
|
| 514 |
+
|
| 515 |
+
|
| 516 |
+
def from_civitai(self, state_dict):
|
| 517 |
+
return self.from_diffusers(state_dict)
|
| 518 |
+
|
diffsynth/models/downloader.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from huggingface_hub import hf_hub_download
|
| 2 |
+
from modelscope import snapshot_download
|
| 3 |
+
import os, shutil
|
| 4 |
+
from typing_extensions import Literal, TypeAlias
|
| 5 |
+
from typing import List
|
| 6 |
+
from ..configs.model_config import preset_models_on_huggingface, preset_models_on_modelscope, Preset_model_id
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def download_from_modelscope(model_id, origin_file_path, local_dir):
|
| 10 |
+
os.makedirs(local_dir, exist_ok=True)
|
| 11 |
+
file_name = os.path.basename(origin_file_path)
|
| 12 |
+
if file_name in os.listdir(local_dir):
|
| 13 |
+
print(f" {file_name} has been already in {local_dir}.")
|
| 14 |
+
else:
|
| 15 |
+
print(f" Start downloading {os.path.join(local_dir, file_name)}")
|
| 16 |
+
snapshot_download(model_id, allow_file_pattern=origin_file_path, local_dir=local_dir)
|
| 17 |
+
downloaded_file_path = os.path.join(local_dir, origin_file_path)
|
| 18 |
+
target_file_path = os.path.join(local_dir, os.path.split(origin_file_path)[-1])
|
| 19 |
+
if downloaded_file_path != target_file_path:
|
| 20 |
+
shutil.move(downloaded_file_path, target_file_path)
|
| 21 |
+
shutil.rmtree(os.path.join(local_dir, origin_file_path.split("/")[0]))
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def download_from_huggingface(model_id, origin_file_path, local_dir):
|
| 25 |
+
os.makedirs(local_dir, exist_ok=True)
|
| 26 |
+
file_name = os.path.basename(origin_file_path)
|
| 27 |
+
if file_name in os.listdir(local_dir):
|
| 28 |
+
print(f" {file_name} has been already in {local_dir}.")
|
| 29 |
+
else:
|
| 30 |
+
print(f" Start downloading {os.path.join(local_dir, file_name)}")
|
| 31 |
+
hf_hub_download(model_id, origin_file_path, local_dir=local_dir)
|
| 32 |
+
downloaded_file_path = os.path.join(local_dir, origin_file_path)
|
| 33 |
+
target_file_path = os.path.join(local_dir, file_name)
|
| 34 |
+
if downloaded_file_path != target_file_path:
|
| 35 |
+
shutil.move(downloaded_file_path, target_file_path)
|
| 36 |
+
shutil.rmtree(os.path.join(local_dir, origin_file_path.split("/")[0]))
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
Preset_model_website: TypeAlias = Literal[
|
| 40 |
+
"HuggingFace",
|
| 41 |
+
"ModelScope",
|
| 42 |
+
]
|
| 43 |
+
website_to_preset_models = {
|
| 44 |
+
"HuggingFace": preset_models_on_huggingface,
|
| 45 |
+
"ModelScope": preset_models_on_modelscope,
|
| 46 |
+
}
|
| 47 |
+
website_to_download_fn = {
|
| 48 |
+
"HuggingFace": download_from_huggingface,
|
| 49 |
+
"ModelScope": download_from_modelscope,
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def download_customized_models(
|
| 54 |
+
model_id,
|
| 55 |
+
origin_file_path,
|
| 56 |
+
local_dir,
|
| 57 |
+
downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],
|
| 58 |
+
):
|
| 59 |
+
downloaded_files = []
|
| 60 |
+
for website in downloading_priority:
|
| 61 |
+
# Check if the file is downloaded.
|
| 62 |
+
file_to_download = os.path.join(local_dir, os.path.basename(origin_file_path))
|
| 63 |
+
if file_to_download in downloaded_files:
|
| 64 |
+
continue
|
| 65 |
+
# Download
|
| 66 |
+
website_to_download_fn[website](model_id, origin_file_path, local_dir)
|
| 67 |
+
if os.path.basename(origin_file_path) in os.listdir(local_dir):
|
| 68 |
+
downloaded_files.append(file_to_download)
|
| 69 |
+
return downloaded_files
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def download_models(
|
| 73 |
+
model_id_list: List[Preset_model_id] = [],
|
| 74 |
+
downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],
|
| 75 |
+
):
|
| 76 |
+
print(f"Downloading models: {model_id_list}")
|
| 77 |
+
downloaded_files = []
|
| 78 |
+
load_files = []
|
| 79 |
+
|
| 80 |
+
for model_id in model_id_list:
|
| 81 |
+
for website in downloading_priority:
|
| 82 |
+
if model_id in website_to_preset_models[website]:
|
| 83 |
+
|
| 84 |
+
# Parse model metadata
|
| 85 |
+
model_metadata = website_to_preset_models[website][model_id]
|
| 86 |
+
if isinstance(model_metadata, list):
|
| 87 |
+
file_data = model_metadata
|
| 88 |
+
else:
|
| 89 |
+
file_data = model_metadata.get("file_list", [])
|
| 90 |
+
|
| 91 |
+
# Try downloading the model from this website.
|
| 92 |
+
model_files = []
|
| 93 |
+
for model_id, origin_file_path, local_dir in file_data:
|
| 94 |
+
# Check if the file is downloaded.
|
| 95 |
+
file_to_download = os.path.join(local_dir, os.path.basename(origin_file_path))
|
| 96 |
+
if file_to_download in downloaded_files:
|
| 97 |
+
continue
|
| 98 |
+
# Download
|
| 99 |
+
website_to_download_fn[website](model_id, origin_file_path, local_dir)
|
| 100 |
+
if os.path.basename(origin_file_path) in os.listdir(local_dir):
|
| 101 |
+
downloaded_files.append(file_to_download)
|
| 102 |
+
model_files.append(file_to_download)
|
| 103 |
+
|
| 104 |
+
# If the model is successfully downloaded, break.
|
| 105 |
+
if len(model_files) > 0:
|
| 106 |
+
if isinstance(model_metadata, dict) and "load_path" in model_metadata:
|
| 107 |
+
model_files = model_metadata["load_path"]
|
| 108 |
+
load_files.extend(model_files)
|
| 109 |
+
break
|
| 110 |
+
|
| 111 |
+
return load_files
|
diffsynth/models/flux_controlnet.py
ADDED
|
@@ -0,0 +1,331 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from einops import rearrange, repeat
|
| 3 |
+
from .flux_dit import RoPEEmbedding, TimestepEmbeddings, FluxJointTransformerBlock, FluxSingleTransformerBlock, RMSNorm
|
| 4 |
+
from .utils import hash_state_dict_keys, init_weights_on_device
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class FluxControlNet(torch.nn.Module):
|
| 9 |
+
def __init__(self, disable_guidance_embedder=False, num_joint_blocks=5, num_single_blocks=10, num_mode=0, mode_dict={}, additional_input_dim=0):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.pos_embedder = RoPEEmbedding(3072, 10000, [16, 56, 56])
|
| 12 |
+
self.time_embedder = TimestepEmbeddings(256, 3072)
|
| 13 |
+
self.guidance_embedder = None if disable_guidance_embedder else TimestepEmbeddings(256, 3072)
|
| 14 |
+
self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(768, 3072), torch.nn.SiLU(), torch.nn.Linear(3072, 3072))
|
| 15 |
+
self.context_embedder = torch.nn.Linear(4096, 3072)
|
| 16 |
+
self.x_embedder = torch.nn.Linear(64, 3072)
|
| 17 |
+
|
| 18 |
+
self.blocks = torch.nn.ModuleList([FluxJointTransformerBlock(3072, 24) for _ in range(num_joint_blocks)])
|
| 19 |
+
self.single_blocks = torch.nn.ModuleList([FluxSingleTransformerBlock(3072, 24) for _ in range(num_single_blocks)])
|
| 20 |
+
|
| 21 |
+
self.controlnet_blocks = torch.nn.ModuleList([torch.nn.Linear(3072, 3072) for _ in range(num_joint_blocks)])
|
| 22 |
+
self.controlnet_single_blocks = torch.nn.ModuleList([torch.nn.Linear(3072, 3072) for _ in range(num_single_blocks)])
|
| 23 |
+
|
| 24 |
+
self.mode_dict = mode_dict
|
| 25 |
+
self.controlnet_mode_embedder = torch.nn.Embedding(num_mode, 3072) if len(mode_dict) > 0 else None
|
| 26 |
+
self.controlnet_x_embedder = torch.nn.Linear(64 + additional_input_dim, 3072)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def prepare_image_ids(self, latents):
|
| 30 |
+
batch_size, _, height, width = latents.shape
|
| 31 |
+
latent_image_ids = torch.zeros(height // 2, width // 2, 3)
|
| 32 |
+
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height // 2)[:, None]
|
| 33 |
+
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width // 2)[None, :]
|
| 34 |
+
|
| 35 |
+
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
|
| 36 |
+
|
| 37 |
+
latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1)
|
| 38 |
+
latent_image_ids = latent_image_ids.reshape(
|
| 39 |
+
batch_size, latent_image_id_height * latent_image_id_width, latent_image_id_channels
|
| 40 |
+
)
|
| 41 |
+
latent_image_ids = latent_image_ids.to(device=latents.device, dtype=latents.dtype)
|
| 42 |
+
|
| 43 |
+
return latent_image_ids
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def patchify(self, hidden_states):
|
| 47 |
+
hidden_states = rearrange(hidden_states, "B C (H P) (W Q) -> B (H W) (C P Q)", P=2, Q=2)
|
| 48 |
+
return hidden_states
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def align_res_stack_to_original_blocks(self, res_stack, num_blocks, hidden_states):
|
| 52 |
+
if len(res_stack) == 0:
|
| 53 |
+
return [torch.zeros_like(hidden_states)] * num_blocks
|
| 54 |
+
interval = (num_blocks + len(res_stack) - 1) // len(res_stack)
|
| 55 |
+
aligned_res_stack = [res_stack[block_id // interval] for block_id in range(num_blocks)]
|
| 56 |
+
return aligned_res_stack
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def forward(
|
| 60 |
+
self,
|
| 61 |
+
hidden_states,
|
| 62 |
+
controlnet_conditioning,
|
| 63 |
+
timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None,
|
| 64 |
+
processor_id=None,
|
| 65 |
+
tiled=False, tile_size=128, tile_stride=64,
|
| 66 |
+
**kwargs
|
| 67 |
+
):
|
| 68 |
+
if image_ids is None:
|
| 69 |
+
image_ids = self.prepare_image_ids(hidden_states)
|
| 70 |
+
|
| 71 |
+
conditioning = self.time_embedder(timestep, hidden_states.dtype) + self.pooled_text_embedder(pooled_prompt_emb)
|
| 72 |
+
if self.guidance_embedder is not None:
|
| 73 |
+
guidance = guidance * 1000
|
| 74 |
+
conditioning = conditioning + self.guidance_embedder(guidance, hidden_states.dtype)
|
| 75 |
+
prompt_emb = self.context_embedder(prompt_emb)
|
| 76 |
+
if self.controlnet_mode_embedder is not None: # Different from FluxDiT
|
| 77 |
+
processor_id = torch.tensor([self.mode_dict[processor_id]], dtype=torch.int)
|
| 78 |
+
processor_id = repeat(processor_id, "D -> B D", B=1).to(text_ids.device)
|
| 79 |
+
prompt_emb = torch.concat([self.controlnet_mode_embedder(processor_id), prompt_emb], dim=1)
|
| 80 |
+
text_ids = torch.cat([text_ids[:, :1], text_ids], dim=1)
|
| 81 |
+
image_rotary_emb = self.pos_embedder(torch.cat((text_ids, image_ids), dim=1))
|
| 82 |
+
|
| 83 |
+
hidden_states = self.patchify(hidden_states)
|
| 84 |
+
hidden_states = self.x_embedder(hidden_states)
|
| 85 |
+
controlnet_conditioning = self.patchify(controlnet_conditioning) # Different from FluxDiT
|
| 86 |
+
hidden_states = hidden_states + self.controlnet_x_embedder(controlnet_conditioning) # Different from FluxDiT
|
| 87 |
+
|
| 88 |
+
controlnet_res_stack = []
|
| 89 |
+
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
|
| 90 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)
|
| 91 |
+
controlnet_res_stack.append(controlnet_block(hidden_states))
|
| 92 |
+
|
| 93 |
+
controlnet_single_res_stack = []
|
| 94 |
+
hidden_states = torch.cat([prompt_emb, hidden_states], dim=1)
|
| 95 |
+
for block, controlnet_block in zip(self.single_blocks, self.controlnet_single_blocks):
|
| 96 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb)
|
| 97 |
+
controlnet_single_res_stack.append(controlnet_block(hidden_states[:, prompt_emb.shape[1]:]))
|
| 98 |
+
|
| 99 |
+
controlnet_res_stack = self.align_res_stack_to_original_blocks(controlnet_res_stack, 19, hidden_states[:, prompt_emb.shape[1]:])
|
| 100 |
+
controlnet_single_res_stack = self.align_res_stack_to_original_blocks(controlnet_single_res_stack, 38, hidden_states[:, prompt_emb.shape[1]:])
|
| 101 |
+
|
| 102 |
+
return controlnet_res_stack, controlnet_single_res_stack
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
@staticmethod
|
| 106 |
+
def state_dict_converter():
|
| 107 |
+
return FluxControlNetStateDictConverter()
|
| 108 |
+
|
| 109 |
+
def quantize(self):
|
| 110 |
+
def cast_to(weight, dtype=None, device=None, copy=False):
|
| 111 |
+
if device is None or weight.device == device:
|
| 112 |
+
if not copy:
|
| 113 |
+
if dtype is None or weight.dtype == dtype:
|
| 114 |
+
return weight
|
| 115 |
+
return weight.to(dtype=dtype, copy=copy)
|
| 116 |
+
|
| 117 |
+
r = torch.empty_like(weight, dtype=dtype, device=device)
|
| 118 |
+
r.copy_(weight)
|
| 119 |
+
return r
|
| 120 |
+
|
| 121 |
+
def cast_weight(s, input=None, dtype=None, device=None):
|
| 122 |
+
if input is not None:
|
| 123 |
+
if dtype is None:
|
| 124 |
+
dtype = input.dtype
|
| 125 |
+
if device is None:
|
| 126 |
+
device = input.device
|
| 127 |
+
weight = cast_to(s.weight, dtype, device)
|
| 128 |
+
return weight
|
| 129 |
+
|
| 130 |
+
def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
|
| 131 |
+
if input is not None:
|
| 132 |
+
if dtype is None:
|
| 133 |
+
dtype = input.dtype
|
| 134 |
+
if bias_dtype is None:
|
| 135 |
+
bias_dtype = dtype
|
| 136 |
+
if device is None:
|
| 137 |
+
device = input.device
|
| 138 |
+
bias = None
|
| 139 |
+
weight = cast_to(s.weight, dtype, device)
|
| 140 |
+
bias = cast_to(s.bias, bias_dtype, device)
|
| 141 |
+
return weight, bias
|
| 142 |
+
|
| 143 |
+
class quantized_layer:
|
| 144 |
+
class QLinear(torch.nn.Linear):
|
| 145 |
+
def __init__(self, *args, **kwargs):
|
| 146 |
+
super().__init__(*args, **kwargs)
|
| 147 |
+
|
| 148 |
+
def forward(self,input,**kwargs):
|
| 149 |
+
weight,bias= cast_bias_weight(self,input)
|
| 150 |
+
return torch.nn.functional.linear(input,weight,bias)
|
| 151 |
+
|
| 152 |
+
class QRMSNorm(torch.nn.Module):
|
| 153 |
+
def __init__(self, module):
|
| 154 |
+
super().__init__()
|
| 155 |
+
self.module = module
|
| 156 |
+
|
| 157 |
+
def forward(self,hidden_states,**kwargs):
|
| 158 |
+
weight= cast_weight(self.module,hidden_states)
|
| 159 |
+
input_dtype = hidden_states.dtype
|
| 160 |
+
variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)
|
| 161 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.module.eps)
|
| 162 |
+
hidden_states = hidden_states.to(input_dtype) * weight
|
| 163 |
+
return hidden_states
|
| 164 |
+
|
| 165 |
+
class QEmbedding(torch.nn.Embedding):
|
| 166 |
+
def __init__(self, *args, **kwargs):
|
| 167 |
+
super().__init__(*args, **kwargs)
|
| 168 |
+
|
| 169 |
+
def forward(self,input,**kwargs):
|
| 170 |
+
weight= cast_weight(self,input)
|
| 171 |
+
return torch.nn.functional.embedding(
|
| 172 |
+
input, weight, self.padding_idx, self.max_norm,
|
| 173 |
+
self.norm_type, self.scale_grad_by_freq, self.sparse)
|
| 174 |
+
|
| 175 |
+
def replace_layer(model):
|
| 176 |
+
for name, module in model.named_children():
|
| 177 |
+
if isinstance(module,quantized_layer.QRMSNorm):
|
| 178 |
+
continue
|
| 179 |
+
if isinstance(module, torch.nn.Linear):
|
| 180 |
+
with init_weights_on_device():
|
| 181 |
+
new_layer = quantized_layer.QLinear(module.in_features,module.out_features)
|
| 182 |
+
new_layer.weight = module.weight
|
| 183 |
+
if module.bias is not None:
|
| 184 |
+
new_layer.bias = module.bias
|
| 185 |
+
setattr(model, name, new_layer)
|
| 186 |
+
elif isinstance(module, RMSNorm):
|
| 187 |
+
if hasattr(module,"quantized"):
|
| 188 |
+
continue
|
| 189 |
+
module.quantized= True
|
| 190 |
+
new_layer = quantized_layer.QRMSNorm(module)
|
| 191 |
+
setattr(model, name, new_layer)
|
| 192 |
+
elif isinstance(module,torch.nn.Embedding):
|
| 193 |
+
rows, cols = module.weight.shape
|
| 194 |
+
new_layer = quantized_layer.QEmbedding(
|
| 195 |
+
num_embeddings=rows,
|
| 196 |
+
embedding_dim=cols,
|
| 197 |
+
_weight=module.weight,
|
| 198 |
+
# _freeze=module.freeze,
|
| 199 |
+
padding_idx=module.padding_idx,
|
| 200 |
+
max_norm=module.max_norm,
|
| 201 |
+
norm_type=module.norm_type,
|
| 202 |
+
scale_grad_by_freq=module.scale_grad_by_freq,
|
| 203 |
+
sparse=module.sparse)
|
| 204 |
+
setattr(model, name, new_layer)
|
| 205 |
+
else:
|
| 206 |
+
replace_layer(module)
|
| 207 |
+
|
| 208 |
+
replace_layer(self)
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
class FluxControlNetStateDictConverter:
|
| 213 |
+
def __init__(self):
|
| 214 |
+
pass
|
| 215 |
+
|
| 216 |
+
def from_diffusers(self, state_dict):
|
| 217 |
+
hash_value = hash_state_dict_keys(state_dict)
|
| 218 |
+
global_rename_dict = {
|
| 219 |
+
"context_embedder": "context_embedder",
|
| 220 |
+
"x_embedder": "x_embedder",
|
| 221 |
+
"time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",
|
| 222 |
+
"time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",
|
| 223 |
+
"time_text_embed.guidance_embedder.linear_1": "guidance_embedder.timestep_embedder.0",
|
| 224 |
+
"time_text_embed.guidance_embedder.linear_2": "guidance_embedder.timestep_embedder.2",
|
| 225 |
+
"time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",
|
| 226 |
+
"time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",
|
| 227 |
+
"norm_out.linear": "final_norm_out.linear",
|
| 228 |
+
"proj_out": "final_proj_out",
|
| 229 |
+
}
|
| 230 |
+
rename_dict = {
|
| 231 |
+
"proj_out": "proj_out",
|
| 232 |
+
"norm1.linear": "norm1_a.linear",
|
| 233 |
+
"norm1_context.linear": "norm1_b.linear",
|
| 234 |
+
"attn.to_q": "attn.a_to_q",
|
| 235 |
+
"attn.to_k": "attn.a_to_k",
|
| 236 |
+
"attn.to_v": "attn.a_to_v",
|
| 237 |
+
"attn.to_out.0": "attn.a_to_out",
|
| 238 |
+
"attn.add_q_proj": "attn.b_to_q",
|
| 239 |
+
"attn.add_k_proj": "attn.b_to_k",
|
| 240 |
+
"attn.add_v_proj": "attn.b_to_v",
|
| 241 |
+
"attn.to_add_out": "attn.b_to_out",
|
| 242 |
+
"ff.net.0.proj": "ff_a.0",
|
| 243 |
+
"ff.net.2": "ff_a.2",
|
| 244 |
+
"ff_context.net.0.proj": "ff_b.0",
|
| 245 |
+
"ff_context.net.2": "ff_b.2",
|
| 246 |
+
"attn.norm_q": "attn.norm_q_a",
|
| 247 |
+
"attn.norm_k": "attn.norm_k_a",
|
| 248 |
+
"attn.norm_added_q": "attn.norm_q_b",
|
| 249 |
+
"attn.norm_added_k": "attn.norm_k_b",
|
| 250 |
+
}
|
| 251 |
+
rename_dict_single = {
|
| 252 |
+
"attn.to_q": "a_to_q",
|
| 253 |
+
"attn.to_k": "a_to_k",
|
| 254 |
+
"attn.to_v": "a_to_v",
|
| 255 |
+
"attn.norm_q": "norm_q_a",
|
| 256 |
+
"attn.norm_k": "norm_k_a",
|
| 257 |
+
"norm.linear": "norm.linear",
|
| 258 |
+
"proj_mlp": "proj_in_besides_attn",
|
| 259 |
+
"proj_out": "proj_out",
|
| 260 |
+
}
|
| 261 |
+
state_dict_ = {}
|
| 262 |
+
for name, param in state_dict.items():
|
| 263 |
+
if name.endswith(".weight") or name.endswith(".bias"):
|
| 264 |
+
suffix = ".weight" if name.endswith(".weight") else ".bias"
|
| 265 |
+
prefix = name[:-len(suffix)]
|
| 266 |
+
if prefix in global_rename_dict:
|
| 267 |
+
state_dict_[global_rename_dict[prefix] + suffix] = param
|
| 268 |
+
elif prefix.startswith("transformer_blocks."):
|
| 269 |
+
names = prefix.split(".")
|
| 270 |
+
names[0] = "blocks"
|
| 271 |
+
middle = ".".join(names[2:])
|
| 272 |
+
if middle in rename_dict:
|
| 273 |
+
name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])
|
| 274 |
+
state_dict_[name_] = param
|
| 275 |
+
elif prefix.startswith("single_transformer_blocks."):
|
| 276 |
+
names = prefix.split(".")
|
| 277 |
+
names[0] = "single_blocks"
|
| 278 |
+
middle = ".".join(names[2:])
|
| 279 |
+
if middle in rename_dict_single:
|
| 280 |
+
name_ = ".".join(names[:2] + [rename_dict_single[middle]] + [suffix[1:]])
|
| 281 |
+
state_dict_[name_] = param
|
| 282 |
+
else:
|
| 283 |
+
state_dict_[name] = param
|
| 284 |
+
else:
|
| 285 |
+
state_dict_[name] = param
|
| 286 |
+
for name in list(state_dict_.keys()):
|
| 287 |
+
if ".proj_in_besides_attn." in name:
|
| 288 |
+
name_ = name.replace(".proj_in_besides_attn.", ".to_qkv_mlp.")
|
| 289 |
+
param = torch.concat([
|
| 290 |
+
state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_q.")],
|
| 291 |
+
state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_k.")],
|
| 292 |
+
state_dict_[name.replace(".proj_in_besides_attn.", f".a_to_v.")],
|
| 293 |
+
state_dict_[name],
|
| 294 |
+
], dim=0)
|
| 295 |
+
state_dict_[name_] = param
|
| 296 |
+
state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_q."))
|
| 297 |
+
state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_k."))
|
| 298 |
+
state_dict_.pop(name.replace(".proj_in_besides_attn.", f".a_to_v."))
|
| 299 |
+
state_dict_.pop(name)
|
| 300 |
+
for name in list(state_dict_.keys()):
|
| 301 |
+
for component in ["a", "b"]:
|
| 302 |
+
if f".{component}_to_q." in name:
|
| 303 |
+
name_ = name.replace(f".{component}_to_q.", f".{component}_to_qkv.")
|
| 304 |
+
param = torch.concat([
|
| 305 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_q.")],
|
| 306 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_k.")],
|
| 307 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_v.")],
|
| 308 |
+
], dim=0)
|
| 309 |
+
state_dict_[name_] = param
|
| 310 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_q."))
|
| 311 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_k."))
|
| 312 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_v."))
|
| 313 |
+
if hash_value == "78d18b9101345ff695f312e7e62538c0":
|
| 314 |
+
extra_kwargs = {"num_mode": 10, "mode_dict": {"canny": 0, "tile": 1, "depth": 2, "blur": 3, "pose": 4, "gray": 5, "lq": 6}}
|
| 315 |
+
elif hash_value == "b001c89139b5f053c715fe772362dd2a":
|
| 316 |
+
extra_kwargs = {"num_single_blocks": 0}
|
| 317 |
+
elif hash_value == "52357cb26250681367488a8954c271e8":
|
| 318 |
+
extra_kwargs = {"num_joint_blocks": 6, "num_single_blocks": 0, "additional_input_dim": 4}
|
| 319 |
+
elif hash_value == "0cfd1740758423a2a854d67c136d1e8c":
|
| 320 |
+
extra_kwargs = {"num_joint_blocks": 4, "num_single_blocks": 1}
|
| 321 |
+
elif hash_value == "7f9583eb8ba86642abb9a21a4b2c9e16":
|
| 322 |
+
extra_kwargs = {"num_joint_blocks": 4, "num_single_blocks": 10}
|
| 323 |
+
elif hash_value == "43ad5aaa27dd4ee01b832ed16773fa52":
|
| 324 |
+
extra_kwargs = {"num_joint_blocks": 6, "num_single_blocks": 0}
|
| 325 |
+
else:
|
| 326 |
+
extra_kwargs = {}
|
| 327 |
+
return state_dict_, extra_kwargs
|
| 328 |
+
|
| 329 |
+
|
| 330 |
+
def from_civitai(self, state_dict):
|
| 331 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/flux_dit.py
ADDED
|
@@ -0,0 +1,749 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd3_dit import TimestepEmbeddings, AdaLayerNorm, RMSNorm
|
| 3 |
+
from einops import rearrange
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
from .utils import init_weights_on_device, hash_state_dict_keys
|
| 6 |
+
|
| 7 |
+
def interact_with_ipadapter(hidden_states, q, ip_k, ip_v, scale=1.0):
|
| 8 |
+
batch_size, num_tokens = hidden_states.shape[0:2]
|
| 9 |
+
ip_hidden_states = torch.nn.functional.scaled_dot_product_attention(q, ip_k, ip_v)
|
| 10 |
+
ip_hidden_states = ip_hidden_states.transpose(1, 2).reshape(batch_size, num_tokens, -1)
|
| 11 |
+
hidden_states = hidden_states + scale * ip_hidden_states
|
| 12 |
+
return hidden_states
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class RoPEEmbedding(torch.nn.Module):
|
| 16 |
+
def __init__(self, dim, theta, axes_dim):
|
| 17 |
+
super().__init__()
|
| 18 |
+
self.dim = dim
|
| 19 |
+
self.theta = theta
|
| 20 |
+
self.axes_dim = axes_dim
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def rope(self, pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
|
| 24 |
+
assert dim % 2 == 0, "The dimension must be even."
|
| 25 |
+
|
| 26 |
+
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
| 27 |
+
omega = 1.0 / (theta**scale)
|
| 28 |
+
|
| 29 |
+
batch_size, seq_length = pos.shape
|
| 30 |
+
out = torch.einsum("...n,d->...nd", pos, omega)
|
| 31 |
+
cos_out = torch.cos(out)
|
| 32 |
+
sin_out = torch.sin(out)
|
| 33 |
+
|
| 34 |
+
stacked_out = torch.stack([cos_out, -sin_out, sin_out, cos_out], dim=-1)
|
| 35 |
+
out = stacked_out.view(batch_size, -1, dim // 2, 2, 2)
|
| 36 |
+
return out.float()
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def forward(self, ids):
|
| 40 |
+
n_axes = ids.shape[-1]
|
| 41 |
+
emb = torch.cat([self.rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)], dim=-3)
|
| 42 |
+
return emb.unsqueeze(1)
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class FluxJointAttention(torch.nn.Module):
|
| 47 |
+
def __init__(self, dim_a, dim_b, num_heads, head_dim, only_out_a=False):
|
| 48 |
+
super().__init__()
|
| 49 |
+
self.num_heads = num_heads
|
| 50 |
+
self.head_dim = head_dim
|
| 51 |
+
self.only_out_a = only_out_a
|
| 52 |
+
|
| 53 |
+
self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)
|
| 54 |
+
self.b_to_qkv = torch.nn.Linear(dim_b, dim_b * 3)
|
| 55 |
+
|
| 56 |
+
self.norm_q_a = RMSNorm(head_dim, eps=1e-6)
|
| 57 |
+
self.norm_k_a = RMSNorm(head_dim, eps=1e-6)
|
| 58 |
+
self.norm_q_b = RMSNorm(head_dim, eps=1e-6)
|
| 59 |
+
self.norm_k_b = RMSNorm(head_dim, eps=1e-6)
|
| 60 |
+
|
| 61 |
+
self.a_to_out = torch.nn.Linear(dim_a, dim_a)
|
| 62 |
+
if not only_out_a:
|
| 63 |
+
self.b_to_out = torch.nn.Linear(dim_b, dim_b)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def apply_rope(self, xq, xk, freqs_cis):
|
| 67 |
+
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
| 68 |
+
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
| 69 |
+
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
| 70 |
+
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
| 71 |
+
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
| 72 |
+
|
| 73 |
+
def forward(self, hidden_states_a, hidden_states_b, image_rotary_emb, attn_mask=None, ipadapter_kwargs_list=None):
|
| 74 |
+
batch_size = hidden_states_a.shape[0]
|
| 75 |
+
|
| 76 |
+
# Part A
|
| 77 |
+
qkv_a = self.a_to_qkv(hidden_states_a)
|
| 78 |
+
qkv_a = qkv_a.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 79 |
+
q_a, k_a, v_a = qkv_a.chunk(3, dim=1)
|
| 80 |
+
q_a, k_a = self.norm_q_a(q_a), self.norm_k_a(k_a)
|
| 81 |
+
|
| 82 |
+
# Part B
|
| 83 |
+
qkv_b = self.b_to_qkv(hidden_states_b)
|
| 84 |
+
qkv_b = qkv_b.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 85 |
+
q_b, k_b, v_b = qkv_b.chunk(3, dim=1)
|
| 86 |
+
q_b, k_b = self.norm_q_b(q_b), self.norm_k_b(k_b)
|
| 87 |
+
|
| 88 |
+
q = torch.concat([q_b, q_a], dim=2)
|
| 89 |
+
k = torch.concat([k_b, k_a], dim=2)
|
| 90 |
+
v = torch.concat([v_b, v_a], dim=2)
|
| 91 |
+
|
| 92 |
+
q, k = self.apply_rope(q, k, image_rotary_emb)
|
| 93 |
+
|
| 94 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 95 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 96 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 97 |
+
hidden_states_b, hidden_states_a = hidden_states[:, :hidden_states_b.shape[1]], hidden_states[:, hidden_states_b.shape[1]:]
|
| 98 |
+
if ipadapter_kwargs_list is not None:
|
| 99 |
+
hidden_states_a = interact_with_ipadapter(hidden_states_a, q_a, **ipadapter_kwargs_list)
|
| 100 |
+
hidden_states_a = self.a_to_out(hidden_states_a)
|
| 101 |
+
if self.only_out_a:
|
| 102 |
+
return hidden_states_a
|
| 103 |
+
else:
|
| 104 |
+
hidden_states_b = self.b_to_out(hidden_states_b)
|
| 105 |
+
return hidden_states_a, hidden_states_b
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
class FluxJointTransformerBlock(torch.nn.Module):
|
| 110 |
+
def __init__(self, dim, num_attention_heads):
|
| 111 |
+
super().__init__()
|
| 112 |
+
self.norm1_a = AdaLayerNorm(dim)
|
| 113 |
+
self.norm1_b = AdaLayerNorm(dim)
|
| 114 |
+
|
| 115 |
+
self.attn = FluxJointAttention(dim, dim, num_attention_heads, dim // num_attention_heads)
|
| 116 |
+
|
| 117 |
+
self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 118 |
+
self.ff_a = torch.nn.Sequential(
|
| 119 |
+
torch.nn.Linear(dim, dim*4),
|
| 120 |
+
torch.nn.GELU(approximate="tanh"),
|
| 121 |
+
torch.nn.Linear(dim*4, dim)
|
| 122 |
+
)
|
| 123 |
+
|
| 124 |
+
self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 125 |
+
self.ff_b = torch.nn.Sequential(
|
| 126 |
+
torch.nn.Linear(dim, dim*4),
|
| 127 |
+
torch.nn.GELU(approximate="tanh"),
|
| 128 |
+
torch.nn.Linear(dim*4, dim)
|
| 129 |
+
)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def forward(self, hidden_states_a, hidden_states_b, temb, image_rotary_emb, attn_mask=None, ipadapter_kwargs_list=None):
|
| 133 |
+
norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)
|
| 134 |
+
norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)
|
| 135 |
+
|
| 136 |
+
# Attention
|
| 137 |
+
attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b, image_rotary_emb, attn_mask, ipadapter_kwargs_list)
|
| 138 |
+
|
| 139 |
+
# Part A
|
| 140 |
+
hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a
|
| 141 |
+
norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a
|
| 142 |
+
hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)
|
| 143 |
+
|
| 144 |
+
# Part B
|
| 145 |
+
hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b
|
| 146 |
+
norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b
|
| 147 |
+
hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)
|
| 148 |
+
|
| 149 |
+
return hidden_states_a, hidden_states_b
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
class FluxSingleAttention(torch.nn.Module):
|
| 154 |
+
def __init__(self, dim_a, dim_b, num_heads, head_dim):
|
| 155 |
+
super().__init__()
|
| 156 |
+
self.num_heads = num_heads
|
| 157 |
+
self.head_dim = head_dim
|
| 158 |
+
|
| 159 |
+
self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)
|
| 160 |
+
|
| 161 |
+
self.norm_q_a = RMSNorm(head_dim, eps=1e-6)
|
| 162 |
+
self.norm_k_a = RMSNorm(head_dim, eps=1e-6)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def apply_rope(self, xq, xk, freqs_cis):
|
| 166 |
+
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
| 167 |
+
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
| 168 |
+
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
| 169 |
+
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
| 170 |
+
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def forward(self, hidden_states, image_rotary_emb):
|
| 174 |
+
batch_size = hidden_states.shape[0]
|
| 175 |
+
|
| 176 |
+
qkv_a = self.a_to_qkv(hidden_states)
|
| 177 |
+
qkv_a = qkv_a.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 178 |
+
q_a, k_a, v = qkv_a.chunk(3, dim=1)
|
| 179 |
+
q_a, k_a = self.norm_q_a(q_a), self.norm_k_a(k_a)
|
| 180 |
+
|
| 181 |
+
q, k = self.apply_rope(q_a, k_a, image_rotary_emb)
|
| 182 |
+
|
| 183 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
| 184 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 185 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 186 |
+
return hidden_states
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
class AdaLayerNormSingle(torch.nn.Module):
|
| 191 |
+
def __init__(self, dim):
|
| 192 |
+
super().__init__()
|
| 193 |
+
self.silu = torch.nn.SiLU()
|
| 194 |
+
self.linear = torch.nn.Linear(dim, 3 * dim, bias=True)
|
| 195 |
+
self.norm = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
def forward(self, x, emb):
|
| 199 |
+
emb = self.linear(self.silu(emb))
|
| 200 |
+
shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1)
|
| 201 |
+
x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]
|
| 202 |
+
return x, gate_msa
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
class FluxSingleTransformerBlock(torch.nn.Module):
|
| 207 |
+
def __init__(self, dim, num_attention_heads):
|
| 208 |
+
super().__init__()
|
| 209 |
+
self.num_heads = num_attention_heads
|
| 210 |
+
self.head_dim = dim // num_attention_heads
|
| 211 |
+
self.dim = dim
|
| 212 |
+
|
| 213 |
+
self.norm = AdaLayerNormSingle(dim)
|
| 214 |
+
self.to_qkv_mlp = torch.nn.Linear(dim, dim * (3 + 4))
|
| 215 |
+
self.norm_q_a = RMSNorm(self.head_dim, eps=1e-6)
|
| 216 |
+
self.norm_k_a = RMSNorm(self.head_dim, eps=1e-6)
|
| 217 |
+
|
| 218 |
+
self.proj_out = torch.nn.Linear(dim * 5, dim)
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def apply_rope(self, xq, xk, freqs_cis):
|
| 222 |
+
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
| 223 |
+
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
| 224 |
+
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
| 225 |
+
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
| 226 |
+
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
def process_attention(self, hidden_states, image_rotary_emb, attn_mask=None, ipadapter_kwargs_list=None):
|
| 230 |
+
batch_size = hidden_states.shape[0]
|
| 231 |
+
|
| 232 |
+
qkv = hidden_states.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 233 |
+
q, k, v = qkv.chunk(3, dim=1)
|
| 234 |
+
q, k = self.norm_q_a(q), self.norm_k_a(k)
|
| 235 |
+
|
| 236 |
+
q, k = self.apply_rope(q, k, image_rotary_emb)
|
| 237 |
+
|
| 238 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 239 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 240 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 241 |
+
if ipadapter_kwargs_list is not None:
|
| 242 |
+
hidden_states = interact_with_ipadapter(hidden_states, q, **ipadapter_kwargs_list)
|
| 243 |
+
return hidden_states
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def forward(self, hidden_states_a, hidden_states_b, temb, image_rotary_emb, attn_mask=None, ipadapter_kwargs_list=None):
|
| 247 |
+
residual = hidden_states_a
|
| 248 |
+
norm_hidden_states, gate = self.norm(hidden_states_a, emb=temb)
|
| 249 |
+
hidden_states_a = self.to_qkv_mlp(norm_hidden_states)
|
| 250 |
+
attn_output, mlp_hidden_states = hidden_states_a[:, :, :self.dim * 3], hidden_states_a[:, :, self.dim * 3:]
|
| 251 |
+
|
| 252 |
+
attn_output = self.process_attention(attn_output, image_rotary_emb, attn_mask, ipadapter_kwargs_list)
|
| 253 |
+
mlp_hidden_states = torch.nn.functional.gelu(mlp_hidden_states, approximate="tanh")
|
| 254 |
+
|
| 255 |
+
hidden_states_a = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
| 256 |
+
hidden_states_a = gate.unsqueeze(1) * self.proj_out(hidden_states_a)
|
| 257 |
+
hidden_states_a = residual + hidden_states_a
|
| 258 |
+
|
| 259 |
+
return hidden_states_a, hidden_states_b
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
|
| 263 |
+
class AdaLayerNormContinuous(torch.nn.Module):
|
| 264 |
+
def __init__(self, dim):
|
| 265 |
+
super().__init__()
|
| 266 |
+
self.silu = torch.nn.SiLU()
|
| 267 |
+
self.linear = torch.nn.Linear(dim, dim * 2, bias=True)
|
| 268 |
+
self.norm = torch.nn.LayerNorm(dim, eps=1e-6, elementwise_affine=False)
|
| 269 |
+
|
| 270 |
+
def forward(self, x, conditioning):
|
| 271 |
+
emb = self.linear(self.silu(conditioning))
|
| 272 |
+
scale, shift = torch.chunk(emb, 2, dim=1)
|
| 273 |
+
x = self.norm(x) * (1 + scale)[:, None] + shift[:, None]
|
| 274 |
+
return x
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
|
| 278 |
+
class FluxDiT(torch.nn.Module):
|
| 279 |
+
def __init__(self, disable_guidance_embedder=False, input_dim=64, num_blocks=19):
|
| 280 |
+
super().__init__()
|
| 281 |
+
self.pos_embedder = RoPEEmbedding(3072, 10000, [16, 56, 56])
|
| 282 |
+
self.time_embedder = TimestepEmbeddings(256, 3072)
|
| 283 |
+
self.guidance_embedder = None if disable_guidance_embedder else TimestepEmbeddings(256, 3072)
|
| 284 |
+
self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(768, 3072), torch.nn.SiLU(), torch.nn.Linear(3072, 3072))
|
| 285 |
+
self.context_embedder = torch.nn.Linear(4096, 3072)
|
| 286 |
+
self.x_embedder = torch.nn.Linear(input_dim, 3072)
|
| 287 |
+
|
| 288 |
+
self.blocks = torch.nn.ModuleList([FluxJointTransformerBlock(3072, 24) for _ in range(num_blocks)])
|
| 289 |
+
self.single_blocks = torch.nn.ModuleList([FluxSingleTransformerBlock(3072, 24) for _ in range(38)])
|
| 290 |
+
|
| 291 |
+
self.final_norm_out = AdaLayerNormContinuous(3072)
|
| 292 |
+
self.final_proj_out = torch.nn.Linear(3072, 64)
|
| 293 |
+
|
| 294 |
+
self.input_dim = input_dim
|
| 295 |
+
|
| 296 |
+
|
| 297 |
+
def patchify(self, hidden_states):
|
| 298 |
+
hidden_states = rearrange(hidden_states, "B C (H P) (W Q) -> B (H W) (C P Q)", P=2, Q=2)
|
| 299 |
+
return hidden_states
|
| 300 |
+
|
| 301 |
+
|
| 302 |
+
def unpatchify(self, hidden_states, height, width):
|
| 303 |
+
hidden_states = rearrange(hidden_states, "B (H W) (C P Q) -> B C (H P) (W Q)", P=2, Q=2, H=height//2, W=width//2)
|
| 304 |
+
return hidden_states
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
def prepare_image_ids(self, latents):
|
| 308 |
+
batch_size, _, height, width = latents.shape
|
| 309 |
+
latent_image_ids = torch.zeros(height // 2, width // 2, 3)
|
| 310 |
+
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height // 2)[:, None]
|
| 311 |
+
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width // 2)[None, :]
|
| 312 |
+
|
| 313 |
+
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
|
| 314 |
+
|
| 315 |
+
latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1)
|
| 316 |
+
latent_image_ids = latent_image_ids.reshape(
|
| 317 |
+
batch_size, latent_image_id_height * latent_image_id_width, latent_image_id_channels
|
| 318 |
+
)
|
| 319 |
+
latent_image_ids = latent_image_ids.to(device=latents.device, dtype=latents.dtype)
|
| 320 |
+
|
| 321 |
+
return latent_image_ids
|
| 322 |
+
|
| 323 |
+
|
| 324 |
+
def tiled_forward(
|
| 325 |
+
self,
|
| 326 |
+
hidden_states,
|
| 327 |
+
timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids,
|
| 328 |
+
tile_size=128, tile_stride=64,
|
| 329 |
+
**kwargs
|
| 330 |
+
):
|
| 331 |
+
# Due to the global positional embedding, we cannot implement layer-wise tiled forward.
|
| 332 |
+
hidden_states = TileWorker().tiled_forward(
|
| 333 |
+
lambda x: self.forward(x, timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None),
|
| 334 |
+
hidden_states,
|
| 335 |
+
tile_size,
|
| 336 |
+
tile_stride,
|
| 337 |
+
tile_device=hidden_states.device,
|
| 338 |
+
tile_dtype=hidden_states.dtype
|
| 339 |
+
)
|
| 340 |
+
return hidden_states
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
def construct_mask(self, entity_masks, prompt_seq_len, image_seq_len):
|
| 344 |
+
N = len(entity_masks)
|
| 345 |
+
batch_size = entity_masks[0].shape[0]
|
| 346 |
+
total_seq_len = N * prompt_seq_len + image_seq_len
|
| 347 |
+
patched_masks = [self.patchify(entity_masks[i]) for i in range(N)]
|
| 348 |
+
attention_mask = torch.ones((batch_size, total_seq_len, total_seq_len), dtype=torch.bool).to(device=entity_masks[0].device)
|
| 349 |
+
|
| 350 |
+
image_start = N * prompt_seq_len
|
| 351 |
+
image_end = N * prompt_seq_len + image_seq_len
|
| 352 |
+
# prompt-image mask
|
| 353 |
+
for i in range(N):
|
| 354 |
+
prompt_start = i * prompt_seq_len
|
| 355 |
+
prompt_end = (i + 1) * prompt_seq_len
|
| 356 |
+
image_mask = torch.sum(patched_masks[i], dim=-1) > 0
|
| 357 |
+
image_mask = image_mask.unsqueeze(1).repeat(1, prompt_seq_len, 1)
|
| 358 |
+
# prompt update with image
|
| 359 |
+
attention_mask[:, prompt_start:prompt_end, image_start:image_end] = image_mask
|
| 360 |
+
# image update with prompt
|
| 361 |
+
attention_mask[:, image_start:image_end, prompt_start:prompt_end] = image_mask.transpose(1, 2)
|
| 362 |
+
# prompt-prompt mask
|
| 363 |
+
for i in range(N):
|
| 364 |
+
for j in range(N):
|
| 365 |
+
if i != j:
|
| 366 |
+
prompt_start_i = i * prompt_seq_len
|
| 367 |
+
prompt_end_i = (i + 1) * prompt_seq_len
|
| 368 |
+
prompt_start_j = j * prompt_seq_len
|
| 369 |
+
prompt_end_j = (j + 1) * prompt_seq_len
|
| 370 |
+
attention_mask[:, prompt_start_i:prompt_end_i, prompt_start_j:prompt_end_j] = False
|
| 371 |
+
|
| 372 |
+
attention_mask = attention_mask.float()
|
| 373 |
+
attention_mask[attention_mask == 0] = float('-inf')
|
| 374 |
+
attention_mask[attention_mask == 1] = 0
|
| 375 |
+
return attention_mask
|
| 376 |
+
|
| 377 |
+
|
| 378 |
+
def process_entity_masks(self, hidden_states, prompt_emb, entity_prompt_emb, entity_masks, text_ids, image_ids):
|
| 379 |
+
repeat_dim = hidden_states.shape[1]
|
| 380 |
+
max_masks = 0
|
| 381 |
+
attention_mask = None
|
| 382 |
+
prompt_embs = [prompt_emb]
|
| 383 |
+
if entity_masks is not None:
|
| 384 |
+
# entity_masks
|
| 385 |
+
batch_size, max_masks = entity_masks.shape[0], entity_masks.shape[1]
|
| 386 |
+
entity_masks = entity_masks.repeat(1, 1, repeat_dim, 1, 1)
|
| 387 |
+
entity_masks = [entity_masks[:, i, None].squeeze(1) for i in range(max_masks)]
|
| 388 |
+
# global mask
|
| 389 |
+
global_mask = torch.ones_like(entity_masks[0]).to(device=hidden_states.device, dtype=hidden_states.dtype)
|
| 390 |
+
entity_masks = entity_masks + [global_mask] # append global to last
|
| 391 |
+
# attention mask
|
| 392 |
+
attention_mask = self.construct_mask(entity_masks, prompt_emb.shape[1], hidden_states.shape[1])
|
| 393 |
+
attention_mask = attention_mask.to(device=hidden_states.device, dtype=hidden_states.dtype)
|
| 394 |
+
attention_mask = attention_mask.unsqueeze(1)
|
| 395 |
+
# embds: n_masks * b * seq * d
|
| 396 |
+
local_embs = [entity_prompt_emb[:, i, None].squeeze(1) for i in range(max_masks)]
|
| 397 |
+
prompt_embs = local_embs + prompt_embs # append global to last
|
| 398 |
+
prompt_embs = [self.context_embedder(prompt_emb) for prompt_emb in prompt_embs]
|
| 399 |
+
prompt_emb = torch.cat(prompt_embs, dim=1)
|
| 400 |
+
|
| 401 |
+
# positional embedding
|
| 402 |
+
text_ids = torch.cat([text_ids] * (max_masks + 1), dim=1)
|
| 403 |
+
image_rotary_emb = self.pos_embedder(torch.cat((text_ids, image_ids), dim=1))
|
| 404 |
+
return prompt_emb, image_rotary_emb, attention_mask
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
def forward(
|
| 408 |
+
self,
|
| 409 |
+
hidden_states,
|
| 410 |
+
timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids, image_ids=None,
|
| 411 |
+
tiled=False, tile_size=128, tile_stride=64, entity_prompt_emb=None, entity_masks=None,
|
| 412 |
+
use_gradient_checkpointing=False,
|
| 413 |
+
**kwargs
|
| 414 |
+
):
|
| 415 |
+
if tiled:
|
| 416 |
+
return self.tiled_forward(
|
| 417 |
+
hidden_states,
|
| 418 |
+
timestep, prompt_emb, pooled_prompt_emb, guidance, text_ids,
|
| 419 |
+
tile_size=tile_size, tile_stride=tile_stride,
|
| 420 |
+
**kwargs
|
| 421 |
+
)
|
| 422 |
+
|
| 423 |
+
if image_ids is None:
|
| 424 |
+
image_ids = self.prepare_image_ids(hidden_states)
|
| 425 |
+
|
| 426 |
+
conditioning = self.time_embedder(timestep, hidden_states.dtype) + self.pooled_text_embedder(pooled_prompt_emb)
|
| 427 |
+
if self.guidance_embedder is not None:
|
| 428 |
+
guidance = guidance * 1000
|
| 429 |
+
conditioning = conditioning + self.guidance_embedder(guidance, hidden_states.dtype)
|
| 430 |
+
|
| 431 |
+
height, width = hidden_states.shape[-2:]
|
| 432 |
+
hidden_states = self.patchify(hidden_states)
|
| 433 |
+
hidden_states = self.x_embedder(hidden_states)
|
| 434 |
+
|
| 435 |
+
if entity_prompt_emb is not None and entity_masks is not None:
|
| 436 |
+
prompt_emb, image_rotary_emb, attention_mask = self.process_entity_masks(hidden_states, prompt_emb, entity_prompt_emb, entity_masks, text_ids, image_ids)
|
| 437 |
+
else:
|
| 438 |
+
prompt_emb = self.context_embedder(prompt_emb)
|
| 439 |
+
image_rotary_emb = self.pos_embedder(torch.cat((text_ids, image_ids), dim=1))
|
| 440 |
+
attention_mask = None
|
| 441 |
+
|
| 442 |
+
def create_custom_forward(module):
|
| 443 |
+
def custom_forward(*inputs):
|
| 444 |
+
return module(*inputs)
|
| 445 |
+
return custom_forward
|
| 446 |
+
|
| 447 |
+
for block in self.blocks:
|
| 448 |
+
if self.training and use_gradient_checkpointing:
|
| 449 |
+
hidden_states, prompt_emb = torch.utils.checkpoint.checkpoint(
|
| 450 |
+
create_custom_forward(block),
|
| 451 |
+
hidden_states, prompt_emb, conditioning, image_rotary_emb, attention_mask,
|
| 452 |
+
use_reentrant=False,
|
| 453 |
+
)
|
| 454 |
+
else:
|
| 455 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb, attention_mask)
|
| 456 |
+
|
| 457 |
+
hidden_states = torch.cat([prompt_emb, hidden_states], dim=1)
|
| 458 |
+
for block in self.single_blocks:
|
| 459 |
+
if self.training and use_gradient_checkpointing:
|
| 460 |
+
hidden_states, prompt_emb = torch.utils.checkpoint.checkpoint(
|
| 461 |
+
create_custom_forward(block),
|
| 462 |
+
hidden_states, prompt_emb, conditioning, image_rotary_emb, attention_mask,
|
| 463 |
+
use_reentrant=False,
|
| 464 |
+
)
|
| 465 |
+
else:
|
| 466 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning, image_rotary_emb, attention_mask)
|
| 467 |
+
hidden_states = hidden_states[:, prompt_emb.shape[1]:]
|
| 468 |
+
|
| 469 |
+
hidden_states = self.final_norm_out(hidden_states, conditioning)
|
| 470 |
+
hidden_states = self.final_proj_out(hidden_states)
|
| 471 |
+
hidden_states = self.unpatchify(hidden_states, height, width)
|
| 472 |
+
|
| 473 |
+
return hidden_states
|
| 474 |
+
|
| 475 |
+
|
| 476 |
+
def quantize(self):
|
| 477 |
+
def cast_to(weight, dtype=None, device=None, copy=False):
|
| 478 |
+
if device is None or weight.device == device:
|
| 479 |
+
if not copy:
|
| 480 |
+
if dtype is None or weight.dtype == dtype:
|
| 481 |
+
return weight
|
| 482 |
+
return weight.to(dtype=dtype, copy=copy)
|
| 483 |
+
|
| 484 |
+
r = torch.empty_like(weight, dtype=dtype, device=device)
|
| 485 |
+
r.copy_(weight)
|
| 486 |
+
return r
|
| 487 |
+
|
| 488 |
+
def cast_weight(s, input=None, dtype=None, device=None):
|
| 489 |
+
if input is not None:
|
| 490 |
+
if dtype is None:
|
| 491 |
+
dtype = input.dtype
|
| 492 |
+
if device is None:
|
| 493 |
+
device = input.device
|
| 494 |
+
weight = cast_to(s.weight, dtype, device)
|
| 495 |
+
return weight
|
| 496 |
+
|
| 497 |
+
def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
|
| 498 |
+
if input is not None:
|
| 499 |
+
if dtype is None:
|
| 500 |
+
dtype = input.dtype
|
| 501 |
+
if bias_dtype is None:
|
| 502 |
+
bias_dtype = dtype
|
| 503 |
+
if device is None:
|
| 504 |
+
device = input.device
|
| 505 |
+
bias = None
|
| 506 |
+
weight = cast_to(s.weight, dtype, device)
|
| 507 |
+
bias = cast_to(s.bias, bias_dtype, device)
|
| 508 |
+
return weight, bias
|
| 509 |
+
|
| 510 |
+
class quantized_layer:
|
| 511 |
+
class Linear(torch.nn.Linear):
|
| 512 |
+
def __init__(self, *args, **kwargs):
|
| 513 |
+
super().__init__(*args, **kwargs)
|
| 514 |
+
|
| 515 |
+
def forward(self,input,**kwargs):
|
| 516 |
+
weight,bias= cast_bias_weight(self,input)
|
| 517 |
+
return torch.nn.functional.linear(input,weight,bias)
|
| 518 |
+
|
| 519 |
+
class RMSNorm(torch.nn.Module):
|
| 520 |
+
def __init__(self, module):
|
| 521 |
+
super().__init__()
|
| 522 |
+
self.module = module
|
| 523 |
+
|
| 524 |
+
def forward(self,hidden_states,**kwargs):
|
| 525 |
+
weight= cast_weight(self.module,hidden_states)
|
| 526 |
+
input_dtype = hidden_states.dtype
|
| 527 |
+
variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)
|
| 528 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.module.eps)
|
| 529 |
+
hidden_states = hidden_states.to(input_dtype) * weight
|
| 530 |
+
return hidden_states
|
| 531 |
+
|
| 532 |
+
def replace_layer(model):
|
| 533 |
+
for name, module in model.named_children():
|
| 534 |
+
if isinstance(module, torch.nn.Linear):
|
| 535 |
+
with init_weights_on_device():
|
| 536 |
+
new_layer = quantized_layer.Linear(module.in_features,module.out_features)
|
| 537 |
+
new_layer.weight = module.weight
|
| 538 |
+
if module.bias is not None:
|
| 539 |
+
new_layer.bias = module.bias
|
| 540 |
+
# del module
|
| 541 |
+
setattr(model, name, new_layer)
|
| 542 |
+
elif isinstance(module, RMSNorm):
|
| 543 |
+
if hasattr(module,"quantized"):
|
| 544 |
+
continue
|
| 545 |
+
module.quantized= True
|
| 546 |
+
new_layer = quantized_layer.RMSNorm(module)
|
| 547 |
+
setattr(model, name, new_layer)
|
| 548 |
+
else:
|
| 549 |
+
replace_layer(module)
|
| 550 |
+
|
| 551 |
+
replace_layer(self)
|
| 552 |
+
|
| 553 |
+
|
| 554 |
+
@staticmethod
|
| 555 |
+
def state_dict_converter():
|
| 556 |
+
return FluxDiTStateDictConverter()
|
| 557 |
+
|
| 558 |
+
|
| 559 |
+
class FluxDiTStateDictConverter:
|
| 560 |
+
def __init__(self):
|
| 561 |
+
pass
|
| 562 |
+
|
| 563 |
+
def from_diffusers(self, state_dict):
|
| 564 |
+
global_rename_dict = {
|
| 565 |
+
"context_embedder": "context_embedder",
|
| 566 |
+
"x_embedder": "x_embedder",
|
| 567 |
+
"time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",
|
| 568 |
+
"time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",
|
| 569 |
+
"time_text_embed.guidance_embedder.linear_1": "guidance_embedder.timestep_embedder.0",
|
| 570 |
+
"time_text_embed.guidance_embedder.linear_2": "guidance_embedder.timestep_embedder.2",
|
| 571 |
+
"time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",
|
| 572 |
+
"time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",
|
| 573 |
+
"norm_out.linear": "final_norm_out.linear",
|
| 574 |
+
"proj_out": "final_proj_out",
|
| 575 |
+
}
|
| 576 |
+
rename_dict = {
|
| 577 |
+
"proj_out": "proj_out",
|
| 578 |
+
"norm1.linear": "norm1_a.linear",
|
| 579 |
+
"norm1_context.linear": "norm1_b.linear",
|
| 580 |
+
"attn.to_q": "attn.a_to_q",
|
| 581 |
+
"attn.to_k": "attn.a_to_k",
|
| 582 |
+
"attn.to_v": "attn.a_to_v",
|
| 583 |
+
"attn.to_out.0": "attn.a_to_out",
|
| 584 |
+
"attn.add_q_proj": "attn.b_to_q",
|
| 585 |
+
"attn.add_k_proj": "attn.b_to_k",
|
| 586 |
+
"attn.add_v_proj": "attn.b_to_v",
|
| 587 |
+
"attn.to_add_out": "attn.b_to_out",
|
| 588 |
+
"ff.net.0.proj": "ff_a.0",
|
| 589 |
+
"ff.net.2": "ff_a.2",
|
| 590 |
+
"ff_context.net.0.proj": "ff_b.0",
|
| 591 |
+
"ff_context.net.2": "ff_b.2",
|
| 592 |
+
"attn.norm_q": "attn.norm_q_a",
|
| 593 |
+
"attn.norm_k": "attn.norm_k_a",
|
| 594 |
+
"attn.norm_added_q": "attn.norm_q_b",
|
| 595 |
+
"attn.norm_added_k": "attn.norm_k_b",
|
| 596 |
+
}
|
| 597 |
+
rename_dict_single = {
|
| 598 |
+
"attn.to_q": "a_to_q",
|
| 599 |
+
"attn.to_k": "a_to_k",
|
| 600 |
+
"attn.to_v": "a_to_v",
|
| 601 |
+
"attn.norm_q": "norm_q_a",
|
| 602 |
+
"attn.norm_k": "norm_k_a",
|
| 603 |
+
"norm.linear": "norm.linear",
|
| 604 |
+
"proj_mlp": "proj_in_besides_attn",
|
| 605 |
+
"proj_out": "proj_out",
|
| 606 |
+
}
|
| 607 |
+
state_dict_ = {}
|
| 608 |
+
for name, param in state_dict.items():
|
| 609 |
+
if name.endswith(".weight") or name.endswith(".bias"):
|
| 610 |
+
suffix = ".weight" if name.endswith(".weight") else ".bias"
|
| 611 |
+
prefix = name[:-len(suffix)]
|
| 612 |
+
if prefix in global_rename_dict:
|
| 613 |
+
state_dict_[global_rename_dict[prefix] + suffix] = param
|
| 614 |
+
elif prefix.startswith("transformer_blocks."):
|
| 615 |
+
names = prefix.split(".")
|
| 616 |
+
names[0] = "blocks"
|
| 617 |
+
middle = ".".join(names[2:])
|
| 618 |
+
if middle in rename_dict:
|
| 619 |
+
name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])
|
| 620 |
+
state_dict_[name_] = param
|
| 621 |
+
elif prefix.startswith("single_transformer_blocks."):
|
| 622 |
+
names = prefix.split(".")
|
| 623 |
+
names[0] = "single_blocks"
|
| 624 |
+
middle = ".".join(names[2:])
|
| 625 |
+
if middle in rename_dict_single:
|
| 626 |
+
name_ = ".".join(names[:2] + [rename_dict_single[middle]] + [suffix[1:]])
|
| 627 |
+
state_dict_[name_] = param
|
| 628 |
+
else:
|
| 629 |
+
pass
|
| 630 |
+
else:
|
| 631 |
+
pass
|
| 632 |
+
for name in list(state_dict_.keys()):
|
| 633 |
+
if "single_blocks." in name and ".a_to_q." in name:
|
| 634 |
+
mlp = state_dict_.get(name.replace(".a_to_q.", ".proj_in_besides_attn."), None)
|
| 635 |
+
if mlp is None:
|
| 636 |
+
mlp = torch.zeros(4 * state_dict_[name].shape[0],
|
| 637 |
+
*state_dict_[name].shape[1:],
|
| 638 |
+
dtype=state_dict_[name].dtype)
|
| 639 |
+
else:
|
| 640 |
+
state_dict_.pop(name.replace(".a_to_q.", ".proj_in_besides_attn."))
|
| 641 |
+
param = torch.concat([
|
| 642 |
+
state_dict_.pop(name),
|
| 643 |
+
state_dict_.pop(name.replace(".a_to_q.", ".a_to_k.")),
|
| 644 |
+
state_dict_.pop(name.replace(".a_to_q.", ".a_to_v.")),
|
| 645 |
+
mlp,
|
| 646 |
+
], dim=0)
|
| 647 |
+
name_ = name.replace(".a_to_q.", ".to_qkv_mlp.")
|
| 648 |
+
state_dict_[name_] = param
|
| 649 |
+
for name in list(state_dict_.keys()):
|
| 650 |
+
for component in ["a", "b"]:
|
| 651 |
+
if f".{component}_to_q." in name:
|
| 652 |
+
name_ = name.replace(f".{component}_to_q.", f".{component}_to_qkv.")
|
| 653 |
+
param = torch.concat([
|
| 654 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_q.")],
|
| 655 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_k.")],
|
| 656 |
+
state_dict_[name.replace(f".{component}_to_q.", f".{component}_to_v.")],
|
| 657 |
+
], dim=0)
|
| 658 |
+
state_dict_[name_] = param
|
| 659 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_q."))
|
| 660 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_k."))
|
| 661 |
+
state_dict_.pop(name.replace(f".{component}_to_q.", f".{component}_to_v."))
|
| 662 |
+
return state_dict_
|
| 663 |
+
|
| 664 |
+
def from_civitai(self, state_dict):
|
| 665 |
+
if hash_state_dict_keys(state_dict, with_shape=True) in ["3e6c61b0f9471135fc9c6d6a98e98b6d", "63c969fd37cce769a90aa781fbff5f81"]:
|
| 666 |
+
dit_state_dict = {key.replace("pipe.dit.", ""): value for key, value in state_dict.items() if key.startswith('pipe.dit.')}
|
| 667 |
+
return dit_state_dict
|
| 668 |
+
rename_dict = {
|
| 669 |
+
"time_in.in_layer.bias": "time_embedder.timestep_embedder.0.bias",
|
| 670 |
+
"time_in.in_layer.weight": "time_embedder.timestep_embedder.0.weight",
|
| 671 |
+
"time_in.out_layer.bias": "time_embedder.timestep_embedder.2.bias",
|
| 672 |
+
"time_in.out_layer.weight": "time_embedder.timestep_embedder.2.weight",
|
| 673 |
+
"txt_in.bias": "context_embedder.bias",
|
| 674 |
+
"txt_in.weight": "context_embedder.weight",
|
| 675 |
+
"vector_in.in_layer.bias": "pooled_text_embedder.0.bias",
|
| 676 |
+
"vector_in.in_layer.weight": "pooled_text_embedder.0.weight",
|
| 677 |
+
"vector_in.out_layer.bias": "pooled_text_embedder.2.bias",
|
| 678 |
+
"vector_in.out_layer.weight": "pooled_text_embedder.2.weight",
|
| 679 |
+
"final_layer.linear.bias": "final_proj_out.bias",
|
| 680 |
+
"final_layer.linear.weight": "final_proj_out.weight",
|
| 681 |
+
"guidance_in.in_layer.bias": "guidance_embedder.timestep_embedder.0.bias",
|
| 682 |
+
"guidance_in.in_layer.weight": "guidance_embedder.timestep_embedder.0.weight",
|
| 683 |
+
"guidance_in.out_layer.bias": "guidance_embedder.timestep_embedder.2.bias",
|
| 684 |
+
"guidance_in.out_layer.weight": "guidance_embedder.timestep_embedder.2.weight",
|
| 685 |
+
"img_in.bias": "x_embedder.bias",
|
| 686 |
+
"img_in.weight": "x_embedder.weight",
|
| 687 |
+
"final_layer.adaLN_modulation.1.weight": "final_norm_out.linear.weight",
|
| 688 |
+
"final_layer.adaLN_modulation.1.bias": "final_norm_out.linear.bias",
|
| 689 |
+
}
|
| 690 |
+
suffix_rename_dict = {
|
| 691 |
+
"img_attn.norm.key_norm.scale": "attn.norm_k_a.weight",
|
| 692 |
+
"img_attn.norm.query_norm.scale": "attn.norm_q_a.weight",
|
| 693 |
+
"img_attn.proj.bias": "attn.a_to_out.bias",
|
| 694 |
+
"img_attn.proj.weight": "attn.a_to_out.weight",
|
| 695 |
+
"img_attn.qkv.bias": "attn.a_to_qkv.bias",
|
| 696 |
+
"img_attn.qkv.weight": "attn.a_to_qkv.weight",
|
| 697 |
+
"img_mlp.0.bias": "ff_a.0.bias",
|
| 698 |
+
"img_mlp.0.weight": "ff_a.0.weight",
|
| 699 |
+
"img_mlp.2.bias": "ff_a.2.bias",
|
| 700 |
+
"img_mlp.2.weight": "ff_a.2.weight",
|
| 701 |
+
"img_mod.lin.bias": "norm1_a.linear.bias",
|
| 702 |
+
"img_mod.lin.weight": "norm1_a.linear.weight",
|
| 703 |
+
"txt_attn.norm.key_norm.scale": "attn.norm_k_b.weight",
|
| 704 |
+
"txt_attn.norm.query_norm.scale": "attn.norm_q_b.weight",
|
| 705 |
+
"txt_attn.proj.bias": "attn.b_to_out.bias",
|
| 706 |
+
"txt_attn.proj.weight": "attn.b_to_out.weight",
|
| 707 |
+
"txt_attn.qkv.bias": "attn.b_to_qkv.bias",
|
| 708 |
+
"txt_attn.qkv.weight": "attn.b_to_qkv.weight",
|
| 709 |
+
"txt_mlp.0.bias": "ff_b.0.bias",
|
| 710 |
+
"txt_mlp.0.weight": "ff_b.0.weight",
|
| 711 |
+
"txt_mlp.2.bias": "ff_b.2.bias",
|
| 712 |
+
"txt_mlp.2.weight": "ff_b.2.weight",
|
| 713 |
+
"txt_mod.lin.bias": "norm1_b.linear.bias",
|
| 714 |
+
"txt_mod.lin.weight": "norm1_b.linear.weight",
|
| 715 |
+
|
| 716 |
+
"linear1.bias": "to_qkv_mlp.bias",
|
| 717 |
+
"linear1.weight": "to_qkv_mlp.weight",
|
| 718 |
+
"linear2.bias": "proj_out.bias",
|
| 719 |
+
"linear2.weight": "proj_out.weight",
|
| 720 |
+
"modulation.lin.bias": "norm.linear.bias",
|
| 721 |
+
"modulation.lin.weight": "norm.linear.weight",
|
| 722 |
+
"norm.key_norm.scale": "norm_k_a.weight",
|
| 723 |
+
"norm.query_norm.scale": "norm_q_a.weight",
|
| 724 |
+
}
|
| 725 |
+
state_dict_ = {}
|
| 726 |
+
for name, param in state_dict.items():
|
| 727 |
+
if name.startswith("model.diffusion_model."):
|
| 728 |
+
name = name[len("model.diffusion_model."):]
|
| 729 |
+
names = name.split(".")
|
| 730 |
+
if name in rename_dict:
|
| 731 |
+
rename = rename_dict[name]
|
| 732 |
+
if name.startswith("final_layer.adaLN_modulation.1."):
|
| 733 |
+
param = torch.concat([param[3072:], param[:3072]], dim=0)
|
| 734 |
+
state_dict_[rename] = param
|
| 735 |
+
elif names[0] == "double_blocks":
|
| 736 |
+
rename = f"blocks.{names[1]}." + suffix_rename_dict[".".join(names[2:])]
|
| 737 |
+
state_dict_[rename] = param
|
| 738 |
+
elif names[0] == "single_blocks":
|
| 739 |
+
if ".".join(names[2:]) in suffix_rename_dict:
|
| 740 |
+
rename = f"single_blocks.{names[1]}." + suffix_rename_dict[".".join(names[2:])]
|
| 741 |
+
state_dict_[rename] = param
|
| 742 |
+
else:
|
| 743 |
+
pass
|
| 744 |
+
if "guidance_embedder.timestep_embedder.0.weight" not in state_dict_:
|
| 745 |
+
return state_dict_, {"disable_guidance_embedder": True}
|
| 746 |
+
elif "blocks.8.attn.norm_k_a.weight" not in state_dict_:
|
| 747 |
+
return state_dict_, {"input_dim": 196, "num_blocks": 8}
|
| 748 |
+
else:
|
| 749 |
+
return state_dict_
|
diffsynth/models/flux_infiniteyou.py
ADDED
|
@@ -0,0 +1,129 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import math
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
# FFN
|
| 7 |
+
def FeedForward(dim, mult=4):
|
| 8 |
+
inner_dim = int(dim * mult)
|
| 9 |
+
return nn.Sequential(
|
| 10 |
+
nn.LayerNorm(dim),
|
| 11 |
+
nn.Linear(dim, inner_dim, bias=False),
|
| 12 |
+
nn.GELU(),
|
| 13 |
+
nn.Linear(inner_dim, dim, bias=False),
|
| 14 |
+
)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def reshape_tensor(x, heads):
|
| 18 |
+
bs, length, width = x.shape
|
| 19 |
+
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
| 20 |
+
x = x.view(bs, length, heads, -1)
|
| 21 |
+
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
| 22 |
+
x = x.transpose(1, 2)
|
| 23 |
+
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
| 24 |
+
x = x.reshape(bs, heads, length, -1)
|
| 25 |
+
return x
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class PerceiverAttention(nn.Module):
|
| 29 |
+
|
| 30 |
+
def __init__(self, *, dim, dim_head=64, heads=8):
|
| 31 |
+
super().__init__()
|
| 32 |
+
self.scale = dim_head**-0.5
|
| 33 |
+
self.dim_head = dim_head
|
| 34 |
+
self.heads = heads
|
| 35 |
+
inner_dim = dim_head * heads
|
| 36 |
+
|
| 37 |
+
self.norm1 = nn.LayerNorm(dim)
|
| 38 |
+
self.norm2 = nn.LayerNorm(dim)
|
| 39 |
+
|
| 40 |
+
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
| 41 |
+
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
| 42 |
+
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
| 43 |
+
|
| 44 |
+
def forward(self, x, latents):
|
| 45 |
+
"""
|
| 46 |
+
Args:
|
| 47 |
+
x (torch.Tensor): image features
|
| 48 |
+
shape (b, n1, D)
|
| 49 |
+
latent (torch.Tensor): latent features
|
| 50 |
+
shape (b, n2, D)
|
| 51 |
+
"""
|
| 52 |
+
x = self.norm1(x)
|
| 53 |
+
latents = self.norm2(latents)
|
| 54 |
+
|
| 55 |
+
b, l, _ = latents.shape
|
| 56 |
+
|
| 57 |
+
q = self.to_q(latents)
|
| 58 |
+
kv_input = torch.cat((x, latents), dim=-2)
|
| 59 |
+
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
| 60 |
+
|
| 61 |
+
q = reshape_tensor(q, self.heads)
|
| 62 |
+
k = reshape_tensor(k, self.heads)
|
| 63 |
+
v = reshape_tensor(v, self.heads)
|
| 64 |
+
|
| 65 |
+
# attention
|
| 66 |
+
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
| 67 |
+
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
| 68 |
+
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
| 69 |
+
out = weight @ v
|
| 70 |
+
|
| 71 |
+
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
| 72 |
+
|
| 73 |
+
return self.to_out(out)
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
class InfiniteYouImageProjector(nn.Module):
|
| 77 |
+
|
| 78 |
+
def __init__(
|
| 79 |
+
self,
|
| 80 |
+
dim=1280,
|
| 81 |
+
depth=4,
|
| 82 |
+
dim_head=64,
|
| 83 |
+
heads=20,
|
| 84 |
+
num_queries=8,
|
| 85 |
+
embedding_dim=512,
|
| 86 |
+
output_dim=4096,
|
| 87 |
+
ff_mult=4,
|
| 88 |
+
):
|
| 89 |
+
super().__init__()
|
| 90 |
+
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
|
| 91 |
+
self.proj_in = nn.Linear(embedding_dim, dim)
|
| 92 |
+
|
| 93 |
+
self.proj_out = nn.Linear(dim, output_dim)
|
| 94 |
+
self.norm_out = nn.LayerNorm(output_dim)
|
| 95 |
+
|
| 96 |
+
self.layers = nn.ModuleList([])
|
| 97 |
+
for _ in range(depth):
|
| 98 |
+
self.layers.append(
|
| 99 |
+
nn.ModuleList([
|
| 100 |
+
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
| 101 |
+
FeedForward(dim=dim, mult=ff_mult),
|
| 102 |
+
]))
|
| 103 |
+
|
| 104 |
+
def forward(self, x):
|
| 105 |
+
|
| 106 |
+
latents = self.latents.repeat(x.size(0), 1, 1)
|
| 107 |
+
latents = latents.to(dtype=x.dtype, device=x.device)
|
| 108 |
+
|
| 109 |
+
x = self.proj_in(x)
|
| 110 |
+
|
| 111 |
+
for attn, ff in self.layers:
|
| 112 |
+
latents = attn(x, latents) + latents
|
| 113 |
+
latents = ff(latents) + latents
|
| 114 |
+
|
| 115 |
+
latents = self.proj_out(latents)
|
| 116 |
+
return self.norm_out(latents)
|
| 117 |
+
|
| 118 |
+
@staticmethod
|
| 119 |
+
def state_dict_converter():
|
| 120 |
+
return FluxInfiniteYouImageProjectorStateDictConverter()
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
class FluxInfiniteYouImageProjectorStateDictConverter:
|
| 124 |
+
|
| 125 |
+
def __init__(self):
|
| 126 |
+
pass
|
| 127 |
+
|
| 128 |
+
def from_diffusers(self, state_dict):
|
| 129 |
+
return state_dict['image_proj']
|
diffsynth/models/flux_ipadapter.py
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .svd_image_encoder import SVDImageEncoder
|
| 2 |
+
from .sd3_dit import RMSNorm
|
| 3 |
+
from transformers import CLIPImageProcessor
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class MLPProjModel(torch.nn.Module):
|
| 8 |
+
def __init__(self, cross_attention_dim=768, id_embeddings_dim=512, num_tokens=4):
|
| 9 |
+
super().__init__()
|
| 10 |
+
|
| 11 |
+
self.cross_attention_dim = cross_attention_dim
|
| 12 |
+
self.num_tokens = num_tokens
|
| 13 |
+
|
| 14 |
+
self.proj = torch.nn.Sequential(
|
| 15 |
+
torch.nn.Linear(id_embeddings_dim, id_embeddings_dim*2),
|
| 16 |
+
torch.nn.GELU(),
|
| 17 |
+
torch.nn.Linear(id_embeddings_dim*2, cross_attention_dim*num_tokens),
|
| 18 |
+
)
|
| 19 |
+
self.norm = torch.nn.LayerNorm(cross_attention_dim)
|
| 20 |
+
|
| 21 |
+
def forward(self, id_embeds):
|
| 22 |
+
x = self.proj(id_embeds)
|
| 23 |
+
x = x.reshape(-1, self.num_tokens, self.cross_attention_dim)
|
| 24 |
+
x = self.norm(x)
|
| 25 |
+
return x
|
| 26 |
+
|
| 27 |
+
class IpAdapterModule(torch.nn.Module):
|
| 28 |
+
def __init__(self, num_attention_heads, attention_head_dim, input_dim):
|
| 29 |
+
super().__init__()
|
| 30 |
+
self.num_heads = num_attention_heads
|
| 31 |
+
self.head_dim = attention_head_dim
|
| 32 |
+
output_dim = num_attention_heads * attention_head_dim
|
| 33 |
+
self.to_k_ip = torch.nn.Linear(input_dim, output_dim, bias=False)
|
| 34 |
+
self.to_v_ip = torch.nn.Linear(input_dim, output_dim, bias=False)
|
| 35 |
+
self.norm_added_k = RMSNorm(attention_head_dim, eps=1e-5, elementwise_affine=False)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def forward(self, hidden_states):
|
| 39 |
+
batch_size = hidden_states.shape[0]
|
| 40 |
+
# ip_k
|
| 41 |
+
ip_k = self.to_k_ip(hidden_states)
|
| 42 |
+
ip_k = ip_k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 43 |
+
ip_k = self.norm_added_k(ip_k)
|
| 44 |
+
# ip_v
|
| 45 |
+
ip_v = self.to_v_ip(hidden_states)
|
| 46 |
+
ip_v = ip_v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 47 |
+
return ip_k, ip_v
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class FluxIpAdapter(torch.nn.Module):
|
| 51 |
+
def __init__(self, num_attention_heads=24, attention_head_dim=128, cross_attention_dim=4096, num_tokens=128, num_blocks=57):
|
| 52 |
+
super().__init__()
|
| 53 |
+
self.ipadapter_modules = torch.nn.ModuleList([IpAdapterModule(num_attention_heads, attention_head_dim, cross_attention_dim) for _ in range(num_blocks)])
|
| 54 |
+
self.image_proj = MLPProjModel(cross_attention_dim=cross_attention_dim, id_embeddings_dim=1152, num_tokens=num_tokens)
|
| 55 |
+
self.set_adapter()
|
| 56 |
+
|
| 57 |
+
def set_adapter(self):
|
| 58 |
+
self.call_block_id = {i:i for i in range(len(self.ipadapter_modules))}
|
| 59 |
+
|
| 60 |
+
def forward(self, hidden_states, scale=1.0):
|
| 61 |
+
hidden_states = self.image_proj(hidden_states)
|
| 62 |
+
hidden_states = hidden_states.view(1, -1, hidden_states.shape[-1])
|
| 63 |
+
ip_kv_dict = {}
|
| 64 |
+
for block_id in self.call_block_id:
|
| 65 |
+
ipadapter_id = self.call_block_id[block_id]
|
| 66 |
+
ip_k, ip_v = self.ipadapter_modules[ipadapter_id](hidden_states)
|
| 67 |
+
ip_kv_dict[block_id] = {
|
| 68 |
+
"ip_k": ip_k,
|
| 69 |
+
"ip_v": ip_v,
|
| 70 |
+
"scale": scale
|
| 71 |
+
}
|
| 72 |
+
return ip_kv_dict
|
| 73 |
+
|
| 74 |
+
@staticmethod
|
| 75 |
+
def state_dict_converter():
|
| 76 |
+
return FluxIpAdapterStateDictConverter()
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class FluxIpAdapterStateDictConverter:
|
| 80 |
+
def __init__(self):
|
| 81 |
+
pass
|
| 82 |
+
|
| 83 |
+
def from_diffusers(self, state_dict):
|
| 84 |
+
state_dict_ = {}
|
| 85 |
+
for name in state_dict["ip_adapter"]:
|
| 86 |
+
name_ = 'ipadapter_modules.' + name
|
| 87 |
+
state_dict_[name_] = state_dict["ip_adapter"][name]
|
| 88 |
+
for name in state_dict["image_proj"]:
|
| 89 |
+
name_ = "image_proj." + name
|
| 90 |
+
state_dict_[name_] = state_dict["image_proj"][name]
|
| 91 |
+
return state_dict_
|
| 92 |
+
|
| 93 |
+
def from_civitai(self, state_dict):
|
| 94 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/flux_lora_encoder.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_text_encoder import CLIPEncoderLayer
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class LoRALayerBlock(torch.nn.Module):
|
| 6 |
+
def __init__(self, L, dim_in, dim_out):
|
| 7 |
+
super().__init__()
|
| 8 |
+
self.x = torch.nn.Parameter(torch.randn(1, L, dim_in))
|
| 9 |
+
self.layer_norm = torch.nn.LayerNorm(dim_out)
|
| 10 |
+
|
| 11 |
+
def forward(self, lora_A, lora_B):
|
| 12 |
+
x = self.x @ lora_A.T @ lora_B.T
|
| 13 |
+
x = self.layer_norm(x)
|
| 14 |
+
return x
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class LoRAEmbedder(torch.nn.Module):
|
| 18 |
+
def __init__(self, lora_patterns=None, L=1, out_dim=2048):
|
| 19 |
+
super().__init__()
|
| 20 |
+
if lora_patterns is None:
|
| 21 |
+
lora_patterns = self.default_lora_patterns()
|
| 22 |
+
|
| 23 |
+
model_dict = {}
|
| 24 |
+
for lora_pattern in lora_patterns:
|
| 25 |
+
name, dim = lora_pattern["name"], lora_pattern["dim"]
|
| 26 |
+
model_dict[name.replace(".", "___")] = LoRALayerBlock(L, dim[0], dim[1])
|
| 27 |
+
self.model_dict = torch.nn.ModuleDict(model_dict)
|
| 28 |
+
|
| 29 |
+
proj_dict = {}
|
| 30 |
+
for lora_pattern in lora_patterns:
|
| 31 |
+
layer_type, dim = lora_pattern["type"], lora_pattern["dim"]
|
| 32 |
+
if layer_type not in proj_dict:
|
| 33 |
+
proj_dict[layer_type.replace(".", "___")] = torch.nn.Linear(dim[1], out_dim)
|
| 34 |
+
self.proj_dict = torch.nn.ModuleDict(proj_dict)
|
| 35 |
+
|
| 36 |
+
self.lora_patterns = lora_patterns
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def default_lora_patterns(self):
|
| 40 |
+
lora_patterns = []
|
| 41 |
+
lora_dict = {
|
| 42 |
+
"attn.a_to_qkv": (3072, 9216), "attn.a_to_out": (3072, 3072), "ff_a.0": (3072, 12288), "ff_a.2": (12288, 3072), "norm1_a.linear": (3072, 18432),
|
| 43 |
+
"attn.b_to_qkv": (3072, 9216), "attn.b_to_out": (3072, 3072), "ff_b.0": (3072, 12288), "ff_b.2": (12288, 3072), "norm1_b.linear": (3072, 18432),
|
| 44 |
+
}
|
| 45 |
+
for i in range(19):
|
| 46 |
+
for suffix in lora_dict:
|
| 47 |
+
lora_patterns.append({
|
| 48 |
+
"name": f"blocks.{i}.{suffix}",
|
| 49 |
+
"dim": lora_dict[suffix],
|
| 50 |
+
"type": suffix,
|
| 51 |
+
})
|
| 52 |
+
lora_dict = {"to_qkv_mlp": (3072, 21504), "proj_out": (15360, 3072), "norm.linear": (3072, 9216)}
|
| 53 |
+
for i in range(38):
|
| 54 |
+
for suffix in lora_dict:
|
| 55 |
+
lora_patterns.append({
|
| 56 |
+
"name": f"single_blocks.{i}.{suffix}",
|
| 57 |
+
"dim": lora_dict[suffix],
|
| 58 |
+
"type": suffix,
|
| 59 |
+
})
|
| 60 |
+
return lora_patterns
|
| 61 |
+
|
| 62 |
+
def forward(self, lora):
|
| 63 |
+
lora_emb = []
|
| 64 |
+
for lora_pattern in self.lora_patterns:
|
| 65 |
+
name, layer_type = lora_pattern["name"], lora_pattern["type"]
|
| 66 |
+
lora_A = lora[name + ".lora_A.default.weight"]
|
| 67 |
+
lora_B = lora[name + ".lora_B.default.weight"]
|
| 68 |
+
lora_out = self.model_dict[name.replace(".", "___")](lora_A, lora_B)
|
| 69 |
+
lora_out = self.proj_dict[layer_type.replace(".", "___")](lora_out)
|
| 70 |
+
lora_emb.append(lora_out)
|
| 71 |
+
lora_emb = torch.concat(lora_emb, dim=1)
|
| 72 |
+
return lora_emb
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class FluxLoRAEncoder(torch.nn.Module):
|
| 76 |
+
def __init__(self, embed_dim=4096, encoder_intermediate_size=8192, num_encoder_layers=1, num_embeds_per_lora=16, num_special_embeds=1):
|
| 77 |
+
super().__init__()
|
| 78 |
+
self.num_embeds_per_lora = num_embeds_per_lora
|
| 79 |
+
# embedder
|
| 80 |
+
self.embedder = LoRAEmbedder(L=num_embeds_per_lora, out_dim=embed_dim)
|
| 81 |
+
|
| 82 |
+
# encoders
|
| 83 |
+
self.encoders = torch.nn.ModuleList([CLIPEncoderLayer(embed_dim, encoder_intermediate_size, num_heads=32, head_dim=128) for _ in range(num_encoder_layers)])
|
| 84 |
+
|
| 85 |
+
# special embedding
|
| 86 |
+
self.special_embeds = torch.nn.Parameter(torch.randn(1, num_special_embeds, embed_dim))
|
| 87 |
+
self.num_special_embeds = num_special_embeds
|
| 88 |
+
|
| 89 |
+
# final layer
|
| 90 |
+
self.final_layer_norm = torch.nn.LayerNorm(embed_dim)
|
| 91 |
+
self.final_linear = torch.nn.Linear(embed_dim, embed_dim)
|
| 92 |
+
|
| 93 |
+
def forward(self, lora):
|
| 94 |
+
lora_embeds = self.embedder(lora)
|
| 95 |
+
special_embeds = self.special_embeds.to(dtype=lora_embeds.dtype, device=lora_embeds.device)
|
| 96 |
+
embeds = torch.concat([special_embeds, lora_embeds], dim=1)
|
| 97 |
+
for encoder_id, encoder in enumerate(self.encoders):
|
| 98 |
+
embeds = encoder(embeds)
|
| 99 |
+
embeds = embeds[:, :self.num_special_embeds]
|
| 100 |
+
embeds = self.final_layer_norm(embeds)
|
| 101 |
+
embeds = self.final_linear(embeds)
|
| 102 |
+
return embeds
|
| 103 |
+
|
| 104 |
+
@staticmethod
|
| 105 |
+
def state_dict_converter():
|
| 106 |
+
return FluxLoRAEncoderStateDictConverter()
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
class FluxLoRAEncoderStateDictConverter:
|
| 110 |
+
def from_civitai(self, state_dict):
|
| 111 |
+
return state_dict
|
diffsynth/models/flux_text_encoder.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from transformers import T5EncoderModel, T5Config
|
| 3 |
+
from .sd_text_encoder import SDTextEncoder
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class FluxTextEncoder2(T5EncoderModel):
|
| 8 |
+
def __init__(self, config):
|
| 9 |
+
super().__init__(config)
|
| 10 |
+
self.eval()
|
| 11 |
+
|
| 12 |
+
def forward(self, input_ids):
|
| 13 |
+
outputs = super().forward(input_ids=input_ids)
|
| 14 |
+
prompt_emb = outputs.last_hidden_state
|
| 15 |
+
return prompt_emb
|
| 16 |
+
|
| 17 |
+
@staticmethod
|
| 18 |
+
def state_dict_converter():
|
| 19 |
+
return FluxTextEncoder2StateDictConverter()
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class FluxTextEncoder2StateDictConverter():
|
| 24 |
+
def __init__(self):
|
| 25 |
+
pass
|
| 26 |
+
|
| 27 |
+
def from_diffusers(self, state_dict):
|
| 28 |
+
state_dict_ = state_dict
|
| 29 |
+
return state_dict_
|
| 30 |
+
|
| 31 |
+
def from_civitai(self, state_dict):
|
| 32 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/flux_vae.py
ADDED
|
@@ -0,0 +1,303 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sd3_vae_encoder import SD3VAEEncoder, SDVAEEncoderStateDictConverter
|
| 2 |
+
from .sd3_vae_decoder import SD3VAEDecoder, SDVAEDecoderStateDictConverter
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class FluxVAEEncoder(SD3VAEEncoder):
|
| 6 |
+
def __init__(self):
|
| 7 |
+
super().__init__()
|
| 8 |
+
self.scaling_factor = 0.3611
|
| 9 |
+
self.shift_factor = 0.1159
|
| 10 |
+
|
| 11 |
+
@staticmethod
|
| 12 |
+
def state_dict_converter():
|
| 13 |
+
return FluxVAEEncoderStateDictConverter()
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class FluxVAEDecoder(SD3VAEDecoder):
|
| 17 |
+
def __init__(self):
|
| 18 |
+
super().__init__()
|
| 19 |
+
self.scaling_factor = 0.3611
|
| 20 |
+
self.shift_factor = 0.1159
|
| 21 |
+
|
| 22 |
+
@staticmethod
|
| 23 |
+
def state_dict_converter():
|
| 24 |
+
return FluxVAEDecoderStateDictConverter()
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class FluxVAEEncoderStateDictConverter(SDVAEEncoderStateDictConverter):
|
| 28 |
+
def __init__(self):
|
| 29 |
+
pass
|
| 30 |
+
|
| 31 |
+
def from_civitai(self, state_dict):
|
| 32 |
+
rename_dict = {
|
| 33 |
+
"encoder.conv_in.bias": "conv_in.bias",
|
| 34 |
+
"encoder.conv_in.weight": "conv_in.weight",
|
| 35 |
+
"encoder.conv_out.bias": "conv_out.bias",
|
| 36 |
+
"encoder.conv_out.weight": "conv_out.weight",
|
| 37 |
+
"encoder.down.0.block.0.conv1.bias": "blocks.0.conv1.bias",
|
| 38 |
+
"encoder.down.0.block.0.conv1.weight": "blocks.0.conv1.weight",
|
| 39 |
+
"encoder.down.0.block.0.conv2.bias": "blocks.0.conv2.bias",
|
| 40 |
+
"encoder.down.0.block.0.conv2.weight": "blocks.0.conv2.weight",
|
| 41 |
+
"encoder.down.0.block.0.norm1.bias": "blocks.0.norm1.bias",
|
| 42 |
+
"encoder.down.0.block.0.norm1.weight": "blocks.0.norm1.weight",
|
| 43 |
+
"encoder.down.0.block.0.norm2.bias": "blocks.0.norm2.bias",
|
| 44 |
+
"encoder.down.0.block.0.norm2.weight": "blocks.0.norm2.weight",
|
| 45 |
+
"encoder.down.0.block.1.conv1.bias": "blocks.1.conv1.bias",
|
| 46 |
+
"encoder.down.0.block.1.conv1.weight": "blocks.1.conv1.weight",
|
| 47 |
+
"encoder.down.0.block.1.conv2.bias": "blocks.1.conv2.bias",
|
| 48 |
+
"encoder.down.0.block.1.conv2.weight": "blocks.1.conv2.weight",
|
| 49 |
+
"encoder.down.0.block.1.norm1.bias": "blocks.1.norm1.bias",
|
| 50 |
+
"encoder.down.0.block.1.norm1.weight": "blocks.1.norm1.weight",
|
| 51 |
+
"encoder.down.0.block.1.norm2.bias": "blocks.1.norm2.bias",
|
| 52 |
+
"encoder.down.0.block.1.norm2.weight": "blocks.1.norm2.weight",
|
| 53 |
+
"encoder.down.0.downsample.conv.bias": "blocks.2.conv.bias",
|
| 54 |
+
"encoder.down.0.downsample.conv.weight": "blocks.2.conv.weight",
|
| 55 |
+
"encoder.down.1.block.0.conv1.bias": "blocks.3.conv1.bias",
|
| 56 |
+
"encoder.down.1.block.0.conv1.weight": "blocks.3.conv1.weight",
|
| 57 |
+
"encoder.down.1.block.0.conv2.bias": "blocks.3.conv2.bias",
|
| 58 |
+
"encoder.down.1.block.0.conv2.weight": "blocks.3.conv2.weight",
|
| 59 |
+
"encoder.down.1.block.0.nin_shortcut.bias": "blocks.3.conv_shortcut.bias",
|
| 60 |
+
"encoder.down.1.block.0.nin_shortcut.weight": "blocks.3.conv_shortcut.weight",
|
| 61 |
+
"encoder.down.1.block.0.norm1.bias": "blocks.3.norm1.bias",
|
| 62 |
+
"encoder.down.1.block.0.norm1.weight": "blocks.3.norm1.weight",
|
| 63 |
+
"encoder.down.1.block.0.norm2.bias": "blocks.3.norm2.bias",
|
| 64 |
+
"encoder.down.1.block.0.norm2.weight": "blocks.3.norm2.weight",
|
| 65 |
+
"encoder.down.1.block.1.conv1.bias": "blocks.4.conv1.bias",
|
| 66 |
+
"encoder.down.1.block.1.conv1.weight": "blocks.4.conv1.weight",
|
| 67 |
+
"encoder.down.1.block.1.conv2.bias": "blocks.4.conv2.bias",
|
| 68 |
+
"encoder.down.1.block.1.conv2.weight": "blocks.4.conv2.weight",
|
| 69 |
+
"encoder.down.1.block.1.norm1.bias": "blocks.4.norm1.bias",
|
| 70 |
+
"encoder.down.1.block.1.norm1.weight": "blocks.4.norm1.weight",
|
| 71 |
+
"encoder.down.1.block.1.norm2.bias": "blocks.4.norm2.bias",
|
| 72 |
+
"encoder.down.1.block.1.norm2.weight": "blocks.4.norm2.weight",
|
| 73 |
+
"encoder.down.1.downsample.conv.bias": "blocks.5.conv.bias",
|
| 74 |
+
"encoder.down.1.downsample.conv.weight": "blocks.5.conv.weight",
|
| 75 |
+
"encoder.down.2.block.0.conv1.bias": "blocks.6.conv1.bias",
|
| 76 |
+
"encoder.down.2.block.0.conv1.weight": "blocks.6.conv1.weight",
|
| 77 |
+
"encoder.down.2.block.0.conv2.bias": "blocks.6.conv2.bias",
|
| 78 |
+
"encoder.down.2.block.0.conv2.weight": "blocks.6.conv2.weight",
|
| 79 |
+
"encoder.down.2.block.0.nin_shortcut.bias": "blocks.6.conv_shortcut.bias",
|
| 80 |
+
"encoder.down.2.block.0.nin_shortcut.weight": "blocks.6.conv_shortcut.weight",
|
| 81 |
+
"encoder.down.2.block.0.norm1.bias": "blocks.6.norm1.bias",
|
| 82 |
+
"encoder.down.2.block.0.norm1.weight": "blocks.6.norm1.weight",
|
| 83 |
+
"encoder.down.2.block.0.norm2.bias": "blocks.6.norm2.bias",
|
| 84 |
+
"encoder.down.2.block.0.norm2.weight": "blocks.6.norm2.weight",
|
| 85 |
+
"encoder.down.2.block.1.conv1.bias": "blocks.7.conv1.bias",
|
| 86 |
+
"encoder.down.2.block.1.conv1.weight": "blocks.7.conv1.weight",
|
| 87 |
+
"encoder.down.2.block.1.conv2.bias": "blocks.7.conv2.bias",
|
| 88 |
+
"encoder.down.2.block.1.conv2.weight": "blocks.7.conv2.weight",
|
| 89 |
+
"encoder.down.2.block.1.norm1.bias": "blocks.7.norm1.bias",
|
| 90 |
+
"encoder.down.2.block.1.norm1.weight": "blocks.7.norm1.weight",
|
| 91 |
+
"encoder.down.2.block.1.norm2.bias": "blocks.7.norm2.bias",
|
| 92 |
+
"encoder.down.2.block.1.norm2.weight": "blocks.7.norm2.weight",
|
| 93 |
+
"encoder.down.2.downsample.conv.bias": "blocks.8.conv.bias",
|
| 94 |
+
"encoder.down.2.downsample.conv.weight": "blocks.8.conv.weight",
|
| 95 |
+
"encoder.down.3.block.0.conv1.bias": "blocks.9.conv1.bias",
|
| 96 |
+
"encoder.down.3.block.0.conv1.weight": "blocks.9.conv1.weight",
|
| 97 |
+
"encoder.down.3.block.0.conv2.bias": "blocks.9.conv2.bias",
|
| 98 |
+
"encoder.down.3.block.0.conv2.weight": "blocks.9.conv2.weight",
|
| 99 |
+
"encoder.down.3.block.0.norm1.bias": "blocks.9.norm1.bias",
|
| 100 |
+
"encoder.down.3.block.0.norm1.weight": "blocks.9.norm1.weight",
|
| 101 |
+
"encoder.down.3.block.0.norm2.bias": "blocks.9.norm2.bias",
|
| 102 |
+
"encoder.down.3.block.0.norm2.weight": "blocks.9.norm2.weight",
|
| 103 |
+
"encoder.down.3.block.1.conv1.bias": "blocks.10.conv1.bias",
|
| 104 |
+
"encoder.down.3.block.1.conv1.weight": "blocks.10.conv1.weight",
|
| 105 |
+
"encoder.down.3.block.1.conv2.bias": "blocks.10.conv2.bias",
|
| 106 |
+
"encoder.down.3.block.1.conv2.weight": "blocks.10.conv2.weight",
|
| 107 |
+
"encoder.down.3.block.1.norm1.bias": "blocks.10.norm1.bias",
|
| 108 |
+
"encoder.down.3.block.1.norm1.weight": "blocks.10.norm1.weight",
|
| 109 |
+
"encoder.down.3.block.1.norm2.bias": "blocks.10.norm2.bias",
|
| 110 |
+
"encoder.down.3.block.1.norm2.weight": "blocks.10.norm2.weight",
|
| 111 |
+
"encoder.mid.attn_1.k.bias": "blocks.12.transformer_blocks.0.to_k.bias",
|
| 112 |
+
"encoder.mid.attn_1.k.weight": "blocks.12.transformer_blocks.0.to_k.weight",
|
| 113 |
+
"encoder.mid.attn_1.norm.bias": "blocks.12.norm.bias",
|
| 114 |
+
"encoder.mid.attn_1.norm.weight": "blocks.12.norm.weight",
|
| 115 |
+
"encoder.mid.attn_1.proj_out.bias": "blocks.12.transformer_blocks.0.to_out.bias",
|
| 116 |
+
"encoder.mid.attn_1.proj_out.weight": "blocks.12.transformer_blocks.0.to_out.weight",
|
| 117 |
+
"encoder.mid.attn_1.q.bias": "blocks.12.transformer_blocks.0.to_q.bias",
|
| 118 |
+
"encoder.mid.attn_1.q.weight": "blocks.12.transformer_blocks.0.to_q.weight",
|
| 119 |
+
"encoder.mid.attn_1.v.bias": "blocks.12.transformer_blocks.0.to_v.bias",
|
| 120 |
+
"encoder.mid.attn_1.v.weight": "blocks.12.transformer_blocks.0.to_v.weight",
|
| 121 |
+
"encoder.mid.block_1.conv1.bias": "blocks.11.conv1.bias",
|
| 122 |
+
"encoder.mid.block_1.conv1.weight": "blocks.11.conv1.weight",
|
| 123 |
+
"encoder.mid.block_1.conv2.bias": "blocks.11.conv2.bias",
|
| 124 |
+
"encoder.mid.block_1.conv2.weight": "blocks.11.conv2.weight",
|
| 125 |
+
"encoder.mid.block_1.norm1.bias": "blocks.11.norm1.bias",
|
| 126 |
+
"encoder.mid.block_1.norm1.weight": "blocks.11.norm1.weight",
|
| 127 |
+
"encoder.mid.block_1.norm2.bias": "blocks.11.norm2.bias",
|
| 128 |
+
"encoder.mid.block_1.norm2.weight": "blocks.11.norm2.weight",
|
| 129 |
+
"encoder.mid.block_2.conv1.bias": "blocks.13.conv1.bias",
|
| 130 |
+
"encoder.mid.block_2.conv1.weight": "blocks.13.conv1.weight",
|
| 131 |
+
"encoder.mid.block_2.conv2.bias": "blocks.13.conv2.bias",
|
| 132 |
+
"encoder.mid.block_2.conv2.weight": "blocks.13.conv2.weight",
|
| 133 |
+
"encoder.mid.block_2.norm1.bias": "blocks.13.norm1.bias",
|
| 134 |
+
"encoder.mid.block_2.norm1.weight": "blocks.13.norm1.weight",
|
| 135 |
+
"encoder.mid.block_2.norm2.bias": "blocks.13.norm2.bias",
|
| 136 |
+
"encoder.mid.block_2.norm2.weight": "blocks.13.norm2.weight",
|
| 137 |
+
"encoder.norm_out.bias": "conv_norm_out.bias",
|
| 138 |
+
"encoder.norm_out.weight": "conv_norm_out.weight",
|
| 139 |
+
}
|
| 140 |
+
state_dict_ = {}
|
| 141 |
+
for name in state_dict:
|
| 142 |
+
if name in rename_dict:
|
| 143 |
+
param = state_dict[name]
|
| 144 |
+
if "transformer_blocks" in rename_dict[name]:
|
| 145 |
+
param = param.squeeze()
|
| 146 |
+
state_dict_[rename_dict[name]] = param
|
| 147 |
+
return state_dict_
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
class FluxVAEDecoderStateDictConverter(SDVAEDecoderStateDictConverter):
|
| 152 |
+
def __init__(self):
|
| 153 |
+
pass
|
| 154 |
+
|
| 155 |
+
def from_civitai(self, state_dict):
|
| 156 |
+
rename_dict = {
|
| 157 |
+
"decoder.conv_in.bias": "conv_in.bias",
|
| 158 |
+
"decoder.conv_in.weight": "conv_in.weight",
|
| 159 |
+
"decoder.conv_out.bias": "conv_out.bias",
|
| 160 |
+
"decoder.conv_out.weight": "conv_out.weight",
|
| 161 |
+
"decoder.mid.attn_1.k.bias": "blocks.1.transformer_blocks.0.to_k.bias",
|
| 162 |
+
"decoder.mid.attn_1.k.weight": "blocks.1.transformer_blocks.0.to_k.weight",
|
| 163 |
+
"decoder.mid.attn_1.norm.bias": "blocks.1.norm.bias",
|
| 164 |
+
"decoder.mid.attn_1.norm.weight": "blocks.1.norm.weight",
|
| 165 |
+
"decoder.mid.attn_1.proj_out.bias": "blocks.1.transformer_blocks.0.to_out.bias",
|
| 166 |
+
"decoder.mid.attn_1.proj_out.weight": "blocks.1.transformer_blocks.0.to_out.weight",
|
| 167 |
+
"decoder.mid.attn_1.q.bias": "blocks.1.transformer_blocks.0.to_q.bias",
|
| 168 |
+
"decoder.mid.attn_1.q.weight": "blocks.1.transformer_blocks.0.to_q.weight",
|
| 169 |
+
"decoder.mid.attn_1.v.bias": "blocks.1.transformer_blocks.0.to_v.bias",
|
| 170 |
+
"decoder.mid.attn_1.v.weight": "blocks.1.transformer_blocks.0.to_v.weight",
|
| 171 |
+
"decoder.mid.block_1.conv1.bias": "blocks.0.conv1.bias",
|
| 172 |
+
"decoder.mid.block_1.conv1.weight": "blocks.0.conv1.weight",
|
| 173 |
+
"decoder.mid.block_1.conv2.bias": "blocks.0.conv2.bias",
|
| 174 |
+
"decoder.mid.block_1.conv2.weight": "blocks.0.conv2.weight",
|
| 175 |
+
"decoder.mid.block_1.norm1.bias": "blocks.0.norm1.bias",
|
| 176 |
+
"decoder.mid.block_1.norm1.weight": "blocks.0.norm1.weight",
|
| 177 |
+
"decoder.mid.block_1.norm2.bias": "blocks.0.norm2.bias",
|
| 178 |
+
"decoder.mid.block_1.norm2.weight": "blocks.0.norm2.weight",
|
| 179 |
+
"decoder.mid.block_2.conv1.bias": "blocks.2.conv1.bias",
|
| 180 |
+
"decoder.mid.block_2.conv1.weight": "blocks.2.conv1.weight",
|
| 181 |
+
"decoder.mid.block_2.conv2.bias": "blocks.2.conv2.bias",
|
| 182 |
+
"decoder.mid.block_2.conv2.weight": "blocks.2.conv2.weight",
|
| 183 |
+
"decoder.mid.block_2.norm1.bias": "blocks.2.norm1.bias",
|
| 184 |
+
"decoder.mid.block_2.norm1.weight": "blocks.2.norm1.weight",
|
| 185 |
+
"decoder.mid.block_2.norm2.bias": "blocks.2.norm2.bias",
|
| 186 |
+
"decoder.mid.block_2.norm2.weight": "blocks.2.norm2.weight",
|
| 187 |
+
"decoder.norm_out.bias": "conv_norm_out.bias",
|
| 188 |
+
"decoder.norm_out.weight": "conv_norm_out.weight",
|
| 189 |
+
"decoder.up.0.block.0.conv1.bias": "blocks.15.conv1.bias",
|
| 190 |
+
"decoder.up.0.block.0.conv1.weight": "blocks.15.conv1.weight",
|
| 191 |
+
"decoder.up.0.block.0.conv2.bias": "blocks.15.conv2.bias",
|
| 192 |
+
"decoder.up.0.block.0.conv2.weight": "blocks.15.conv2.weight",
|
| 193 |
+
"decoder.up.0.block.0.nin_shortcut.bias": "blocks.15.conv_shortcut.bias",
|
| 194 |
+
"decoder.up.0.block.0.nin_shortcut.weight": "blocks.15.conv_shortcut.weight",
|
| 195 |
+
"decoder.up.0.block.0.norm1.bias": "blocks.15.norm1.bias",
|
| 196 |
+
"decoder.up.0.block.0.norm1.weight": "blocks.15.norm1.weight",
|
| 197 |
+
"decoder.up.0.block.0.norm2.bias": "blocks.15.norm2.bias",
|
| 198 |
+
"decoder.up.0.block.0.norm2.weight": "blocks.15.norm2.weight",
|
| 199 |
+
"decoder.up.0.block.1.conv1.bias": "blocks.16.conv1.bias",
|
| 200 |
+
"decoder.up.0.block.1.conv1.weight": "blocks.16.conv1.weight",
|
| 201 |
+
"decoder.up.0.block.1.conv2.bias": "blocks.16.conv2.bias",
|
| 202 |
+
"decoder.up.0.block.1.conv2.weight": "blocks.16.conv2.weight",
|
| 203 |
+
"decoder.up.0.block.1.norm1.bias": "blocks.16.norm1.bias",
|
| 204 |
+
"decoder.up.0.block.1.norm1.weight": "blocks.16.norm1.weight",
|
| 205 |
+
"decoder.up.0.block.1.norm2.bias": "blocks.16.norm2.bias",
|
| 206 |
+
"decoder.up.0.block.1.norm2.weight": "blocks.16.norm2.weight",
|
| 207 |
+
"decoder.up.0.block.2.conv1.bias": "blocks.17.conv1.bias",
|
| 208 |
+
"decoder.up.0.block.2.conv1.weight": "blocks.17.conv1.weight",
|
| 209 |
+
"decoder.up.0.block.2.conv2.bias": "blocks.17.conv2.bias",
|
| 210 |
+
"decoder.up.0.block.2.conv2.weight": "blocks.17.conv2.weight",
|
| 211 |
+
"decoder.up.0.block.2.norm1.bias": "blocks.17.norm1.bias",
|
| 212 |
+
"decoder.up.0.block.2.norm1.weight": "blocks.17.norm1.weight",
|
| 213 |
+
"decoder.up.0.block.2.norm2.bias": "blocks.17.norm2.bias",
|
| 214 |
+
"decoder.up.0.block.2.norm2.weight": "blocks.17.norm2.weight",
|
| 215 |
+
"decoder.up.1.block.0.conv1.bias": "blocks.11.conv1.bias",
|
| 216 |
+
"decoder.up.1.block.0.conv1.weight": "blocks.11.conv1.weight",
|
| 217 |
+
"decoder.up.1.block.0.conv2.bias": "blocks.11.conv2.bias",
|
| 218 |
+
"decoder.up.1.block.0.conv2.weight": "blocks.11.conv2.weight",
|
| 219 |
+
"decoder.up.1.block.0.nin_shortcut.bias": "blocks.11.conv_shortcut.bias",
|
| 220 |
+
"decoder.up.1.block.0.nin_shortcut.weight": "blocks.11.conv_shortcut.weight",
|
| 221 |
+
"decoder.up.1.block.0.norm1.bias": "blocks.11.norm1.bias",
|
| 222 |
+
"decoder.up.1.block.0.norm1.weight": "blocks.11.norm1.weight",
|
| 223 |
+
"decoder.up.1.block.0.norm2.bias": "blocks.11.norm2.bias",
|
| 224 |
+
"decoder.up.1.block.0.norm2.weight": "blocks.11.norm2.weight",
|
| 225 |
+
"decoder.up.1.block.1.conv1.bias": "blocks.12.conv1.bias",
|
| 226 |
+
"decoder.up.1.block.1.conv1.weight": "blocks.12.conv1.weight",
|
| 227 |
+
"decoder.up.1.block.1.conv2.bias": "blocks.12.conv2.bias",
|
| 228 |
+
"decoder.up.1.block.1.conv2.weight": "blocks.12.conv2.weight",
|
| 229 |
+
"decoder.up.1.block.1.norm1.bias": "blocks.12.norm1.bias",
|
| 230 |
+
"decoder.up.1.block.1.norm1.weight": "blocks.12.norm1.weight",
|
| 231 |
+
"decoder.up.1.block.1.norm2.bias": "blocks.12.norm2.bias",
|
| 232 |
+
"decoder.up.1.block.1.norm2.weight": "blocks.12.norm2.weight",
|
| 233 |
+
"decoder.up.1.block.2.conv1.bias": "blocks.13.conv1.bias",
|
| 234 |
+
"decoder.up.1.block.2.conv1.weight": "blocks.13.conv1.weight",
|
| 235 |
+
"decoder.up.1.block.2.conv2.bias": "blocks.13.conv2.bias",
|
| 236 |
+
"decoder.up.1.block.2.conv2.weight": "blocks.13.conv2.weight",
|
| 237 |
+
"decoder.up.1.block.2.norm1.bias": "blocks.13.norm1.bias",
|
| 238 |
+
"decoder.up.1.block.2.norm1.weight": "blocks.13.norm1.weight",
|
| 239 |
+
"decoder.up.1.block.2.norm2.bias": "blocks.13.norm2.bias",
|
| 240 |
+
"decoder.up.1.block.2.norm2.weight": "blocks.13.norm2.weight",
|
| 241 |
+
"decoder.up.1.upsample.conv.bias": "blocks.14.conv.bias",
|
| 242 |
+
"decoder.up.1.upsample.conv.weight": "blocks.14.conv.weight",
|
| 243 |
+
"decoder.up.2.block.0.conv1.bias": "blocks.7.conv1.bias",
|
| 244 |
+
"decoder.up.2.block.0.conv1.weight": "blocks.7.conv1.weight",
|
| 245 |
+
"decoder.up.2.block.0.conv2.bias": "blocks.7.conv2.bias",
|
| 246 |
+
"decoder.up.2.block.0.conv2.weight": "blocks.7.conv2.weight",
|
| 247 |
+
"decoder.up.2.block.0.norm1.bias": "blocks.7.norm1.bias",
|
| 248 |
+
"decoder.up.2.block.0.norm1.weight": "blocks.7.norm1.weight",
|
| 249 |
+
"decoder.up.2.block.0.norm2.bias": "blocks.7.norm2.bias",
|
| 250 |
+
"decoder.up.2.block.0.norm2.weight": "blocks.7.norm2.weight",
|
| 251 |
+
"decoder.up.2.block.1.conv1.bias": "blocks.8.conv1.bias",
|
| 252 |
+
"decoder.up.2.block.1.conv1.weight": "blocks.8.conv1.weight",
|
| 253 |
+
"decoder.up.2.block.1.conv2.bias": "blocks.8.conv2.bias",
|
| 254 |
+
"decoder.up.2.block.1.conv2.weight": "blocks.8.conv2.weight",
|
| 255 |
+
"decoder.up.2.block.1.norm1.bias": "blocks.8.norm1.bias",
|
| 256 |
+
"decoder.up.2.block.1.norm1.weight": "blocks.8.norm1.weight",
|
| 257 |
+
"decoder.up.2.block.1.norm2.bias": "blocks.8.norm2.bias",
|
| 258 |
+
"decoder.up.2.block.1.norm2.weight": "blocks.8.norm2.weight",
|
| 259 |
+
"decoder.up.2.block.2.conv1.bias": "blocks.9.conv1.bias",
|
| 260 |
+
"decoder.up.2.block.2.conv1.weight": "blocks.9.conv1.weight",
|
| 261 |
+
"decoder.up.2.block.2.conv2.bias": "blocks.9.conv2.bias",
|
| 262 |
+
"decoder.up.2.block.2.conv2.weight": "blocks.9.conv2.weight",
|
| 263 |
+
"decoder.up.2.block.2.norm1.bias": "blocks.9.norm1.bias",
|
| 264 |
+
"decoder.up.2.block.2.norm1.weight": "blocks.9.norm1.weight",
|
| 265 |
+
"decoder.up.2.block.2.norm2.bias": "blocks.9.norm2.bias",
|
| 266 |
+
"decoder.up.2.block.2.norm2.weight": "blocks.9.norm2.weight",
|
| 267 |
+
"decoder.up.2.upsample.conv.bias": "blocks.10.conv.bias",
|
| 268 |
+
"decoder.up.2.upsample.conv.weight": "blocks.10.conv.weight",
|
| 269 |
+
"decoder.up.3.block.0.conv1.bias": "blocks.3.conv1.bias",
|
| 270 |
+
"decoder.up.3.block.0.conv1.weight": "blocks.3.conv1.weight",
|
| 271 |
+
"decoder.up.3.block.0.conv2.bias": "blocks.3.conv2.bias",
|
| 272 |
+
"decoder.up.3.block.0.conv2.weight": "blocks.3.conv2.weight",
|
| 273 |
+
"decoder.up.3.block.0.norm1.bias": "blocks.3.norm1.bias",
|
| 274 |
+
"decoder.up.3.block.0.norm1.weight": "blocks.3.norm1.weight",
|
| 275 |
+
"decoder.up.3.block.0.norm2.bias": "blocks.3.norm2.bias",
|
| 276 |
+
"decoder.up.3.block.0.norm2.weight": "blocks.3.norm2.weight",
|
| 277 |
+
"decoder.up.3.block.1.conv1.bias": "blocks.4.conv1.bias",
|
| 278 |
+
"decoder.up.3.block.1.conv1.weight": "blocks.4.conv1.weight",
|
| 279 |
+
"decoder.up.3.block.1.conv2.bias": "blocks.4.conv2.bias",
|
| 280 |
+
"decoder.up.3.block.1.conv2.weight": "blocks.4.conv2.weight",
|
| 281 |
+
"decoder.up.3.block.1.norm1.bias": "blocks.4.norm1.bias",
|
| 282 |
+
"decoder.up.3.block.1.norm1.weight": "blocks.4.norm1.weight",
|
| 283 |
+
"decoder.up.3.block.1.norm2.bias": "blocks.4.norm2.bias",
|
| 284 |
+
"decoder.up.3.block.1.norm2.weight": "blocks.4.norm2.weight",
|
| 285 |
+
"decoder.up.3.block.2.conv1.bias": "blocks.5.conv1.bias",
|
| 286 |
+
"decoder.up.3.block.2.conv1.weight": "blocks.5.conv1.weight",
|
| 287 |
+
"decoder.up.3.block.2.conv2.bias": "blocks.5.conv2.bias",
|
| 288 |
+
"decoder.up.3.block.2.conv2.weight": "blocks.5.conv2.weight",
|
| 289 |
+
"decoder.up.3.block.2.norm1.bias": "blocks.5.norm1.bias",
|
| 290 |
+
"decoder.up.3.block.2.norm1.weight": "blocks.5.norm1.weight",
|
| 291 |
+
"decoder.up.3.block.2.norm2.bias": "blocks.5.norm2.bias",
|
| 292 |
+
"decoder.up.3.block.2.norm2.weight": "blocks.5.norm2.weight",
|
| 293 |
+
"decoder.up.3.upsample.conv.bias": "blocks.6.conv.bias",
|
| 294 |
+
"decoder.up.3.upsample.conv.weight": "blocks.6.conv.weight",
|
| 295 |
+
}
|
| 296 |
+
state_dict_ = {}
|
| 297 |
+
for name in state_dict:
|
| 298 |
+
if name in rename_dict:
|
| 299 |
+
param = state_dict[name]
|
| 300 |
+
if "transformer_blocks" in rename_dict[name]:
|
| 301 |
+
param = param.squeeze()
|
| 302 |
+
state_dict_[rename_dict[name]] = param
|
| 303 |
+
return state_dict_
|
diffsynth/models/flux_value_control.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from diffsynth.models.svd_unet import TemporalTimesteps
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class MultiValueEncoder(torch.nn.Module):
|
| 6 |
+
def __init__(self, encoders=()):
|
| 7 |
+
super().__init__()
|
| 8 |
+
self.encoders = torch.nn.ModuleList(encoders)
|
| 9 |
+
|
| 10 |
+
def __call__(self, values, dtype):
|
| 11 |
+
emb = []
|
| 12 |
+
for encoder, value in zip(self.encoders, values):
|
| 13 |
+
if value is not None:
|
| 14 |
+
value = value.unsqueeze(0)
|
| 15 |
+
emb.append(encoder(value, dtype))
|
| 16 |
+
emb = torch.concat(emb, dim=0)
|
| 17 |
+
return emb
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class SingleValueEncoder(torch.nn.Module):
|
| 21 |
+
def __init__(self, dim_in=256, dim_out=4096, prefer_len=32, computation_device=None):
|
| 22 |
+
super().__init__()
|
| 23 |
+
self.prefer_len = prefer_len
|
| 24 |
+
self.prefer_proj = TemporalTimesteps(num_channels=dim_in, flip_sin_to_cos=True, downscale_freq_shift=0, computation_device=computation_device)
|
| 25 |
+
self.prefer_value_embedder = torch.nn.Sequential(
|
| 26 |
+
torch.nn.Linear(dim_in, dim_out), torch.nn.SiLU(), torch.nn.Linear(dim_out, dim_out)
|
| 27 |
+
)
|
| 28 |
+
self.positional_embedding = torch.nn.Parameter(
|
| 29 |
+
torch.randn(self.prefer_len, dim_out)
|
| 30 |
+
)
|
| 31 |
+
self._initialize_weights()
|
| 32 |
+
|
| 33 |
+
def _initialize_weights(self):
|
| 34 |
+
last_linear = self.prefer_value_embedder[-1]
|
| 35 |
+
torch.nn.init.zeros_(last_linear.weight)
|
| 36 |
+
torch.nn.init.zeros_(last_linear.bias)
|
| 37 |
+
|
| 38 |
+
def forward(self, value, dtype):
|
| 39 |
+
value = value * 1000
|
| 40 |
+
emb = self.prefer_proj(value).to(dtype)
|
| 41 |
+
emb = self.prefer_value_embedder(emb).squeeze(0)
|
| 42 |
+
base_embeddings = emb.expand(self.prefer_len, -1)
|
| 43 |
+
positional_embedding = self.positional_embedding.to(dtype=base_embeddings.dtype, device=base_embeddings.device)
|
| 44 |
+
learned_embeddings = base_embeddings + positional_embedding
|
| 45 |
+
return learned_embeddings
|
| 46 |
+
|
| 47 |
+
@staticmethod
|
| 48 |
+
def state_dict_converter():
|
| 49 |
+
return SingleValueEncoderStateDictConverter()
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class SingleValueEncoderStateDictConverter:
|
| 53 |
+
def __init__(self):
|
| 54 |
+
pass
|
| 55 |
+
|
| 56 |
+
def from_diffusers(self, state_dict):
|
| 57 |
+
return state_dict
|
| 58 |
+
|
| 59 |
+
def from_civitai(self, state_dict):
|
| 60 |
+
return state_dict
|
diffsynth/models/hunyuan_dit.py
ADDED
|
@@ -0,0 +1,451 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .attention import Attention
|
| 2 |
+
from einops import repeat, rearrange
|
| 3 |
+
import math
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class HunyuanDiTRotaryEmbedding(torch.nn.Module):
|
| 8 |
+
|
| 9 |
+
def __init__(self, q_norm_shape=88, k_norm_shape=88, rotary_emb_on_k=True):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.q_norm = torch.nn.LayerNorm((q_norm_shape,), elementwise_affine=True, eps=1e-06)
|
| 12 |
+
self.k_norm = torch.nn.LayerNorm((k_norm_shape,), elementwise_affine=True, eps=1e-06)
|
| 13 |
+
self.rotary_emb_on_k = rotary_emb_on_k
|
| 14 |
+
self.k_cache, self.v_cache = [], []
|
| 15 |
+
|
| 16 |
+
def reshape_for_broadcast(self, freqs_cis, x):
|
| 17 |
+
ndim = x.ndim
|
| 18 |
+
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
| 19 |
+
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
|
| 20 |
+
|
| 21 |
+
def rotate_half(self, x):
|
| 22 |
+
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
|
| 23 |
+
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
| 24 |
+
|
| 25 |
+
def apply_rotary_emb(self, xq, xk, freqs_cis):
|
| 26 |
+
xk_out = None
|
| 27 |
+
cos, sin = self.reshape_for_broadcast(freqs_cis, xq)
|
| 28 |
+
cos, sin = cos.to(xq.device), sin.to(xq.device)
|
| 29 |
+
xq_out = (xq.float() * cos + self.rotate_half(xq.float()) * sin).type_as(xq)
|
| 30 |
+
if xk is not None:
|
| 31 |
+
xk_out = (xk.float() * cos + self.rotate_half(xk.float()) * sin).type_as(xk)
|
| 32 |
+
return xq_out, xk_out
|
| 33 |
+
|
| 34 |
+
def forward(self, q, k, v, freqs_cis_img, to_cache=False):
|
| 35 |
+
# norm
|
| 36 |
+
q = self.q_norm(q)
|
| 37 |
+
k = self.k_norm(k)
|
| 38 |
+
|
| 39 |
+
# RoPE
|
| 40 |
+
if self.rotary_emb_on_k:
|
| 41 |
+
q, k = self.apply_rotary_emb(q, k, freqs_cis_img)
|
| 42 |
+
else:
|
| 43 |
+
q, _ = self.apply_rotary_emb(q, None, freqs_cis_img)
|
| 44 |
+
|
| 45 |
+
if to_cache:
|
| 46 |
+
self.k_cache.append(k)
|
| 47 |
+
self.v_cache.append(v)
|
| 48 |
+
elif len(self.k_cache) > 0 and len(self.v_cache) > 0:
|
| 49 |
+
k = torch.concat([k] + self.k_cache, dim=2)
|
| 50 |
+
v = torch.concat([v] + self.v_cache, dim=2)
|
| 51 |
+
self.k_cache, self.v_cache = [], []
|
| 52 |
+
return q, k, v
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class FP32_Layernorm(torch.nn.LayerNorm):
|
| 56 |
+
def forward(self, inputs):
|
| 57 |
+
origin_dtype = inputs.dtype
|
| 58 |
+
return torch.nn.functional.layer_norm(inputs.float(), self.normalized_shape, self.weight.float(), self.bias.float(), self.eps).to(origin_dtype)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
class FP32_SiLU(torch.nn.SiLU):
|
| 62 |
+
def forward(self, inputs):
|
| 63 |
+
origin_dtype = inputs.dtype
|
| 64 |
+
return torch.nn.functional.silu(inputs.float(), inplace=False).to(origin_dtype)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class HunyuanDiTFinalLayer(torch.nn.Module):
|
| 68 |
+
def __init__(self, final_hidden_size=1408, condition_dim=1408, patch_size=2, out_channels=8):
|
| 69 |
+
super().__init__()
|
| 70 |
+
self.norm_final = torch.nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6)
|
| 71 |
+
self.linear = torch.nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True)
|
| 72 |
+
self.adaLN_modulation = torch.nn.Sequential(
|
| 73 |
+
FP32_SiLU(),
|
| 74 |
+
torch.nn.Linear(condition_dim, 2 * final_hidden_size, bias=True)
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
def modulate(self, x, shift, scale):
|
| 78 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 79 |
+
|
| 80 |
+
def forward(self, hidden_states, condition_emb):
|
| 81 |
+
shift, scale = self.adaLN_modulation(condition_emb).chunk(2, dim=1)
|
| 82 |
+
hidden_states = self.modulate(self.norm_final(hidden_states), shift, scale)
|
| 83 |
+
hidden_states = self.linear(hidden_states)
|
| 84 |
+
return hidden_states
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
class HunyuanDiTBlock(torch.nn.Module):
|
| 88 |
+
|
| 89 |
+
def __init__(
|
| 90 |
+
self,
|
| 91 |
+
hidden_dim=1408,
|
| 92 |
+
condition_dim=1408,
|
| 93 |
+
num_heads=16,
|
| 94 |
+
mlp_ratio=4.3637,
|
| 95 |
+
text_dim=1024,
|
| 96 |
+
skip_connection=False
|
| 97 |
+
):
|
| 98 |
+
super().__init__()
|
| 99 |
+
self.norm1 = FP32_Layernorm((hidden_dim,), eps=1e-6, elementwise_affine=True)
|
| 100 |
+
self.rota1 = HunyuanDiTRotaryEmbedding(hidden_dim//num_heads, hidden_dim//num_heads)
|
| 101 |
+
self.attn1 = Attention(hidden_dim, num_heads, hidden_dim//num_heads, bias_q=True, bias_kv=True, bias_out=True)
|
| 102 |
+
self.norm2 = FP32_Layernorm((hidden_dim,), eps=1e-6, elementwise_affine=True)
|
| 103 |
+
self.rota2 = HunyuanDiTRotaryEmbedding(hidden_dim//num_heads, hidden_dim//num_heads, rotary_emb_on_k=False)
|
| 104 |
+
self.attn2 = Attention(hidden_dim, num_heads, hidden_dim//num_heads, kv_dim=text_dim, bias_q=True, bias_kv=True, bias_out=True)
|
| 105 |
+
self.norm3 = FP32_Layernorm((hidden_dim,), eps=1e-6, elementwise_affine=True)
|
| 106 |
+
self.modulation = torch.nn.Sequential(FP32_SiLU(), torch.nn.Linear(condition_dim, hidden_dim, bias=True))
|
| 107 |
+
self.mlp = torch.nn.Sequential(
|
| 108 |
+
torch.nn.Linear(hidden_dim, int(hidden_dim*mlp_ratio), bias=True),
|
| 109 |
+
torch.nn.GELU(approximate="tanh"),
|
| 110 |
+
torch.nn.Linear(int(hidden_dim*mlp_ratio), hidden_dim, bias=True)
|
| 111 |
+
)
|
| 112 |
+
if skip_connection:
|
| 113 |
+
self.skip_norm = FP32_Layernorm((hidden_dim * 2,), eps=1e-6, elementwise_affine=True)
|
| 114 |
+
self.skip_linear = torch.nn.Linear(hidden_dim * 2, hidden_dim, bias=True)
|
| 115 |
+
else:
|
| 116 |
+
self.skip_norm, self.skip_linear = None, None
|
| 117 |
+
|
| 118 |
+
def forward(self, hidden_states, condition_emb, text_emb, freq_cis_img, residual=None, to_cache=False):
|
| 119 |
+
# Long Skip Connection
|
| 120 |
+
if self.skip_norm is not None and self.skip_linear is not None:
|
| 121 |
+
hidden_states = torch.cat([hidden_states, residual], dim=-1)
|
| 122 |
+
hidden_states = self.skip_norm(hidden_states)
|
| 123 |
+
hidden_states = self.skip_linear(hidden_states)
|
| 124 |
+
|
| 125 |
+
# Self-Attention
|
| 126 |
+
shift_msa = self.modulation(condition_emb).unsqueeze(dim=1)
|
| 127 |
+
attn_input = self.norm1(hidden_states) + shift_msa
|
| 128 |
+
hidden_states = hidden_states + self.attn1(attn_input, qkv_preprocessor=lambda q, k, v: self.rota1(q, k, v, freq_cis_img, to_cache=to_cache))
|
| 129 |
+
|
| 130 |
+
# Cross-Attention
|
| 131 |
+
attn_input = self.norm3(hidden_states)
|
| 132 |
+
hidden_states = hidden_states + self.attn2(attn_input, text_emb, qkv_preprocessor=lambda q, k, v: self.rota2(q, k, v, freq_cis_img))
|
| 133 |
+
|
| 134 |
+
# FFN Layer
|
| 135 |
+
mlp_input = self.norm2(hidden_states)
|
| 136 |
+
hidden_states = hidden_states + self.mlp(mlp_input)
|
| 137 |
+
return hidden_states
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
class AttentionPool(torch.nn.Module):
|
| 141 |
+
def __init__(self, spacial_dim, embed_dim, num_heads, output_dim = None):
|
| 142 |
+
super().__init__()
|
| 143 |
+
self.positional_embedding = torch.nn.Parameter(torch.randn(spacial_dim + 1, embed_dim) / embed_dim ** 0.5)
|
| 144 |
+
self.k_proj = torch.nn.Linear(embed_dim, embed_dim)
|
| 145 |
+
self.q_proj = torch.nn.Linear(embed_dim, embed_dim)
|
| 146 |
+
self.v_proj = torch.nn.Linear(embed_dim, embed_dim)
|
| 147 |
+
self.c_proj = torch.nn.Linear(embed_dim, output_dim or embed_dim)
|
| 148 |
+
self.num_heads = num_heads
|
| 149 |
+
|
| 150 |
+
def forward(self, x):
|
| 151 |
+
x = x.permute(1, 0, 2) # NLC -> LNC
|
| 152 |
+
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (L+1)NC
|
| 153 |
+
x = x + self.positional_embedding[:, None, :].to(x.dtype) # (L+1)NC
|
| 154 |
+
x, _ = torch.nn.functional.multi_head_attention_forward(
|
| 155 |
+
query=x[:1], key=x, value=x,
|
| 156 |
+
embed_dim_to_check=x.shape[-1],
|
| 157 |
+
num_heads=self.num_heads,
|
| 158 |
+
q_proj_weight=self.q_proj.weight,
|
| 159 |
+
k_proj_weight=self.k_proj.weight,
|
| 160 |
+
v_proj_weight=self.v_proj.weight,
|
| 161 |
+
in_proj_weight=None,
|
| 162 |
+
in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
|
| 163 |
+
bias_k=None,
|
| 164 |
+
bias_v=None,
|
| 165 |
+
add_zero_attn=False,
|
| 166 |
+
dropout_p=0,
|
| 167 |
+
out_proj_weight=self.c_proj.weight,
|
| 168 |
+
out_proj_bias=self.c_proj.bias,
|
| 169 |
+
use_separate_proj_weight=True,
|
| 170 |
+
training=self.training,
|
| 171 |
+
need_weights=False
|
| 172 |
+
)
|
| 173 |
+
return x.squeeze(0)
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
class PatchEmbed(torch.nn.Module):
|
| 177 |
+
def __init__(
|
| 178 |
+
self,
|
| 179 |
+
patch_size=(2, 2),
|
| 180 |
+
in_chans=4,
|
| 181 |
+
embed_dim=1408,
|
| 182 |
+
bias=True,
|
| 183 |
+
):
|
| 184 |
+
super().__init__()
|
| 185 |
+
self.proj = torch.nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
|
| 186 |
+
|
| 187 |
+
def forward(self, x):
|
| 188 |
+
x = self.proj(x)
|
| 189 |
+
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
| 190 |
+
return x
|
| 191 |
+
|
| 192 |
+
|
| 193 |
+
def timestep_embedding(t, dim, max_period=10000, repeat_only=False):
|
| 194 |
+
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
| 195 |
+
if not repeat_only:
|
| 196 |
+
half = dim // 2
|
| 197 |
+
freqs = torch.exp(
|
| 198 |
+
-math.log(max_period)
|
| 199 |
+
* torch.arange(start=0, end=half, dtype=torch.float32)
|
| 200 |
+
/ half
|
| 201 |
+
).to(device=t.device) # size: [dim/2], 一个指数衰减的曲线
|
| 202 |
+
args = t[:, None].float() * freqs[None]
|
| 203 |
+
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
| 204 |
+
if dim % 2:
|
| 205 |
+
embedding = torch.cat(
|
| 206 |
+
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
|
| 207 |
+
)
|
| 208 |
+
else:
|
| 209 |
+
embedding = repeat(t, "b -> b d", d=dim)
|
| 210 |
+
return embedding
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
class TimestepEmbedder(torch.nn.Module):
|
| 214 |
+
def __init__(self, hidden_size=1408, frequency_embedding_size=256):
|
| 215 |
+
super().__init__()
|
| 216 |
+
self.mlp = torch.nn.Sequential(
|
| 217 |
+
torch.nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
| 218 |
+
torch.nn.SiLU(),
|
| 219 |
+
torch.nn.Linear(hidden_size, hidden_size, bias=True),
|
| 220 |
+
)
|
| 221 |
+
self.frequency_embedding_size = frequency_embedding_size
|
| 222 |
+
|
| 223 |
+
def forward(self, t):
|
| 224 |
+
t_freq = timestep_embedding(t, self.frequency_embedding_size).type(self.mlp[0].weight.dtype)
|
| 225 |
+
t_emb = self.mlp(t_freq)
|
| 226 |
+
return t_emb
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
class HunyuanDiT(torch.nn.Module):
|
| 230 |
+
def __init__(self, num_layers_down=21, num_layers_up=19, in_channels=4, out_channels=8, hidden_dim=1408, text_dim=1024, t5_dim=2048, text_length=77, t5_length=256):
|
| 231 |
+
super().__init__()
|
| 232 |
+
|
| 233 |
+
# Embedders
|
| 234 |
+
self.text_emb_padding = torch.nn.Parameter(torch.randn(text_length + t5_length, text_dim, dtype=torch.float32))
|
| 235 |
+
self.t5_embedder = torch.nn.Sequential(
|
| 236 |
+
torch.nn.Linear(t5_dim, t5_dim * 4, bias=True),
|
| 237 |
+
FP32_SiLU(),
|
| 238 |
+
torch.nn.Linear(t5_dim * 4, text_dim, bias=True),
|
| 239 |
+
)
|
| 240 |
+
self.t5_pooler = AttentionPool(t5_length, t5_dim, num_heads=8, output_dim=1024)
|
| 241 |
+
self.style_embedder = torch.nn.Parameter(torch.randn(hidden_dim))
|
| 242 |
+
self.patch_embedder = PatchEmbed(in_chans=in_channels)
|
| 243 |
+
self.timestep_embedder = TimestepEmbedder()
|
| 244 |
+
self.extra_embedder = torch.nn.Sequential(
|
| 245 |
+
torch.nn.Linear(256 * 6 + 1024 + hidden_dim, hidden_dim * 4),
|
| 246 |
+
FP32_SiLU(),
|
| 247 |
+
torch.nn.Linear(hidden_dim * 4, hidden_dim),
|
| 248 |
+
)
|
| 249 |
+
|
| 250 |
+
# Transformer blocks
|
| 251 |
+
self.num_layers_down = num_layers_down
|
| 252 |
+
self.num_layers_up = num_layers_up
|
| 253 |
+
self.blocks = torch.nn.ModuleList(
|
| 254 |
+
[HunyuanDiTBlock(skip_connection=False) for _ in range(num_layers_down)] + \
|
| 255 |
+
[HunyuanDiTBlock(skip_connection=True) for _ in range(num_layers_up)]
|
| 256 |
+
)
|
| 257 |
+
|
| 258 |
+
# Output layers
|
| 259 |
+
self.final_layer = HunyuanDiTFinalLayer()
|
| 260 |
+
self.out_channels = out_channels
|
| 261 |
+
|
| 262 |
+
def prepare_text_emb(self, text_emb, text_emb_t5, text_emb_mask, text_emb_mask_t5):
|
| 263 |
+
text_emb_mask = text_emb_mask.bool()
|
| 264 |
+
text_emb_mask_t5 = text_emb_mask_t5.bool()
|
| 265 |
+
text_emb_t5 = self.t5_embedder(text_emb_t5)
|
| 266 |
+
text_emb = torch.cat([text_emb, text_emb_t5], dim=1)
|
| 267 |
+
text_emb_mask = torch.cat([text_emb_mask, text_emb_mask_t5], dim=-1)
|
| 268 |
+
text_emb = torch.where(text_emb_mask.unsqueeze(2), text_emb, self.text_emb_padding.to(text_emb))
|
| 269 |
+
return text_emb
|
| 270 |
+
|
| 271 |
+
def prepare_extra_emb(self, text_emb_t5, timestep, size_emb, dtype, batch_size):
|
| 272 |
+
# Text embedding
|
| 273 |
+
pooled_text_emb_t5 = self.t5_pooler(text_emb_t5)
|
| 274 |
+
|
| 275 |
+
# Timestep embedding
|
| 276 |
+
timestep_emb = self.timestep_embedder(timestep)
|
| 277 |
+
|
| 278 |
+
# Size embedding
|
| 279 |
+
size_emb = timestep_embedding(size_emb.view(-1), 256).to(dtype)
|
| 280 |
+
size_emb = size_emb.view(-1, 6 * 256)
|
| 281 |
+
|
| 282 |
+
# Style embedding
|
| 283 |
+
style_emb = repeat(self.style_embedder, "D -> B D", B=batch_size)
|
| 284 |
+
|
| 285 |
+
# Concatenate all extra vectors
|
| 286 |
+
extra_emb = torch.cat([pooled_text_emb_t5, size_emb, style_emb], dim=1)
|
| 287 |
+
condition_emb = timestep_emb + self.extra_embedder(extra_emb)
|
| 288 |
+
|
| 289 |
+
return condition_emb
|
| 290 |
+
|
| 291 |
+
def unpatchify(self, x, h, w):
|
| 292 |
+
return rearrange(x, "B (H W) (P Q C) -> B C (H P) (W Q)", H=h, W=w, P=2, Q=2)
|
| 293 |
+
|
| 294 |
+
def build_mask(self, data, is_bound):
|
| 295 |
+
_, _, H, W = data.shape
|
| 296 |
+
h = repeat(torch.arange(H), "H -> H W", H=H, W=W)
|
| 297 |
+
w = repeat(torch.arange(W), "W -> H W", H=H, W=W)
|
| 298 |
+
border_width = (H + W) // 4
|
| 299 |
+
pad = torch.ones_like(h) * border_width
|
| 300 |
+
mask = torch.stack([
|
| 301 |
+
pad if is_bound[0] else h + 1,
|
| 302 |
+
pad if is_bound[1] else H - h,
|
| 303 |
+
pad if is_bound[2] else w + 1,
|
| 304 |
+
pad if is_bound[3] else W - w
|
| 305 |
+
]).min(dim=0).values
|
| 306 |
+
mask = mask.clip(1, border_width)
|
| 307 |
+
mask = (mask / border_width).to(dtype=data.dtype, device=data.device)
|
| 308 |
+
mask = rearrange(mask, "H W -> 1 H W")
|
| 309 |
+
return mask
|
| 310 |
+
|
| 311 |
+
def tiled_block_forward(self, block, hidden_states, condition_emb, text_emb, freq_cis_img, residual, torch_dtype, data_device, computation_device, tile_size, tile_stride):
|
| 312 |
+
B, C, H, W = hidden_states.shape
|
| 313 |
+
|
| 314 |
+
weight = torch.zeros((1, 1, H, W), dtype=torch_dtype, device=data_device)
|
| 315 |
+
values = torch.zeros((B, C, H, W), dtype=torch_dtype, device=data_device)
|
| 316 |
+
|
| 317 |
+
# Split tasks
|
| 318 |
+
tasks = []
|
| 319 |
+
for h in range(0, H, tile_stride):
|
| 320 |
+
for w in range(0, W, tile_stride):
|
| 321 |
+
if (h-tile_stride >= 0 and h-tile_stride+tile_size >= H) or (w-tile_stride >= 0 and w-tile_stride+tile_size >= W):
|
| 322 |
+
continue
|
| 323 |
+
h_, w_ = h + tile_size, w + tile_size
|
| 324 |
+
if h_ > H: h, h_ = H - tile_size, H
|
| 325 |
+
if w_ > W: w, w_ = W - tile_size, W
|
| 326 |
+
tasks.append((h, h_, w, w_))
|
| 327 |
+
|
| 328 |
+
# Run
|
| 329 |
+
for hl, hr, wl, wr in tasks:
|
| 330 |
+
hidden_states_batch = hidden_states[:, :, hl:hr, wl:wr].to(computation_device)
|
| 331 |
+
hidden_states_batch = rearrange(hidden_states_batch, "B C H W -> B (H W) C")
|
| 332 |
+
if residual is not None:
|
| 333 |
+
residual_batch = residual[:, :, hl:hr, wl:wr].to(computation_device)
|
| 334 |
+
residual_batch = rearrange(residual_batch, "B C H W -> B (H W) C")
|
| 335 |
+
else:
|
| 336 |
+
residual_batch = None
|
| 337 |
+
|
| 338 |
+
# Forward
|
| 339 |
+
hidden_states_batch = block(hidden_states_batch, condition_emb, text_emb, freq_cis_img, residual_batch).to(data_device)
|
| 340 |
+
hidden_states_batch = rearrange(hidden_states_batch, "B (H W) C -> B C H W", H=hr-hl)
|
| 341 |
+
|
| 342 |
+
mask = self.build_mask(hidden_states_batch, is_bound=(hl==0, hr>=H, wl==0, wr>=W))
|
| 343 |
+
values[:, :, hl:hr, wl:wr] += hidden_states_batch * mask
|
| 344 |
+
weight[:, :, hl:hr, wl:wr] += mask
|
| 345 |
+
values /= weight
|
| 346 |
+
return values
|
| 347 |
+
|
| 348 |
+
def forward(
|
| 349 |
+
self, hidden_states, text_emb, text_emb_t5, text_emb_mask, text_emb_mask_t5, timestep, size_emb, freq_cis_img,
|
| 350 |
+
tiled=False, tile_size=64, tile_stride=32,
|
| 351 |
+
to_cache=False,
|
| 352 |
+
use_gradient_checkpointing=False,
|
| 353 |
+
):
|
| 354 |
+
# Embeddings
|
| 355 |
+
text_emb = self.prepare_text_emb(text_emb, text_emb_t5, text_emb_mask, text_emb_mask_t5)
|
| 356 |
+
condition_emb = self.prepare_extra_emb(text_emb_t5, timestep, size_emb, hidden_states.dtype, hidden_states.shape[0])
|
| 357 |
+
|
| 358 |
+
# Input
|
| 359 |
+
height, width = hidden_states.shape[-2], hidden_states.shape[-1]
|
| 360 |
+
hidden_states = self.patch_embedder(hidden_states)
|
| 361 |
+
|
| 362 |
+
# Blocks
|
| 363 |
+
def create_custom_forward(module):
|
| 364 |
+
def custom_forward(*inputs):
|
| 365 |
+
return module(*inputs)
|
| 366 |
+
return custom_forward
|
| 367 |
+
if tiled:
|
| 368 |
+
hidden_states = rearrange(hidden_states, "B (H W) C -> B C H W", H=height//2)
|
| 369 |
+
residuals = []
|
| 370 |
+
for block_id, block in enumerate(self.blocks):
|
| 371 |
+
residual = residuals.pop() if block_id >= self.num_layers_down else None
|
| 372 |
+
hidden_states = self.tiled_block_forward(
|
| 373 |
+
block, hidden_states, condition_emb, text_emb, freq_cis_img, residual,
|
| 374 |
+
torch_dtype=hidden_states.dtype, data_device=hidden_states.device, computation_device=hidden_states.device,
|
| 375 |
+
tile_size=tile_size, tile_stride=tile_stride
|
| 376 |
+
)
|
| 377 |
+
if block_id < self.num_layers_down - 2:
|
| 378 |
+
residuals.append(hidden_states)
|
| 379 |
+
hidden_states = rearrange(hidden_states, "B C H W -> B (H W) C")
|
| 380 |
+
else:
|
| 381 |
+
residuals = []
|
| 382 |
+
for block_id, block in enumerate(self.blocks):
|
| 383 |
+
residual = residuals.pop() if block_id >= self.num_layers_down else None
|
| 384 |
+
if self.training and use_gradient_checkpointing:
|
| 385 |
+
hidden_states = torch.utils.checkpoint.checkpoint(
|
| 386 |
+
create_custom_forward(block),
|
| 387 |
+
hidden_states, condition_emb, text_emb, freq_cis_img, residual,
|
| 388 |
+
use_reentrant=False,
|
| 389 |
+
)
|
| 390 |
+
else:
|
| 391 |
+
hidden_states = block(hidden_states, condition_emb, text_emb, freq_cis_img, residual, to_cache=to_cache)
|
| 392 |
+
if block_id < self.num_layers_down - 2:
|
| 393 |
+
residuals.append(hidden_states)
|
| 394 |
+
|
| 395 |
+
# Output
|
| 396 |
+
hidden_states = self.final_layer(hidden_states, condition_emb)
|
| 397 |
+
hidden_states = self.unpatchify(hidden_states, height//2, width//2)
|
| 398 |
+
hidden_states, _ = hidden_states.chunk(2, dim=1)
|
| 399 |
+
return hidden_states
|
| 400 |
+
|
| 401 |
+
@staticmethod
|
| 402 |
+
def state_dict_converter():
|
| 403 |
+
return HunyuanDiTStateDictConverter()
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
|
| 407 |
+
class HunyuanDiTStateDictConverter():
|
| 408 |
+
def __init__(self):
|
| 409 |
+
pass
|
| 410 |
+
|
| 411 |
+
def from_diffusers(self, state_dict):
|
| 412 |
+
state_dict_ = {}
|
| 413 |
+
for name, param in state_dict.items():
|
| 414 |
+
name_ = name
|
| 415 |
+
name_ = name_.replace(".default_modulation.", ".modulation.")
|
| 416 |
+
name_ = name_.replace(".mlp.fc1.", ".mlp.0.")
|
| 417 |
+
name_ = name_.replace(".mlp.fc2.", ".mlp.2.")
|
| 418 |
+
name_ = name_.replace(".attn1.q_norm.", ".rota1.q_norm.")
|
| 419 |
+
name_ = name_.replace(".attn2.q_norm.", ".rota2.q_norm.")
|
| 420 |
+
name_ = name_.replace(".attn1.k_norm.", ".rota1.k_norm.")
|
| 421 |
+
name_ = name_.replace(".attn2.k_norm.", ".rota2.k_norm.")
|
| 422 |
+
name_ = name_.replace(".q_proj.", ".to_q.")
|
| 423 |
+
name_ = name_.replace(".out_proj.", ".to_out.")
|
| 424 |
+
name_ = name_.replace("text_embedding_padding", "text_emb_padding")
|
| 425 |
+
name_ = name_.replace("mlp_t5.0.", "t5_embedder.0.")
|
| 426 |
+
name_ = name_.replace("mlp_t5.2.", "t5_embedder.2.")
|
| 427 |
+
name_ = name_.replace("pooler.", "t5_pooler.")
|
| 428 |
+
name_ = name_.replace("x_embedder.", "patch_embedder.")
|
| 429 |
+
name_ = name_.replace("t_embedder.", "timestep_embedder.")
|
| 430 |
+
name_ = name_.replace("t5_pooler.to_q.", "t5_pooler.q_proj.")
|
| 431 |
+
name_ = name_.replace("style_embedder.weight", "style_embedder")
|
| 432 |
+
if ".kv_proj." in name_:
|
| 433 |
+
param_k = param[:param.shape[0]//2]
|
| 434 |
+
param_v = param[param.shape[0]//2:]
|
| 435 |
+
state_dict_[name_.replace(".kv_proj.", ".to_k.")] = param_k
|
| 436 |
+
state_dict_[name_.replace(".kv_proj.", ".to_v.")] = param_v
|
| 437 |
+
elif ".Wqkv." in name_:
|
| 438 |
+
param_q = param[:param.shape[0]//3]
|
| 439 |
+
param_k = param[param.shape[0]//3:param.shape[0]//3*2]
|
| 440 |
+
param_v = param[param.shape[0]//3*2:]
|
| 441 |
+
state_dict_[name_.replace(".Wqkv.", ".to_q.")] = param_q
|
| 442 |
+
state_dict_[name_.replace(".Wqkv.", ".to_k.")] = param_k
|
| 443 |
+
state_dict_[name_.replace(".Wqkv.", ".to_v.")] = param_v
|
| 444 |
+
elif "style_embedder" in name_:
|
| 445 |
+
state_dict_[name_] = param.squeeze()
|
| 446 |
+
else:
|
| 447 |
+
state_dict_[name_] = param
|
| 448 |
+
return state_dict_
|
| 449 |
+
|
| 450 |
+
def from_civitai(self, state_dict):
|
| 451 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/hunyuan_dit_text_encoder.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from transformers import BertModel, BertConfig, T5EncoderModel, T5Config
|
| 2 |
+
import torch
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class HunyuanDiTCLIPTextEncoder(BertModel):
|
| 7 |
+
def __init__(self):
|
| 8 |
+
config = BertConfig(
|
| 9 |
+
_name_or_path = "",
|
| 10 |
+
architectures = ["BertModel"],
|
| 11 |
+
attention_probs_dropout_prob = 0.1,
|
| 12 |
+
bos_token_id = 0,
|
| 13 |
+
classifier_dropout = None,
|
| 14 |
+
directionality = "bidi",
|
| 15 |
+
eos_token_id = 2,
|
| 16 |
+
hidden_act = "gelu",
|
| 17 |
+
hidden_dropout_prob = 0.1,
|
| 18 |
+
hidden_size = 1024,
|
| 19 |
+
initializer_range = 0.02,
|
| 20 |
+
intermediate_size = 4096,
|
| 21 |
+
layer_norm_eps = 1e-12,
|
| 22 |
+
max_position_embeddings = 512,
|
| 23 |
+
model_type = "bert",
|
| 24 |
+
num_attention_heads = 16,
|
| 25 |
+
num_hidden_layers = 24,
|
| 26 |
+
output_past = True,
|
| 27 |
+
pad_token_id = 0,
|
| 28 |
+
pooler_fc_size = 768,
|
| 29 |
+
pooler_num_attention_heads = 12,
|
| 30 |
+
pooler_num_fc_layers = 3,
|
| 31 |
+
pooler_size_per_head = 128,
|
| 32 |
+
pooler_type = "first_token_transform",
|
| 33 |
+
position_embedding_type = "absolute",
|
| 34 |
+
torch_dtype = "float32",
|
| 35 |
+
transformers_version = "4.37.2",
|
| 36 |
+
type_vocab_size = 2,
|
| 37 |
+
use_cache = True,
|
| 38 |
+
vocab_size = 47020
|
| 39 |
+
)
|
| 40 |
+
super().__init__(config, add_pooling_layer=False)
|
| 41 |
+
self.eval()
|
| 42 |
+
|
| 43 |
+
def forward(self, input_ids, attention_mask, clip_skip=1):
|
| 44 |
+
input_shape = input_ids.size()
|
| 45 |
+
|
| 46 |
+
batch_size, seq_length = input_shape
|
| 47 |
+
device = input_ids.device
|
| 48 |
+
|
| 49 |
+
past_key_values_length = 0
|
| 50 |
+
|
| 51 |
+
if attention_mask is None:
|
| 52 |
+
attention_mask = torch.ones(((batch_size, seq_length + past_key_values_length)), device=device)
|
| 53 |
+
|
| 54 |
+
extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(attention_mask, input_shape)
|
| 55 |
+
|
| 56 |
+
embedding_output = self.embeddings(
|
| 57 |
+
input_ids=input_ids,
|
| 58 |
+
position_ids=None,
|
| 59 |
+
token_type_ids=None,
|
| 60 |
+
inputs_embeds=None,
|
| 61 |
+
past_key_values_length=0,
|
| 62 |
+
)
|
| 63 |
+
encoder_outputs = self.encoder(
|
| 64 |
+
embedding_output,
|
| 65 |
+
attention_mask=extended_attention_mask,
|
| 66 |
+
head_mask=None,
|
| 67 |
+
encoder_hidden_states=None,
|
| 68 |
+
encoder_attention_mask=None,
|
| 69 |
+
past_key_values=None,
|
| 70 |
+
use_cache=False,
|
| 71 |
+
output_attentions=False,
|
| 72 |
+
output_hidden_states=True,
|
| 73 |
+
return_dict=True,
|
| 74 |
+
)
|
| 75 |
+
all_hidden_states = encoder_outputs.hidden_states
|
| 76 |
+
prompt_emb = all_hidden_states[-clip_skip]
|
| 77 |
+
if clip_skip > 1:
|
| 78 |
+
mean, std = all_hidden_states[-1].mean(), all_hidden_states[-1].std()
|
| 79 |
+
prompt_emb = (prompt_emb - prompt_emb.mean()) / prompt_emb.std() * std + mean
|
| 80 |
+
return prompt_emb
|
| 81 |
+
|
| 82 |
+
@staticmethod
|
| 83 |
+
def state_dict_converter():
|
| 84 |
+
return HunyuanDiTCLIPTextEncoderStateDictConverter()
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
class HunyuanDiTT5TextEncoder(T5EncoderModel):
|
| 89 |
+
def __init__(self):
|
| 90 |
+
config = T5Config(
|
| 91 |
+
_name_or_path = "../HunyuanDiT/t2i/mt5",
|
| 92 |
+
architectures = ["MT5ForConditionalGeneration"],
|
| 93 |
+
classifier_dropout = 0.0,
|
| 94 |
+
d_ff = 5120,
|
| 95 |
+
d_kv = 64,
|
| 96 |
+
d_model = 2048,
|
| 97 |
+
decoder_start_token_id = 0,
|
| 98 |
+
dense_act_fn = "gelu_new",
|
| 99 |
+
dropout_rate = 0.1,
|
| 100 |
+
eos_token_id = 1,
|
| 101 |
+
feed_forward_proj = "gated-gelu",
|
| 102 |
+
initializer_factor = 1.0,
|
| 103 |
+
is_encoder_decoder = True,
|
| 104 |
+
is_gated_act = True,
|
| 105 |
+
layer_norm_epsilon = 1e-06,
|
| 106 |
+
model_type = "t5",
|
| 107 |
+
num_decoder_layers = 24,
|
| 108 |
+
num_heads = 32,
|
| 109 |
+
num_layers = 24,
|
| 110 |
+
output_past = True,
|
| 111 |
+
pad_token_id = 0,
|
| 112 |
+
relative_attention_max_distance = 128,
|
| 113 |
+
relative_attention_num_buckets = 32,
|
| 114 |
+
tie_word_embeddings = False,
|
| 115 |
+
tokenizer_class = "T5Tokenizer",
|
| 116 |
+
transformers_version = "4.37.2",
|
| 117 |
+
use_cache = True,
|
| 118 |
+
vocab_size = 250112
|
| 119 |
+
)
|
| 120 |
+
super().__init__(config)
|
| 121 |
+
self.eval()
|
| 122 |
+
|
| 123 |
+
def forward(self, input_ids, attention_mask, clip_skip=1):
|
| 124 |
+
outputs = super().forward(
|
| 125 |
+
input_ids=input_ids,
|
| 126 |
+
attention_mask=attention_mask,
|
| 127 |
+
output_hidden_states=True,
|
| 128 |
+
)
|
| 129 |
+
prompt_emb = outputs.hidden_states[-clip_skip]
|
| 130 |
+
if clip_skip > 1:
|
| 131 |
+
mean, std = outputs.hidden_states[-1].mean(), outputs.hidden_states[-1].std()
|
| 132 |
+
prompt_emb = (prompt_emb - prompt_emb.mean()) / prompt_emb.std() * std + mean
|
| 133 |
+
return prompt_emb
|
| 134 |
+
|
| 135 |
+
@staticmethod
|
| 136 |
+
def state_dict_converter():
|
| 137 |
+
return HunyuanDiTT5TextEncoderStateDictConverter()
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
class HunyuanDiTCLIPTextEncoderStateDictConverter():
|
| 142 |
+
def __init__(self):
|
| 143 |
+
pass
|
| 144 |
+
|
| 145 |
+
def from_diffusers(self, state_dict):
|
| 146 |
+
state_dict_ = {name[5:]: param for name, param in state_dict.items() if name.startswith("bert.")}
|
| 147 |
+
return state_dict_
|
| 148 |
+
|
| 149 |
+
def from_civitai(self, state_dict):
|
| 150 |
+
return self.from_diffusers(state_dict)
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
class HunyuanDiTT5TextEncoderStateDictConverter():
|
| 154 |
+
def __init__(self):
|
| 155 |
+
pass
|
| 156 |
+
|
| 157 |
+
def from_diffusers(self, state_dict):
|
| 158 |
+
state_dict_ = {name: param for name, param in state_dict.items() if name.startswith("encoder.")}
|
| 159 |
+
state_dict_["shared.weight"] = state_dict["shared.weight"]
|
| 160 |
+
return state_dict_
|
| 161 |
+
|
| 162 |
+
def from_civitai(self, state_dict):
|
| 163 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/hunyuan_video_dit.py
ADDED
|
@@ -0,0 +1,920 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd3_dit import TimestepEmbeddings, RMSNorm
|
| 3 |
+
from .utils import init_weights_on_device
|
| 4 |
+
from einops import rearrange, repeat
|
| 5 |
+
from tqdm import tqdm
|
| 6 |
+
from typing import Union, Tuple, List
|
| 7 |
+
from .utils import hash_state_dict_keys
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def HunyuanVideoRope(latents):
|
| 11 |
+
def _to_tuple(x, dim=2):
|
| 12 |
+
if isinstance(x, int):
|
| 13 |
+
return (x,) * dim
|
| 14 |
+
elif len(x) == dim:
|
| 15 |
+
return x
|
| 16 |
+
else:
|
| 17 |
+
raise ValueError(f"Expected length {dim} or int, but got {x}")
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def get_meshgrid_nd(start, *args, dim=2):
|
| 21 |
+
"""
|
| 22 |
+
Get n-D meshgrid with start, stop and num.
|
| 23 |
+
|
| 24 |
+
Args:
|
| 25 |
+
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
|
| 26 |
+
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
|
| 27 |
+
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
|
| 28 |
+
n-tuples.
|
| 29 |
+
*args: See above.
|
| 30 |
+
dim (int): Dimension of the meshgrid. Defaults to 2.
|
| 31 |
+
|
| 32 |
+
Returns:
|
| 33 |
+
grid (np.ndarray): [dim, ...]
|
| 34 |
+
"""
|
| 35 |
+
if len(args) == 0:
|
| 36 |
+
# start is grid_size
|
| 37 |
+
num = _to_tuple(start, dim=dim)
|
| 38 |
+
start = (0,) * dim
|
| 39 |
+
stop = num
|
| 40 |
+
elif len(args) == 1:
|
| 41 |
+
# start is start, args[0] is stop, step is 1
|
| 42 |
+
start = _to_tuple(start, dim=dim)
|
| 43 |
+
stop = _to_tuple(args[0], dim=dim)
|
| 44 |
+
num = [stop[i] - start[i] for i in range(dim)]
|
| 45 |
+
elif len(args) == 2:
|
| 46 |
+
# start is start, args[0] is stop, args[1] is num
|
| 47 |
+
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
|
| 48 |
+
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
|
| 49 |
+
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
|
| 50 |
+
else:
|
| 51 |
+
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
|
| 52 |
+
|
| 53 |
+
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
|
| 54 |
+
axis_grid = []
|
| 55 |
+
for i in range(dim):
|
| 56 |
+
a, b, n = start[i], stop[i], num[i]
|
| 57 |
+
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
|
| 58 |
+
axis_grid.append(g)
|
| 59 |
+
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
|
| 60 |
+
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
|
| 61 |
+
|
| 62 |
+
return grid
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def get_1d_rotary_pos_embed(
|
| 66 |
+
dim: int,
|
| 67 |
+
pos: Union[torch.FloatTensor, int],
|
| 68 |
+
theta: float = 10000.0,
|
| 69 |
+
use_real: bool = False,
|
| 70 |
+
theta_rescale_factor: float = 1.0,
|
| 71 |
+
interpolation_factor: float = 1.0,
|
| 72 |
+
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
| 73 |
+
"""
|
| 74 |
+
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
| 75 |
+
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
|
| 76 |
+
|
| 77 |
+
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
|
| 78 |
+
and the end index 'end'. The 'theta' parameter scales the frequencies.
|
| 79 |
+
The returned tensor contains complex values in complex64 data type.
|
| 80 |
+
|
| 81 |
+
Args:
|
| 82 |
+
dim (int): Dimension of the frequency tensor.
|
| 83 |
+
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
|
| 84 |
+
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
|
| 85 |
+
use_real (bool, optional): If True, return real part and imaginary part separately.
|
| 86 |
+
Otherwise, return complex numbers.
|
| 87 |
+
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
|
| 88 |
+
|
| 89 |
+
Returns:
|
| 90 |
+
freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]
|
| 91 |
+
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
|
| 92 |
+
"""
|
| 93 |
+
if isinstance(pos, int):
|
| 94 |
+
pos = torch.arange(pos).float()
|
| 95 |
+
|
| 96 |
+
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
| 97 |
+
# has some connection to NTK literature
|
| 98 |
+
if theta_rescale_factor != 1.0:
|
| 99 |
+
theta *= theta_rescale_factor ** (dim / (dim - 2))
|
| 100 |
+
|
| 101 |
+
freqs = 1.0 / (
|
| 102 |
+
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
|
| 103 |
+
) # [D/2]
|
| 104 |
+
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
|
| 105 |
+
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
| 106 |
+
if use_real:
|
| 107 |
+
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
| 108 |
+
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
|
| 109 |
+
return freqs_cos, freqs_sin
|
| 110 |
+
else:
|
| 111 |
+
freqs_cis = torch.polar(
|
| 112 |
+
torch.ones_like(freqs), freqs
|
| 113 |
+
) # complex64 # [S, D/2]
|
| 114 |
+
return freqs_cis
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def get_nd_rotary_pos_embed(
|
| 118 |
+
rope_dim_list,
|
| 119 |
+
start,
|
| 120 |
+
*args,
|
| 121 |
+
theta=10000.0,
|
| 122 |
+
use_real=False,
|
| 123 |
+
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
| 124 |
+
interpolation_factor: Union[float, List[float]] = 1.0,
|
| 125 |
+
):
|
| 126 |
+
"""
|
| 127 |
+
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
| 128 |
+
|
| 129 |
+
Args:
|
| 130 |
+
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
|
| 131 |
+
sum(rope_dim_list) should equal to head_dim of attention layer.
|
| 132 |
+
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
|
| 133 |
+
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
|
| 134 |
+
*args: See above.
|
| 135 |
+
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
|
| 136 |
+
use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.
|
| 137 |
+
Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real
|
| 138 |
+
part and an imaginary part separately.
|
| 139 |
+
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
|
| 140 |
+
|
| 141 |
+
Returns:
|
| 142 |
+
pos_embed (torch.Tensor): [HW, D/2]
|
| 143 |
+
"""
|
| 144 |
+
|
| 145 |
+
grid = get_meshgrid_nd(
|
| 146 |
+
start, *args, dim=len(rope_dim_list)
|
| 147 |
+
) # [3, W, H, D] / [2, W, H]
|
| 148 |
+
|
| 149 |
+
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
|
| 150 |
+
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
| 151 |
+
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
|
| 152 |
+
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
|
| 153 |
+
assert len(theta_rescale_factor) == len(
|
| 154 |
+
rope_dim_list
|
| 155 |
+
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
| 156 |
+
|
| 157 |
+
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
|
| 158 |
+
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
| 159 |
+
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
|
| 160 |
+
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
|
| 161 |
+
assert len(interpolation_factor) == len(
|
| 162 |
+
rope_dim_list
|
| 163 |
+
), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
| 164 |
+
|
| 165 |
+
# use 1/ndim of dimensions to encode grid_axis
|
| 166 |
+
embs = []
|
| 167 |
+
for i in range(len(rope_dim_list)):
|
| 168 |
+
emb = get_1d_rotary_pos_embed(
|
| 169 |
+
rope_dim_list[i],
|
| 170 |
+
grid[i].reshape(-1),
|
| 171 |
+
theta,
|
| 172 |
+
use_real=use_real,
|
| 173 |
+
theta_rescale_factor=theta_rescale_factor[i],
|
| 174 |
+
interpolation_factor=interpolation_factor[i],
|
| 175 |
+
) # 2 x [WHD, rope_dim_list[i]]
|
| 176 |
+
embs.append(emb)
|
| 177 |
+
|
| 178 |
+
if use_real:
|
| 179 |
+
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
|
| 180 |
+
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
|
| 181 |
+
return cos, sin
|
| 182 |
+
else:
|
| 183 |
+
emb = torch.cat(embs, dim=1) # (WHD, D/2)
|
| 184 |
+
return emb
|
| 185 |
+
|
| 186 |
+
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
| 187 |
+
[16, 56, 56],
|
| 188 |
+
[latents.shape[2], latents.shape[3] // 2, latents.shape[4] // 2],
|
| 189 |
+
theta=256,
|
| 190 |
+
use_real=True,
|
| 191 |
+
theta_rescale_factor=1,
|
| 192 |
+
)
|
| 193 |
+
return freqs_cos, freqs_sin
|
| 194 |
+
|
| 195 |
+
|
| 196 |
+
class PatchEmbed(torch.nn.Module):
|
| 197 |
+
def __init__(self, patch_size=(1, 2, 2), in_channels=16, embed_dim=3072):
|
| 198 |
+
super().__init__()
|
| 199 |
+
self.proj = torch.nn.Conv3d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
|
| 200 |
+
|
| 201 |
+
def forward(self, x):
|
| 202 |
+
x = self.proj(x)
|
| 203 |
+
x = x.flatten(2).transpose(1, 2)
|
| 204 |
+
return x
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
class IndividualTokenRefinerBlock(torch.nn.Module):
|
| 208 |
+
def __init__(self, hidden_size=3072, num_heads=24):
|
| 209 |
+
super().__init__()
|
| 210 |
+
self.num_heads = num_heads
|
| 211 |
+
self.norm1 = torch.nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
|
| 212 |
+
self.self_attn_qkv = torch.nn.Linear(hidden_size, hidden_size * 3)
|
| 213 |
+
self.self_attn_proj = torch.nn.Linear(hidden_size, hidden_size)
|
| 214 |
+
|
| 215 |
+
self.norm2 = torch.nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
|
| 216 |
+
self.mlp = torch.nn.Sequential(
|
| 217 |
+
torch.nn.Linear(hidden_size, hidden_size * 4),
|
| 218 |
+
torch.nn.SiLU(),
|
| 219 |
+
torch.nn.Linear(hidden_size * 4, hidden_size)
|
| 220 |
+
)
|
| 221 |
+
self.adaLN_modulation = torch.nn.Sequential(
|
| 222 |
+
torch.nn.SiLU(),
|
| 223 |
+
torch.nn.Linear(hidden_size, hidden_size * 2, device="cuda", dtype=torch.bfloat16),
|
| 224 |
+
)
|
| 225 |
+
|
| 226 |
+
def forward(self, x, c, attn_mask=None):
|
| 227 |
+
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 228 |
+
|
| 229 |
+
norm_x = self.norm1(x)
|
| 230 |
+
qkv = self.self_attn_qkv(norm_x)
|
| 231 |
+
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
| 232 |
+
|
| 233 |
+
attn = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 234 |
+
attn = rearrange(attn, "B H L D -> B L (H D)")
|
| 235 |
+
|
| 236 |
+
x = x + self.self_attn_proj(attn) * gate_msa.unsqueeze(1)
|
| 237 |
+
x = x + self.mlp(self.norm2(x)) * gate_mlp.unsqueeze(1)
|
| 238 |
+
|
| 239 |
+
return x
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
class SingleTokenRefiner(torch.nn.Module):
|
| 243 |
+
def __init__(self, in_channels=4096, hidden_size=3072, depth=2):
|
| 244 |
+
super().__init__()
|
| 245 |
+
self.input_embedder = torch.nn.Linear(in_channels, hidden_size, bias=True)
|
| 246 |
+
self.t_embedder = TimestepEmbeddings(256, hidden_size, computation_device="cpu")
|
| 247 |
+
self.c_embedder = torch.nn.Sequential(
|
| 248 |
+
torch.nn.Linear(in_channels, hidden_size),
|
| 249 |
+
torch.nn.SiLU(),
|
| 250 |
+
torch.nn.Linear(hidden_size, hidden_size)
|
| 251 |
+
)
|
| 252 |
+
self.blocks = torch.nn.ModuleList([IndividualTokenRefinerBlock(hidden_size=hidden_size) for _ in range(depth)])
|
| 253 |
+
|
| 254 |
+
def forward(self, x, t, mask=None):
|
| 255 |
+
timestep_aware_representations = self.t_embedder(t, dtype=torch.float32)
|
| 256 |
+
|
| 257 |
+
mask_float = mask.float().unsqueeze(-1)
|
| 258 |
+
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
| 259 |
+
context_aware_representations = self.c_embedder(context_aware_representations)
|
| 260 |
+
c = timestep_aware_representations + context_aware_representations
|
| 261 |
+
|
| 262 |
+
x = self.input_embedder(x)
|
| 263 |
+
|
| 264 |
+
mask = mask.to(device=x.device, dtype=torch.bool)
|
| 265 |
+
mask = repeat(mask, "B L -> B 1 D L", D=mask.shape[-1])
|
| 266 |
+
mask = mask & mask.transpose(2, 3)
|
| 267 |
+
mask[:, :, :, 0] = True
|
| 268 |
+
|
| 269 |
+
for block in self.blocks:
|
| 270 |
+
x = block(x, c, mask)
|
| 271 |
+
|
| 272 |
+
return x
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
class ModulateDiT(torch.nn.Module):
|
| 276 |
+
def __init__(self, hidden_size, factor=6):
|
| 277 |
+
super().__init__()
|
| 278 |
+
self.act = torch.nn.SiLU()
|
| 279 |
+
self.linear = torch.nn.Linear(hidden_size, factor * hidden_size)
|
| 280 |
+
|
| 281 |
+
def forward(self, x):
|
| 282 |
+
return self.linear(self.act(x))
|
| 283 |
+
|
| 284 |
+
|
| 285 |
+
def modulate(x, shift=None, scale=None, tr_shift=None, tr_scale=None, tr_token=None):
|
| 286 |
+
if tr_shift is not None:
|
| 287 |
+
x_zero = x[:, :tr_token] * (1 + tr_scale.unsqueeze(1)) + tr_shift.unsqueeze(1)
|
| 288 |
+
x_orig = x[:, tr_token:] * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 289 |
+
x = torch.concat((x_zero, x_orig), dim=1)
|
| 290 |
+
return x
|
| 291 |
+
if scale is None and shift is None:
|
| 292 |
+
return x
|
| 293 |
+
elif shift is None:
|
| 294 |
+
return x * (1 + scale.unsqueeze(1))
|
| 295 |
+
elif scale is None:
|
| 296 |
+
return x + shift.unsqueeze(1)
|
| 297 |
+
else:
|
| 298 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 299 |
+
|
| 300 |
+
|
| 301 |
+
def reshape_for_broadcast(
|
| 302 |
+
freqs_cis,
|
| 303 |
+
x: torch.Tensor,
|
| 304 |
+
head_first=False,
|
| 305 |
+
):
|
| 306 |
+
ndim = x.ndim
|
| 307 |
+
assert 0 <= 1 < ndim
|
| 308 |
+
|
| 309 |
+
if isinstance(freqs_cis, tuple):
|
| 310 |
+
# freqs_cis: (cos, sin) in real space
|
| 311 |
+
if head_first:
|
| 312 |
+
assert freqs_cis[0].shape == (
|
| 313 |
+
x.shape[-2],
|
| 314 |
+
x.shape[-1],
|
| 315 |
+
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
|
| 316 |
+
shape = [
|
| 317 |
+
d if i == ndim - 2 or i == ndim - 1 else 1
|
| 318 |
+
for i, d in enumerate(x.shape)
|
| 319 |
+
]
|
| 320 |
+
else:
|
| 321 |
+
assert freqs_cis[0].shape == (
|
| 322 |
+
x.shape[1],
|
| 323 |
+
x.shape[-1],
|
| 324 |
+
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
|
| 325 |
+
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
| 326 |
+
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
|
| 327 |
+
else:
|
| 328 |
+
# freqs_cis: values in complex space
|
| 329 |
+
if head_first:
|
| 330 |
+
assert freqs_cis.shape == (
|
| 331 |
+
x.shape[-2],
|
| 332 |
+
x.shape[-1],
|
| 333 |
+
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
|
| 334 |
+
shape = [
|
| 335 |
+
d if i == ndim - 2 or i == ndim - 1 else 1
|
| 336 |
+
for i, d in enumerate(x.shape)
|
| 337 |
+
]
|
| 338 |
+
else:
|
| 339 |
+
assert freqs_cis.shape == (
|
| 340 |
+
x.shape[1],
|
| 341 |
+
x.shape[-1],
|
| 342 |
+
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
|
| 343 |
+
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
| 344 |
+
return freqs_cis.view(*shape)
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
def rotate_half(x):
|
| 348 |
+
x_real, x_imag = (
|
| 349 |
+
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
|
| 350 |
+
) # [B, S, H, D//2]
|
| 351 |
+
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
def apply_rotary_emb(
|
| 355 |
+
xq: torch.Tensor,
|
| 356 |
+
xk: torch.Tensor,
|
| 357 |
+
freqs_cis,
|
| 358 |
+
head_first: bool = False,
|
| 359 |
+
):
|
| 360 |
+
xk_out = None
|
| 361 |
+
if isinstance(freqs_cis, tuple):
|
| 362 |
+
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
|
| 363 |
+
cos, sin = cos.to(xq.device), sin.to(xq.device)
|
| 364 |
+
# real * cos - imag * sin
|
| 365 |
+
# imag * cos + real * sin
|
| 366 |
+
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
|
| 367 |
+
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
|
| 368 |
+
else:
|
| 369 |
+
# view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)
|
| 370 |
+
xq_ = torch.view_as_complex(
|
| 371 |
+
xq.float().reshape(*xq.shape[:-1], -1, 2)
|
| 372 |
+
) # [B, S, H, D//2]
|
| 373 |
+
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(
|
| 374 |
+
xq.device
|
| 375 |
+
) # [S, D//2] --> [1, S, 1, D//2]
|
| 376 |
+
# (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)
|
| 377 |
+
# view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)
|
| 378 |
+
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
|
| 379 |
+
xk_ = torch.view_as_complex(
|
| 380 |
+
xk.float().reshape(*xk.shape[:-1], -1, 2)
|
| 381 |
+
) # [B, S, H, D//2]
|
| 382 |
+
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
|
| 383 |
+
|
| 384 |
+
return xq_out, xk_out
|
| 385 |
+
|
| 386 |
+
|
| 387 |
+
def attention(q, k, v):
|
| 388 |
+
q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
|
| 389 |
+
x = torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
| 390 |
+
x = x.transpose(1, 2).flatten(2, 3)
|
| 391 |
+
return x
|
| 392 |
+
|
| 393 |
+
|
| 394 |
+
def apply_gate(x, gate, tr_gate=None, tr_token=None):
|
| 395 |
+
if tr_gate is not None:
|
| 396 |
+
x_zero = x[:, :tr_token] * tr_gate.unsqueeze(1)
|
| 397 |
+
x_orig = x[:, tr_token:] * gate.unsqueeze(1)
|
| 398 |
+
return torch.concat((x_zero, x_orig), dim=1)
|
| 399 |
+
else:
|
| 400 |
+
return x * gate.unsqueeze(1)
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
class MMDoubleStreamBlockComponent(torch.nn.Module):
|
| 404 |
+
def __init__(self, hidden_size=3072, heads_num=24, mlp_width_ratio=4):
|
| 405 |
+
super().__init__()
|
| 406 |
+
self.heads_num = heads_num
|
| 407 |
+
|
| 408 |
+
self.mod = ModulateDiT(hidden_size)
|
| 409 |
+
self.norm1 = torch.nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 410 |
+
|
| 411 |
+
self.to_qkv = torch.nn.Linear(hidden_size, hidden_size * 3)
|
| 412 |
+
self.norm_q = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 413 |
+
self.norm_k = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 414 |
+
self.to_out = torch.nn.Linear(hidden_size, hidden_size)
|
| 415 |
+
|
| 416 |
+
self.norm2 = torch.nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 417 |
+
self.ff = torch.nn.Sequential(
|
| 418 |
+
torch.nn.Linear(hidden_size, hidden_size * mlp_width_ratio),
|
| 419 |
+
torch.nn.GELU(approximate="tanh"),
|
| 420 |
+
torch.nn.Linear(hidden_size * mlp_width_ratio, hidden_size)
|
| 421 |
+
)
|
| 422 |
+
|
| 423 |
+
def forward(self, hidden_states, conditioning, freqs_cis=None, token_replace_vec=None, tr_token=None):
|
| 424 |
+
mod1_shift, mod1_scale, mod1_gate, mod2_shift, mod2_scale, mod2_gate = self.mod(conditioning).chunk(6, dim=-1)
|
| 425 |
+
if token_replace_vec is not None:
|
| 426 |
+
assert tr_token is not None
|
| 427 |
+
tr_mod1_shift, tr_mod1_scale, tr_mod1_gate, tr_mod2_shift, tr_mod2_scale, tr_mod2_gate = self.mod(token_replace_vec).chunk(6, dim=-1)
|
| 428 |
+
else:
|
| 429 |
+
tr_mod1_shift, tr_mod1_scale, tr_mod1_gate, tr_mod2_shift, tr_mod2_scale, tr_mod2_gate = None, None, None, None, None, None
|
| 430 |
+
|
| 431 |
+
norm_hidden_states = self.norm1(hidden_states)
|
| 432 |
+
norm_hidden_states = modulate(norm_hidden_states, shift=mod1_shift, scale=mod1_scale,
|
| 433 |
+
tr_shift=tr_mod1_shift, tr_scale=tr_mod1_scale, tr_token=tr_token)
|
| 434 |
+
qkv = self.to_qkv(norm_hidden_states)
|
| 435 |
+
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
| 436 |
+
|
| 437 |
+
q = self.norm_q(q)
|
| 438 |
+
k = self.norm_k(k)
|
| 439 |
+
|
| 440 |
+
if freqs_cis is not None:
|
| 441 |
+
q, k = apply_rotary_emb(q, k, freqs_cis, head_first=False)
|
| 442 |
+
return (q, k, v), (mod1_gate, mod2_shift, mod2_scale, mod2_gate), (tr_mod1_gate, tr_mod2_shift, tr_mod2_scale, tr_mod2_gate)
|
| 443 |
+
|
| 444 |
+
def process_ff(self, hidden_states, attn_output, mod, mod_tr=None, tr_token=None):
|
| 445 |
+
mod1_gate, mod2_shift, mod2_scale, mod2_gate = mod
|
| 446 |
+
if mod_tr is not None:
|
| 447 |
+
tr_mod1_gate, tr_mod2_shift, tr_mod2_scale, tr_mod2_gate = mod_tr
|
| 448 |
+
else:
|
| 449 |
+
tr_mod1_gate, tr_mod2_shift, tr_mod2_scale, tr_mod2_gate = None, None, None, None
|
| 450 |
+
hidden_states = hidden_states + apply_gate(self.to_out(attn_output), mod1_gate, tr_mod1_gate, tr_token)
|
| 451 |
+
x = self.ff(modulate(self.norm2(hidden_states), shift=mod2_shift, scale=mod2_scale, tr_shift=tr_mod2_shift, tr_scale=tr_mod2_scale, tr_token=tr_token))
|
| 452 |
+
hidden_states = hidden_states + apply_gate(x, mod2_gate, tr_mod2_gate, tr_token)
|
| 453 |
+
return hidden_states
|
| 454 |
+
|
| 455 |
+
|
| 456 |
+
class MMDoubleStreamBlock(torch.nn.Module):
|
| 457 |
+
def __init__(self, hidden_size=3072, heads_num=24, mlp_width_ratio=4):
|
| 458 |
+
super().__init__()
|
| 459 |
+
self.component_a = MMDoubleStreamBlockComponent(hidden_size, heads_num, mlp_width_ratio)
|
| 460 |
+
self.component_b = MMDoubleStreamBlockComponent(hidden_size, heads_num, mlp_width_ratio)
|
| 461 |
+
|
| 462 |
+
def forward(self, hidden_states_a, hidden_states_b, conditioning, freqs_cis, token_replace_vec=None, tr_token=None, split_token=71):
|
| 463 |
+
(q_a, k_a, v_a), mod_a, mod_tr = self.component_a(hidden_states_a, conditioning, freqs_cis, token_replace_vec, tr_token)
|
| 464 |
+
(q_b, k_b, v_b), mod_b, _ = self.component_b(hidden_states_b, conditioning, freqs_cis=None)
|
| 465 |
+
|
| 466 |
+
q_a, q_b = torch.concat([q_a, q_b[:, :split_token]], dim=1), q_b[:, split_token:].contiguous()
|
| 467 |
+
k_a, k_b = torch.concat([k_a, k_b[:, :split_token]], dim=1), k_b[:, split_token:].contiguous()
|
| 468 |
+
v_a, v_b = torch.concat([v_a, v_b[:, :split_token]], dim=1), v_b[:, split_token:].contiguous()
|
| 469 |
+
attn_output_a = attention(q_a, k_a, v_a)
|
| 470 |
+
attn_output_b = attention(q_b, k_b, v_b)
|
| 471 |
+
attn_output_a, attn_output_b = attn_output_a[:, :-split_token].contiguous(), torch.concat([attn_output_a[:, -split_token:], attn_output_b], dim=1)
|
| 472 |
+
|
| 473 |
+
hidden_states_a = self.component_a.process_ff(hidden_states_a, attn_output_a, mod_a, mod_tr, tr_token)
|
| 474 |
+
hidden_states_b = self.component_b.process_ff(hidden_states_b, attn_output_b, mod_b)
|
| 475 |
+
return hidden_states_a, hidden_states_b
|
| 476 |
+
|
| 477 |
+
|
| 478 |
+
class MMSingleStreamBlockOriginal(torch.nn.Module):
|
| 479 |
+
def __init__(self, hidden_size=3072, heads_num=24, mlp_width_ratio=4):
|
| 480 |
+
super().__init__()
|
| 481 |
+
self.hidden_size = hidden_size
|
| 482 |
+
self.heads_num = heads_num
|
| 483 |
+
self.mlp_hidden_dim = hidden_size * mlp_width_ratio
|
| 484 |
+
|
| 485 |
+
self.linear1 = torch.nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
|
| 486 |
+
self.linear2 = torch.nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
|
| 487 |
+
|
| 488 |
+
self.q_norm = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 489 |
+
self.k_norm = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 490 |
+
|
| 491 |
+
self.pre_norm = torch.nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 492 |
+
|
| 493 |
+
self.mlp_act = torch.nn.GELU(approximate="tanh")
|
| 494 |
+
self.modulation = ModulateDiT(hidden_size, factor=3)
|
| 495 |
+
|
| 496 |
+
def forward(self, x, vec, freqs_cis=None, txt_len=256):
|
| 497 |
+
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
| 498 |
+
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
|
| 499 |
+
qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
|
| 500 |
+
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
| 501 |
+
q = self.q_norm(q)
|
| 502 |
+
k = self.k_norm(k)
|
| 503 |
+
|
| 504 |
+
q_a, q_b = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
|
| 505 |
+
k_a, k_b = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
|
| 506 |
+
q_a, k_a = apply_rotary_emb(q_a, k_a, freqs_cis, head_first=False)
|
| 507 |
+
q = torch.cat((q_a, q_b), dim=1)
|
| 508 |
+
k = torch.cat((k_a, k_b), dim=1)
|
| 509 |
+
|
| 510 |
+
attn_output_a = attention(q[:, :-185].contiguous(), k[:, :-185].contiguous(), v[:, :-185].contiguous())
|
| 511 |
+
attn_output_b = attention(q[:, -185:].contiguous(), k[:, -185:].contiguous(), v[:, -185:].contiguous())
|
| 512 |
+
attn_output = torch.concat([attn_output_a, attn_output_b], dim=1)
|
| 513 |
+
|
| 514 |
+
output = self.linear2(torch.cat((attn_output, self.mlp_act(mlp)), 2))
|
| 515 |
+
return x + output * mod_gate.unsqueeze(1)
|
| 516 |
+
|
| 517 |
+
|
| 518 |
+
class MMSingleStreamBlock(torch.nn.Module):
|
| 519 |
+
def __init__(self, hidden_size=3072, heads_num=24, mlp_width_ratio=4):
|
| 520 |
+
super().__init__()
|
| 521 |
+
self.heads_num = heads_num
|
| 522 |
+
|
| 523 |
+
self.mod = ModulateDiT(hidden_size, factor=3)
|
| 524 |
+
self.norm = torch.nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 525 |
+
|
| 526 |
+
self.to_qkv = torch.nn.Linear(hidden_size, hidden_size * 3)
|
| 527 |
+
self.norm_q = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 528 |
+
self.norm_k = RMSNorm(dim=hidden_size // heads_num, eps=1e-6)
|
| 529 |
+
self.to_out = torch.nn.Linear(hidden_size, hidden_size)
|
| 530 |
+
|
| 531 |
+
self.ff = torch.nn.Sequential(
|
| 532 |
+
torch.nn.Linear(hidden_size, hidden_size * mlp_width_ratio),
|
| 533 |
+
torch.nn.GELU(approximate="tanh"),
|
| 534 |
+
torch.nn.Linear(hidden_size * mlp_width_ratio, hidden_size, bias=False)
|
| 535 |
+
)
|
| 536 |
+
|
| 537 |
+
def forward(self, hidden_states, conditioning, freqs_cis=None, txt_len=256, token_replace_vec=None, tr_token=None, split_token=71):
|
| 538 |
+
mod_shift, mod_scale, mod_gate = self.mod(conditioning).chunk(3, dim=-1)
|
| 539 |
+
if token_replace_vec is not None:
|
| 540 |
+
assert tr_token is not None
|
| 541 |
+
tr_mod_shift, tr_mod_scale, tr_mod_gate = self.mod(token_replace_vec).chunk(3, dim=-1)
|
| 542 |
+
else:
|
| 543 |
+
tr_mod_shift, tr_mod_scale, tr_mod_gate = None, None, None
|
| 544 |
+
|
| 545 |
+
norm_hidden_states = self.norm(hidden_states)
|
| 546 |
+
norm_hidden_states = modulate(norm_hidden_states, shift=mod_shift, scale=mod_scale,
|
| 547 |
+
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, tr_token=tr_token)
|
| 548 |
+
qkv = self.to_qkv(norm_hidden_states)
|
| 549 |
+
|
| 550 |
+
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
| 551 |
+
|
| 552 |
+
q = self.norm_q(q)
|
| 553 |
+
k = self.norm_k(k)
|
| 554 |
+
|
| 555 |
+
q_a, q_b = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
|
| 556 |
+
k_a, k_b = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
|
| 557 |
+
q_a, k_a = apply_rotary_emb(q_a, k_a, freqs_cis, head_first=False)
|
| 558 |
+
|
| 559 |
+
v_len = txt_len - split_token
|
| 560 |
+
q_a, q_b = torch.concat([q_a, q_b[:, :split_token]], dim=1), q_b[:, split_token:].contiguous()
|
| 561 |
+
k_a, k_b = torch.concat([k_a, k_b[:, :split_token]], dim=1), k_b[:, split_token:].contiguous()
|
| 562 |
+
v_a, v_b = v[:, :-v_len].contiguous(), v[:, -v_len:].contiguous()
|
| 563 |
+
|
| 564 |
+
attn_output_a = attention(q_a, k_a, v_a)
|
| 565 |
+
attn_output_b = attention(q_b, k_b, v_b)
|
| 566 |
+
attn_output = torch.concat([attn_output_a, attn_output_b], dim=1)
|
| 567 |
+
|
| 568 |
+
hidden_states = hidden_states + apply_gate(self.to_out(attn_output), mod_gate, tr_mod_gate, tr_token)
|
| 569 |
+
hidden_states = hidden_states + apply_gate(self.ff(norm_hidden_states), mod_gate, tr_mod_gate, tr_token)
|
| 570 |
+
return hidden_states
|
| 571 |
+
|
| 572 |
+
|
| 573 |
+
class FinalLayer(torch.nn.Module):
|
| 574 |
+
def __init__(self, hidden_size=3072, patch_size=(1, 2, 2), out_channels=16):
|
| 575 |
+
super().__init__()
|
| 576 |
+
|
| 577 |
+
self.norm_final = torch.nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 578 |
+
self.linear = torch.nn.Linear(hidden_size, patch_size[0] * patch_size[1] * patch_size[2] * out_channels)
|
| 579 |
+
|
| 580 |
+
self.adaLN_modulation = torch.nn.Sequential(torch.nn.SiLU(), torch.nn.Linear(hidden_size, 2 * hidden_size))
|
| 581 |
+
|
| 582 |
+
def forward(self, x, c):
|
| 583 |
+
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 584 |
+
x = modulate(self.norm_final(x), shift=shift, scale=scale)
|
| 585 |
+
x = self.linear(x)
|
| 586 |
+
return x
|
| 587 |
+
|
| 588 |
+
|
| 589 |
+
class HunyuanVideoDiT(torch.nn.Module):
|
| 590 |
+
def __init__(self, in_channels=16, hidden_size=3072, text_dim=4096, num_double_blocks=20, num_single_blocks=40, guidance_embed=True):
|
| 591 |
+
super().__init__()
|
| 592 |
+
self.img_in = PatchEmbed(in_channels=in_channels, embed_dim=hidden_size)
|
| 593 |
+
self.txt_in = SingleTokenRefiner(in_channels=text_dim, hidden_size=hidden_size)
|
| 594 |
+
self.time_in = TimestepEmbeddings(256, hidden_size, computation_device="cpu")
|
| 595 |
+
self.vector_in = torch.nn.Sequential(
|
| 596 |
+
torch.nn.Linear(768, hidden_size),
|
| 597 |
+
torch.nn.SiLU(),
|
| 598 |
+
torch.nn.Linear(hidden_size, hidden_size)
|
| 599 |
+
)
|
| 600 |
+
self.guidance_in = TimestepEmbeddings(256, hidden_size, computation_device="cpu") if guidance_embed else None
|
| 601 |
+
self.double_blocks = torch.nn.ModuleList([MMDoubleStreamBlock(hidden_size) for _ in range(num_double_blocks)])
|
| 602 |
+
self.single_blocks = torch.nn.ModuleList([MMSingleStreamBlock(hidden_size) for _ in range(num_single_blocks)])
|
| 603 |
+
self.final_layer = FinalLayer(hidden_size)
|
| 604 |
+
|
| 605 |
+
# TODO: remove these parameters
|
| 606 |
+
self.dtype = torch.bfloat16
|
| 607 |
+
self.patch_size = [1, 2, 2]
|
| 608 |
+
self.hidden_size = 3072
|
| 609 |
+
self.heads_num = 24
|
| 610 |
+
self.rope_dim_list = [16, 56, 56]
|
| 611 |
+
|
| 612 |
+
def unpatchify(self, x, T, H, W):
|
| 613 |
+
x = rearrange(x, "B (T H W) (C pT pH pW) -> B C (T pT) (H pH) (W pW)", H=H, W=W, pT=1, pH=2, pW=2)
|
| 614 |
+
return x
|
| 615 |
+
|
| 616 |
+
def enable_block_wise_offload(self, warm_device="cuda", cold_device="cpu"):
|
| 617 |
+
self.warm_device = warm_device
|
| 618 |
+
self.cold_device = cold_device
|
| 619 |
+
self.to(self.cold_device)
|
| 620 |
+
|
| 621 |
+
def load_models_to_device(self, loadmodel_names=[], device="cpu"):
|
| 622 |
+
for model_name in loadmodel_names:
|
| 623 |
+
model = getattr(self, model_name)
|
| 624 |
+
if model is not None:
|
| 625 |
+
model.to(device)
|
| 626 |
+
torch.cuda.empty_cache()
|
| 627 |
+
|
| 628 |
+
def prepare_freqs(self, latents):
|
| 629 |
+
return HunyuanVideoRope(latents)
|
| 630 |
+
|
| 631 |
+
def forward(
|
| 632 |
+
self,
|
| 633 |
+
x: torch.Tensor,
|
| 634 |
+
t: torch.Tensor,
|
| 635 |
+
prompt_emb: torch.Tensor = None,
|
| 636 |
+
text_mask: torch.Tensor = None,
|
| 637 |
+
pooled_prompt_emb: torch.Tensor = None,
|
| 638 |
+
freqs_cos: torch.Tensor = None,
|
| 639 |
+
freqs_sin: torch.Tensor = None,
|
| 640 |
+
guidance: torch.Tensor = None,
|
| 641 |
+
**kwargs
|
| 642 |
+
):
|
| 643 |
+
B, C, T, H, W = x.shape
|
| 644 |
+
|
| 645 |
+
vec = self.time_in(t, dtype=torch.float32) + self.vector_in(pooled_prompt_emb)
|
| 646 |
+
if self.guidance_in is not None:
|
| 647 |
+
vec += self.guidance_in(guidance * 1000, dtype=torch.float32)
|
| 648 |
+
img = self.img_in(x)
|
| 649 |
+
txt = self.txt_in(prompt_emb, t, text_mask)
|
| 650 |
+
|
| 651 |
+
for block in tqdm(self.double_blocks, desc="Double stream blocks"):
|
| 652 |
+
img, txt = block(img, txt, vec, (freqs_cos, freqs_sin))
|
| 653 |
+
|
| 654 |
+
x = torch.concat([img, txt], dim=1)
|
| 655 |
+
for block in tqdm(self.single_blocks, desc="Single stream blocks"):
|
| 656 |
+
x = block(x, vec, (freqs_cos, freqs_sin))
|
| 657 |
+
|
| 658 |
+
img = x[:, :-256]
|
| 659 |
+
img = self.final_layer(img, vec)
|
| 660 |
+
img = self.unpatchify(img, T=T//1, H=H//2, W=W//2)
|
| 661 |
+
return img
|
| 662 |
+
|
| 663 |
+
|
| 664 |
+
def enable_auto_offload(self, dtype=torch.bfloat16, device="cuda"):
|
| 665 |
+
def cast_to(weight, dtype=None, device=None, copy=False):
|
| 666 |
+
if device is None or weight.device == device:
|
| 667 |
+
if not copy:
|
| 668 |
+
if dtype is None or weight.dtype == dtype:
|
| 669 |
+
return weight
|
| 670 |
+
return weight.to(dtype=dtype, copy=copy)
|
| 671 |
+
|
| 672 |
+
r = torch.empty_like(weight, dtype=dtype, device=device)
|
| 673 |
+
r.copy_(weight)
|
| 674 |
+
return r
|
| 675 |
+
|
| 676 |
+
def cast_weight(s, input=None, dtype=None, device=None):
|
| 677 |
+
if input is not None:
|
| 678 |
+
if dtype is None:
|
| 679 |
+
dtype = input.dtype
|
| 680 |
+
if device is None:
|
| 681 |
+
device = input.device
|
| 682 |
+
weight = cast_to(s.weight, dtype, device)
|
| 683 |
+
return weight
|
| 684 |
+
|
| 685 |
+
def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
|
| 686 |
+
if input is not None:
|
| 687 |
+
if dtype is None:
|
| 688 |
+
dtype = input.dtype
|
| 689 |
+
if bias_dtype is None:
|
| 690 |
+
bias_dtype = dtype
|
| 691 |
+
if device is None:
|
| 692 |
+
device = input.device
|
| 693 |
+
weight = cast_to(s.weight, dtype, device)
|
| 694 |
+
bias = cast_to(s.bias, bias_dtype, device) if s.bias is not None else None
|
| 695 |
+
return weight, bias
|
| 696 |
+
|
| 697 |
+
class quantized_layer:
|
| 698 |
+
class Linear(torch.nn.Linear):
|
| 699 |
+
def __init__(self, *args, dtype=torch.bfloat16, device="cuda", **kwargs):
|
| 700 |
+
super().__init__(*args, **kwargs)
|
| 701 |
+
self.dtype = dtype
|
| 702 |
+
self.device = device
|
| 703 |
+
|
| 704 |
+
def block_forward_(self, x, i, j, dtype, device):
|
| 705 |
+
weight_ = cast_to(
|
| 706 |
+
self.weight[j * self.block_size: (j + 1) * self.block_size, i * self.block_size: (i + 1) * self.block_size],
|
| 707 |
+
dtype=dtype, device=device
|
| 708 |
+
)
|
| 709 |
+
if self.bias is None or i > 0:
|
| 710 |
+
bias_ = None
|
| 711 |
+
else:
|
| 712 |
+
bias_ = cast_to(self.bias[j * self.block_size: (j + 1) * self.block_size], dtype=dtype, device=device)
|
| 713 |
+
x_ = x[..., i * self.block_size: (i + 1) * self.block_size]
|
| 714 |
+
y_ = torch.nn.functional.linear(x_, weight_, bias_)
|
| 715 |
+
del x_, weight_, bias_
|
| 716 |
+
torch.cuda.empty_cache()
|
| 717 |
+
return y_
|
| 718 |
+
|
| 719 |
+
def block_forward(self, x, **kwargs):
|
| 720 |
+
# This feature can only reduce 2GB VRAM, so we disable it.
|
| 721 |
+
y = torch.zeros(x.shape[:-1] + (self.out_features,), dtype=x.dtype, device=x.device)
|
| 722 |
+
for i in range((self.in_features + self.block_size - 1) // self.block_size):
|
| 723 |
+
for j in range((self.out_features + self.block_size - 1) // self.block_size):
|
| 724 |
+
y[..., j * self.block_size: (j + 1) * self.block_size] += self.block_forward_(x, i, j, dtype=x.dtype, device=x.device)
|
| 725 |
+
return y
|
| 726 |
+
|
| 727 |
+
def forward(self, x, **kwargs):
|
| 728 |
+
weight, bias = cast_bias_weight(self, x, dtype=self.dtype, device=self.device)
|
| 729 |
+
return torch.nn.functional.linear(x, weight, bias)
|
| 730 |
+
|
| 731 |
+
|
| 732 |
+
class RMSNorm(torch.nn.Module):
|
| 733 |
+
def __init__(self, module, dtype=torch.bfloat16, device="cuda"):
|
| 734 |
+
super().__init__()
|
| 735 |
+
self.module = module
|
| 736 |
+
self.dtype = dtype
|
| 737 |
+
self.device = device
|
| 738 |
+
|
| 739 |
+
def forward(self, hidden_states, **kwargs):
|
| 740 |
+
input_dtype = hidden_states.dtype
|
| 741 |
+
variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)
|
| 742 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.module.eps)
|
| 743 |
+
hidden_states = hidden_states.to(input_dtype)
|
| 744 |
+
if self.module.weight is not None:
|
| 745 |
+
weight = cast_weight(self.module, hidden_states, dtype=torch.bfloat16, device="cuda")
|
| 746 |
+
hidden_states = hidden_states * weight
|
| 747 |
+
return hidden_states
|
| 748 |
+
|
| 749 |
+
class Conv3d(torch.nn.Conv3d):
|
| 750 |
+
def __init__(self, *args, dtype=torch.bfloat16, device="cuda", **kwargs):
|
| 751 |
+
super().__init__(*args, **kwargs)
|
| 752 |
+
self.dtype = dtype
|
| 753 |
+
self.device = device
|
| 754 |
+
|
| 755 |
+
def forward(self, x):
|
| 756 |
+
weight, bias = cast_bias_weight(self, x, dtype=self.dtype, device=self.device)
|
| 757 |
+
return torch.nn.functional.conv3d(x, weight, bias, self.stride, self.padding, self.dilation, self.groups)
|
| 758 |
+
|
| 759 |
+
class LayerNorm(torch.nn.LayerNorm):
|
| 760 |
+
def __init__(self, *args, dtype=torch.bfloat16, device="cuda", **kwargs):
|
| 761 |
+
super().__init__(*args, **kwargs)
|
| 762 |
+
self.dtype = dtype
|
| 763 |
+
self.device = device
|
| 764 |
+
|
| 765 |
+
def forward(self, x):
|
| 766 |
+
if self.weight is not None and self.bias is not None:
|
| 767 |
+
weight, bias = cast_bias_weight(self, x, dtype=self.dtype, device=self.device)
|
| 768 |
+
return torch.nn.functional.layer_norm(x, self.normalized_shape, weight, bias, self.eps)
|
| 769 |
+
else:
|
| 770 |
+
return torch.nn.functional.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
|
| 771 |
+
|
| 772 |
+
def replace_layer(model, dtype=torch.bfloat16, device="cuda"):
|
| 773 |
+
for name, module in model.named_children():
|
| 774 |
+
if isinstance(module, torch.nn.Linear):
|
| 775 |
+
with init_weights_on_device():
|
| 776 |
+
new_layer = quantized_layer.Linear(
|
| 777 |
+
module.in_features, module.out_features, bias=module.bias is not None,
|
| 778 |
+
dtype=dtype, device=device
|
| 779 |
+
)
|
| 780 |
+
new_layer.load_state_dict(module.state_dict(), assign=True)
|
| 781 |
+
setattr(model, name, new_layer)
|
| 782 |
+
elif isinstance(module, torch.nn.Conv3d):
|
| 783 |
+
with init_weights_on_device():
|
| 784 |
+
new_layer = quantized_layer.Conv3d(
|
| 785 |
+
module.in_channels, module.out_channels, kernel_size=module.kernel_size, stride=module.stride,
|
| 786 |
+
dtype=dtype, device=device
|
| 787 |
+
)
|
| 788 |
+
new_layer.load_state_dict(module.state_dict(), assign=True)
|
| 789 |
+
setattr(model, name, new_layer)
|
| 790 |
+
elif isinstance(module, RMSNorm):
|
| 791 |
+
new_layer = quantized_layer.RMSNorm(
|
| 792 |
+
module,
|
| 793 |
+
dtype=dtype, device=device
|
| 794 |
+
)
|
| 795 |
+
setattr(model, name, new_layer)
|
| 796 |
+
elif isinstance(module, torch.nn.LayerNorm):
|
| 797 |
+
with init_weights_on_device():
|
| 798 |
+
new_layer = quantized_layer.LayerNorm(
|
| 799 |
+
module.normalized_shape, elementwise_affine=module.elementwise_affine, eps=module.eps,
|
| 800 |
+
dtype=dtype, device=device
|
| 801 |
+
)
|
| 802 |
+
new_layer.load_state_dict(module.state_dict(), assign=True)
|
| 803 |
+
setattr(model, name, new_layer)
|
| 804 |
+
else:
|
| 805 |
+
replace_layer(module, dtype=dtype, device=device)
|
| 806 |
+
|
| 807 |
+
replace_layer(self, dtype=dtype, device=device)
|
| 808 |
+
|
| 809 |
+
@staticmethod
|
| 810 |
+
def state_dict_converter():
|
| 811 |
+
return HunyuanVideoDiTStateDictConverter()
|
| 812 |
+
|
| 813 |
+
|
| 814 |
+
class HunyuanVideoDiTStateDictConverter:
|
| 815 |
+
def __init__(self):
|
| 816 |
+
pass
|
| 817 |
+
|
| 818 |
+
def from_civitai(self, state_dict):
|
| 819 |
+
origin_hash_key = hash_state_dict_keys(state_dict, with_shape=True)
|
| 820 |
+
if "module" in state_dict:
|
| 821 |
+
state_dict = state_dict["module"]
|
| 822 |
+
direct_dict = {
|
| 823 |
+
"img_in.proj": "img_in.proj",
|
| 824 |
+
"time_in.mlp.0": "time_in.timestep_embedder.0",
|
| 825 |
+
"time_in.mlp.2": "time_in.timestep_embedder.2",
|
| 826 |
+
"vector_in.in_layer": "vector_in.0",
|
| 827 |
+
"vector_in.out_layer": "vector_in.2",
|
| 828 |
+
"guidance_in.mlp.0": "guidance_in.timestep_embedder.0",
|
| 829 |
+
"guidance_in.mlp.2": "guidance_in.timestep_embedder.2",
|
| 830 |
+
"txt_in.input_embedder": "txt_in.input_embedder",
|
| 831 |
+
"txt_in.t_embedder.mlp.0": "txt_in.t_embedder.timestep_embedder.0",
|
| 832 |
+
"txt_in.t_embedder.mlp.2": "txt_in.t_embedder.timestep_embedder.2",
|
| 833 |
+
"txt_in.c_embedder.linear_1": "txt_in.c_embedder.0",
|
| 834 |
+
"txt_in.c_embedder.linear_2": "txt_in.c_embedder.2",
|
| 835 |
+
"final_layer.linear": "final_layer.linear",
|
| 836 |
+
"final_layer.adaLN_modulation.1": "final_layer.adaLN_modulation.1",
|
| 837 |
+
}
|
| 838 |
+
txt_suffix_dict = {
|
| 839 |
+
"norm1": "norm1",
|
| 840 |
+
"self_attn_qkv": "self_attn_qkv",
|
| 841 |
+
"self_attn_proj": "self_attn_proj",
|
| 842 |
+
"norm2": "norm2",
|
| 843 |
+
"mlp.fc1": "mlp.0",
|
| 844 |
+
"mlp.fc2": "mlp.2",
|
| 845 |
+
"adaLN_modulation.1": "adaLN_modulation.1",
|
| 846 |
+
}
|
| 847 |
+
double_suffix_dict = {
|
| 848 |
+
"img_mod.linear": "component_a.mod.linear",
|
| 849 |
+
"img_attn_qkv": "component_a.to_qkv",
|
| 850 |
+
"img_attn_q_norm": "component_a.norm_q",
|
| 851 |
+
"img_attn_k_norm": "component_a.norm_k",
|
| 852 |
+
"img_attn_proj": "component_a.to_out",
|
| 853 |
+
"img_mlp.fc1": "component_a.ff.0",
|
| 854 |
+
"img_mlp.fc2": "component_a.ff.2",
|
| 855 |
+
"txt_mod.linear": "component_b.mod.linear",
|
| 856 |
+
"txt_attn_qkv": "component_b.to_qkv",
|
| 857 |
+
"txt_attn_q_norm": "component_b.norm_q",
|
| 858 |
+
"txt_attn_k_norm": "component_b.norm_k",
|
| 859 |
+
"txt_attn_proj": "component_b.to_out",
|
| 860 |
+
"txt_mlp.fc1": "component_b.ff.0",
|
| 861 |
+
"txt_mlp.fc2": "component_b.ff.2",
|
| 862 |
+
}
|
| 863 |
+
single_suffix_dict = {
|
| 864 |
+
"linear1": ["to_qkv", "ff.0"],
|
| 865 |
+
"linear2": ["to_out", "ff.2"],
|
| 866 |
+
"q_norm": "norm_q",
|
| 867 |
+
"k_norm": "norm_k",
|
| 868 |
+
"modulation.linear": "mod.linear",
|
| 869 |
+
}
|
| 870 |
+
# single_suffix_dict = {
|
| 871 |
+
# "linear1": "linear1",
|
| 872 |
+
# "linear2": "linear2",
|
| 873 |
+
# "q_norm": "q_norm",
|
| 874 |
+
# "k_norm": "k_norm",
|
| 875 |
+
# "modulation.linear": "modulation.linear",
|
| 876 |
+
# }
|
| 877 |
+
state_dict_ = {}
|
| 878 |
+
for name, param in state_dict.items():
|
| 879 |
+
names = name.split(".")
|
| 880 |
+
direct_name = ".".join(names[:-1])
|
| 881 |
+
if direct_name in direct_dict:
|
| 882 |
+
name_ = direct_dict[direct_name] + "." + names[-1]
|
| 883 |
+
state_dict_[name_] = param
|
| 884 |
+
elif names[0] == "double_blocks":
|
| 885 |
+
prefix = ".".join(names[:2])
|
| 886 |
+
suffix = ".".join(names[2:-1])
|
| 887 |
+
name_ = prefix + "." + double_suffix_dict[suffix] + "." + names[-1]
|
| 888 |
+
state_dict_[name_] = param
|
| 889 |
+
elif names[0] == "single_blocks":
|
| 890 |
+
prefix = ".".join(names[:2])
|
| 891 |
+
suffix = ".".join(names[2:-1])
|
| 892 |
+
if isinstance(single_suffix_dict[suffix], list):
|
| 893 |
+
if suffix == "linear1":
|
| 894 |
+
name_a, name_b = single_suffix_dict[suffix]
|
| 895 |
+
param_a, param_b = torch.split(param, (3072*3, 3072*4), dim=0)
|
| 896 |
+
state_dict_[prefix + "." + name_a + "." + names[-1]] = param_a
|
| 897 |
+
state_dict_[prefix + "." + name_b + "." + names[-1]] = param_b
|
| 898 |
+
elif suffix == "linear2":
|
| 899 |
+
if names[-1] == "weight":
|
| 900 |
+
name_a, name_b = single_suffix_dict[suffix]
|
| 901 |
+
param_a, param_b = torch.split(param, (3072*1, 3072*4), dim=-1)
|
| 902 |
+
state_dict_[prefix + "." + name_a + "." + names[-1]] = param_a
|
| 903 |
+
state_dict_[prefix + "." + name_b + "." + names[-1]] = param_b
|
| 904 |
+
else:
|
| 905 |
+
name_a, name_b = single_suffix_dict[suffix]
|
| 906 |
+
state_dict_[prefix + "." + name_a + "." + names[-1]] = param
|
| 907 |
+
else:
|
| 908 |
+
pass
|
| 909 |
+
else:
|
| 910 |
+
name_ = prefix + "." + single_suffix_dict[suffix] + "." + names[-1]
|
| 911 |
+
state_dict_[name_] = param
|
| 912 |
+
elif names[0] == "txt_in":
|
| 913 |
+
prefix = ".".join(names[:4]).replace(".individual_token_refiner.", ".")
|
| 914 |
+
suffix = ".".join(names[4:-1])
|
| 915 |
+
name_ = prefix + "." + txt_suffix_dict[suffix] + "." + names[-1]
|
| 916 |
+
state_dict_[name_] = param
|
| 917 |
+
else:
|
| 918 |
+
pass
|
| 919 |
+
|
| 920 |
+
return state_dict_
|
diffsynth/models/hunyuan_video_text_encoder.py
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from transformers import LlamaModel, LlamaConfig, DynamicCache, LlavaForConditionalGeneration
|
| 2 |
+
from copy import deepcopy
|
| 3 |
+
import torch
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class HunyuanVideoLLMEncoder(LlamaModel):
|
| 7 |
+
|
| 8 |
+
def __init__(self, config: LlamaConfig):
|
| 9 |
+
super().__init__(config)
|
| 10 |
+
self.auto_offload = False
|
| 11 |
+
|
| 12 |
+
def enable_auto_offload(self, **kwargs):
|
| 13 |
+
self.auto_offload = True
|
| 14 |
+
|
| 15 |
+
def forward(self, input_ids, attention_mask, hidden_state_skip_layer=2):
|
| 16 |
+
embed_tokens = deepcopy(self.embed_tokens).to(input_ids.device) if self.auto_offload else self.embed_tokens
|
| 17 |
+
inputs_embeds = embed_tokens(input_ids)
|
| 18 |
+
|
| 19 |
+
past_key_values = DynamicCache()
|
| 20 |
+
|
| 21 |
+
cache_position = torch.arange(0, inputs_embeds.shape[1], device=inputs_embeds.device)
|
| 22 |
+
position_ids = cache_position.unsqueeze(0)
|
| 23 |
+
|
| 24 |
+
causal_mask = self._update_causal_mask(attention_mask, inputs_embeds, cache_position, None, False)
|
| 25 |
+
hidden_states = inputs_embeds
|
| 26 |
+
|
| 27 |
+
# create position embeddings to be shared across the decoder layers
|
| 28 |
+
rotary_emb = deepcopy(self.rotary_emb).to(input_ids.device) if self.auto_offload else self.rotary_emb
|
| 29 |
+
position_embeddings = rotary_emb(hidden_states, position_ids)
|
| 30 |
+
|
| 31 |
+
# decoder layers
|
| 32 |
+
for layer_id, decoder_layer in enumerate(self.layers):
|
| 33 |
+
if self.auto_offload:
|
| 34 |
+
decoder_layer = deepcopy(decoder_layer).to(hidden_states.device)
|
| 35 |
+
layer_outputs = decoder_layer(
|
| 36 |
+
hidden_states,
|
| 37 |
+
attention_mask=causal_mask,
|
| 38 |
+
position_ids=position_ids,
|
| 39 |
+
past_key_value=past_key_values,
|
| 40 |
+
output_attentions=False,
|
| 41 |
+
use_cache=True,
|
| 42 |
+
cache_position=cache_position,
|
| 43 |
+
position_embeddings=position_embeddings,
|
| 44 |
+
)
|
| 45 |
+
hidden_states = layer_outputs[0]
|
| 46 |
+
if layer_id + hidden_state_skip_layer + 1 >= len(self.layers):
|
| 47 |
+
break
|
| 48 |
+
|
| 49 |
+
return hidden_states
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class HunyuanVideoMLLMEncoder(LlavaForConditionalGeneration):
|
| 53 |
+
|
| 54 |
+
def __init__(self, config):
|
| 55 |
+
super().__init__(config)
|
| 56 |
+
self.auto_offload = False
|
| 57 |
+
|
| 58 |
+
def enable_auto_offload(self, **kwargs):
|
| 59 |
+
self.auto_offload = True
|
| 60 |
+
|
| 61 |
+
# TODO: implement the low VRAM inference for MLLM.
|
| 62 |
+
def forward(self, input_ids, pixel_values, attention_mask, hidden_state_skip_layer=2):
|
| 63 |
+
outputs = super().forward(input_ids=input_ids,
|
| 64 |
+
attention_mask=attention_mask,
|
| 65 |
+
output_hidden_states=True,
|
| 66 |
+
pixel_values=pixel_values)
|
| 67 |
+
hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
| 68 |
+
return hidden_state
|
diffsynth/models/hunyuan_video_vae_decoder.py
ADDED
|
@@ -0,0 +1,507 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from einops import rearrange
|
| 5 |
+
import numpy as np
|
| 6 |
+
from tqdm import tqdm
|
| 7 |
+
from einops import repeat
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class CausalConv3d(nn.Module):
|
| 11 |
+
|
| 12 |
+
def __init__(self, in_channel, out_channel, kernel_size, stride=1, dilation=1, pad_mode='replicate', **kwargs):
|
| 13 |
+
super().__init__()
|
| 14 |
+
self.pad_mode = pad_mode
|
| 15 |
+
self.time_causal_padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0
|
| 16 |
+
) # W, H, T
|
| 17 |
+
self.conv = nn.Conv3d(in_channel, out_channel, kernel_size, stride=stride, dilation=dilation, **kwargs)
|
| 18 |
+
|
| 19 |
+
def forward(self, x):
|
| 20 |
+
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
| 21 |
+
return self.conv(x)
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class UpsampleCausal3D(nn.Module):
|
| 25 |
+
|
| 26 |
+
def __init__(self, channels, use_conv=False, out_channels=None, kernel_size=None, bias=True, upsample_factor=(2, 2, 2)):
|
| 27 |
+
super().__init__()
|
| 28 |
+
self.channels = channels
|
| 29 |
+
self.out_channels = out_channels or channels
|
| 30 |
+
self.upsample_factor = upsample_factor
|
| 31 |
+
self.conv = None
|
| 32 |
+
if use_conv:
|
| 33 |
+
kernel_size = 3 if kernel_size is None else kernel_size
|
| 34 |
+
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, bias=bias)
|
| 35 |
+
|
| 36 |
+
def forward(self, hidden_states):
|
| 37 |
+
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
| 38 |
+
dtype = hidden_states.dtype
|
| 39 |
+
if dtype == torch.bfloat16:
|
| 40 |
+
hidden_states = hidden_states.to(torch.float32)
|
| 41 |
+
|
| 42 |
+
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
| 43 |
+
if hidden_states.shape[0] >= 64:
|
| 44 |
+
hidden_states = hidden_states.contiguous()
|
| 45 |
+
|
| 46 |
+
# interpolate
|
| 47 |
+
B, C, T, H, W = hidden_states.shape
|
| 48 |
+
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
|
| 49 |
+
if T > 1:
|
| 50 |
+
other_h = F.interpolate(other_h, scale_factor=self.upsample_factor, mode="nearest")
|
| 51 |
+
first_h = F.interpolate(first_h.squeeze(2), scale_factor=self.upsample_factor[1:], mode="nearest").unsqueeze(2)
|
| 52 |
+
hidden_states = torch.cat((first_h, other_h), dim=2) if T > 1 else first_h
|
| 53 |
+
|
| 54 |
+
# If the input is bfloat16, we cast back to bfloat16
|
| 55 |
+
if dtype == torch.bfloat16:
|
| 56 |
+
hidden_states = hidden_states.to(dtype)
|
| 57 |
+
|
| 58 |
+
if self.conv:
|
| 59 |
+
hidden_states = self.conv(hidden_states)
|
| 60 |
+
|
| 61 |
+
return hidden_states
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
class ResnetBlockCausal3D(nn.Module):
|
| 65 |
+
|
| 66 |
+
def __init__(self, in_channels, out_channels=None, dropout=0.0, groups=32, eps=1e-6, conv_shortcut_bias=True):
|
| 67 |
+
super().__init__()
|
| 68 |
+
self.pre_norm = True
|
| 69 |
+
self.in_channels = in_channels
|
| 70 |
+
out_channels = in_channels if out_channels is None else out_channels
|
| 71 |
+
self.out_channels = out_channels
|
| 72 |
+
|
| 73 |
+
self.norm1 = nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
| 74 |
+
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
|
| 75 |
+
|
| 76 |
+
self.norm2 = nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
|
| 77 |
+
self.conv2 = CausalConv3d(out_channels, out_channels, kernel_size=3, stride=1)
|
| 78 |
+
|
| 79 |
+
self.dropout = nn.Dropout(dropout)
|
| 80 |
+
self.nonlinearity = nn.SiLU()
|
| 81 |
+
|
| 82 |
+
self.conv_shortcut = None
|
| 83 |
+
if in_channels != out_channels:
|
| 84 |
+
self.conv_shortcut = CausalConv3d(in_channels, out_channels, kernel_size=1, stride=1, bias=conv_shortcut_bias)
|
| 85 |
+
|
| 86 |
+
def forward(self, input_tensor):
|
| 87 |
+
hidden_states = input_tensor
|
| 88 |
+
# conv1
|
| 89 |
+
hidden_states = self.norm1(hidden_states)
|
| 90 |
+
hidden_states = self.nonlinearity(hidden_states)
|
| 91 |
+
hidden_states = self.conv1(hidden_states)
|
| 92 |
+
|
| 93 |
+
# conv2
|
| 94 |
+
hidden_states = self.norm2(hidden_states)
|
| 95 |
+
hidden_states = self.nonlinearity(hidden_states)
|
| 96 |
+
hidden_states = self.dropout(hidden_states)
|
| 97 |
+
hidden_states = self.conv2(hidden_states)
|
| 98 |
+
# shortcut
|
| 99 |
+
if self.conv_shortcut is not None:
|
| 100 |
+
input_tensor = (self.conv_shortcut(input_tensor))
|
| 101 |
+
# shortcut and scale
|
| 102 |
+
output_tensor = input_tensor + hidden_states
|
| 103 |
+
|
| 104 |
+
return output_tensor
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
def prepare_causal_attention_mask(n_frame, n_hw, dtype, device, batch_size=None):
|
| 108 |
+
seq_len = n_frame * n_hw
|
| 109 |
+
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
|
| 110 |
+
for i in range(seq_len):
|
| 111 |
+
i_frame = i // n_hw
|
| 112 |
+
mask[i, :(i_frame + 1) * n_hw] = 0
|
| 113 |
+
if batch_size is not None:
|
| 114 |
+
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
| 115 |
+
return mask
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
class Attention(nn.Module):
|
| 119 |
+
|
| 120 |
+
def __init__(self,
|
| 121 |
+
in_channels,
|
| 122 |
+
num_heads,
|
| 123 |
+
head_dim,
|
| 124 |
+
num_groups=32,
|
| 125 |
+
dropout=0.0,
|
| 126 |
+
eps=1e-6,
|
| 127 |
+
bias=True,
|
| 128 |
+
residual_connection=True):
|
| 129 |
+
super().__init__()
|
| 130 |
+
self.num_heads = num_heads
|
| 131 |
+
self.head_dim = head_dim
|
| 132 |
+
self.residual_connection = residual_connection
|
| 133 |
+
dim_inner = head_dim * num_heads
|
| 134 |
+
self.group_norm = nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=eps, affine=True)
|
| 135 |
+
self.to_q = nn.Linear(in_channels, dim_inner, bias=bias)
|
| 136 |
+
self.to_k = nn.Linear(in_channels, dim_inner, bias=bias)
|
| 137 |
+
self.to_v = nn.Linear(in_channels, dim_inner, bias=bias)
|
| 138 |
+
self.to_out = nn.Sequential(nn.Linear(dim_inner, in_channels, bias=bias), nn.Dropout(dropout))
|
| 139 |
+
|
| 140 |
+
def forward(self, input_tensor, attn_mask=None):
|
| 141 |
+
hidden_states = self.group_norm(input_tensor.transpose(1, 2)).transpose(1, 2)
|
| 142 |
+
batch_size = hidden_states.shape[0]
|
| 143 |
+
|
| 144 |
+
q = self.to_q(hidden_states)
|
| 145 |
+
k = self.to_k(hidden_states)
|
| 146 |
+
v = self.to_v(hidden_states)
|
| 147 |
+
|
| 148 |
+
q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 149 |
+
k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 150 |
+
v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
|
| 151 |
+
|
| 152 |
+
if attn_mask is not None:
|
| 153 |
+
attn_mask = attn_mask.view(batch_size, self.num_heads, -1, attn_mask.shape[-1])
|
| 154 |
+
hidden_states = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 155 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 156 |
+
hidden_states = self.to_out(hidden_states)
|
| 157 |
+
if self.residual_connection:
|
| 158 |
+
output_tensor = input_tensor + hidden_states
|
| 159 |
+
return output_tensor
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
class UNetMidBlockCausal3D(nn.Module):
|
| 163 |
+
|
| 164 |
+
def __init__(self, in_channels, dropout=0.0, num_layers=1, eps=1e-6, num_groups=32, attention_head_dim=None):
|
| 165 |
+
super().__init__()
|
| 166 |
+
resnets = [
|
| 167 |
+
ResnetBlockCausal3D(
|
| 168 |
+
in_channels=in_channels,
|
| 169 |
+
out_channels=in_channels,
|
| 170 |
+
dropout=dropout,
|
| 171 |
+
groups=num_groups,
|
| 172 |
+
eps=eps,
|
| 173 |
+
)
|
| 174 |
+
]
|
| 175 |
+
attentions = []
|
| 176 |
+
attention_head_dim = attention_head_dim or in_channels
|
| 177 |
+
|
| 178 |
+
for _ in range(num_layers):
|
| 179 |
+
attentions.append(
|
| 180 |
+
Attention(
|
| 181 |
+
in_channels,
|
| 182 |
+
num_heads=in_channels // attention_head_dim,
|
| 183 |
+
head_dim=attention_head_dim,
|
| 184 |
+
num_groups=num_groups,
|
| 185 |
+
dropout=dropout,
|
| 186 |
+
eps=eps,
|
| 187 |
+
bias=True,
|
| 188 |
+
residual_connection=True,
|
| 189 |
+
))
|
| 190 |
+
|
| 191 |
+
resnets.append(
|
| 192 |
+
ResnetBlockCausal3D(
|
| 193 |
+
in_channels=in_channels,
|
| 194 |
+
out_channels=in_channels,
|
| 195 |
+
dropout=dropout,
|
| 196 |
+
groups=num_groups,
|
| 197 |
+
eps=eps,
|
| 198 |
+
))
|
| 199 |
+
|
| 200 |
+
self.attentions = nn.ModuleList(attentions)
|
| 201 |
+
self.resnets = nn.ModuleList(resnets)
|
| 202 |
+
|
| 203 |
+
def forward(self, hidden_states):
|
| 204 |
+
hidden_states = self.resnets[0](hidden_states)
|
| 205 |
+
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
| 206 |
+
B, C, T, H, W = hidden_states.shape
|
| 207 |
+
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
|
| 208 |
+
attn_mask = prepare_causal_attention_mask(T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B)
|
| 209 |
+
hidden_states = attn(hidden_states, attn_mask=attn_mask)
|
| 210 |
+
hidden_states = rearrange(hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W)
|
| 211 |
+
hidden_states = resnet(hidden_states)
|
| 212 |
+
|
| 213 |
+
return hidden_states
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
class UpDecoderBlockCausal3D(nn.Module):
|
| 217 |
+
|
| 218 |
+
def __init__(
|
| 219 |
+
self,
|
| 220 |
+
in_channels,
|
| 221 |
+
out_channels,
|
| 222 |
+
dropout=0.0,
|
| 223 |
+
num_layers=1,
|
| 224 |
+
eps=1e-6,
|
| 225 |
+
num_groups=32,
|
| 226 |
+
add_upsample=True,
|
| 227 |
+
upsample_scale_factor=(2, 2, 2),
|
| 228 |
+
):
|
| 229 |
+
super().__init__()
|
| 230 |
+
resnets = []
|
| 231 |
+
for i in range(num_layers):
|
| 232 |
+
cur_in_channel = in_channels if i == 0 else out_channels
|
| 233 |
+
resnets.append(
|
| 234 |
+
ResnetBlockCausal3D(
|
| 235 |
+
in_channels=cur_in_channel,
|
| 236 |
+
out_channels=out_channels,
|
| 237 |
+
groups=num_groups,
|
| 238 |
+
dropout=dropout,
|
| 239 |
+
eps=eps,
|
| 240 |
+
))
|
| 241 |
+
self.resnets = nn.ModuleList(resnets)
|
| 242 |
+
|
| 243 |
+
self.upsamplers = None
|
| 244 |
+
if add_upsample:
|
| 245 |
+
self.upsamplers = nn.ModuleList([
|
| 246 |
+
UpsampleCausal3D(
|
| 247 |
+
out_channels,
|
| 248 |
+
use_conv=True,
|
| 249 |
+
out_channels=out_channels,
|
| 250 |
+
upsample_factor=upsample_scale_factor,
|
| 251 |
+
)
|
| 252 |
+
])
|
| 253 |
+
|
| 254 |
+
def forward(self, hidden_states):
|
| 255 |
+
for resnet in self.resnets:
|
| 256 |
+
hidden_states = resnet(hidden_states)
|
| 257 |
+
if self.upsamplers is not None:
|
| 258 |
+
for upsampler in self.upsamplers:
|
| 259 |
+
hidden_states = upsampler(hidden_states)
|
| 260 |
+
return hidden_states
|
| 261 |
+
|
| 262 |
+
|
| 263 |
+
class DecoderCausal3D(nn.Module):
|
| 264 |
+
|
| 265 |
+
def __init__(
|
| 266 |
+
self,
|
| 267 |
+
in_channels=16,
|
| 268 |
+
out_channels=3,
|
| 269 |
+
eps=1e-6,
|
| 270 |
+
dropout=0.0,
|
| 271 |
+
block_out_channels=[128, 256, 512, 512],
|
| 272 |
+
layers_per_block=2,
|
| 273 |
+
num_groups=32,
|
| 274 |
+
time_compression_ratio=4,
|
| 275 |
+
spatial_compression_ratio=8,
|
| 276 |
+
gradient_checkpointing=False,
|
| 277 |
+
):
|
| 278 |
+
super().__init__()
|
| 279 |
+
self.layers_per_block = layers_per_block
|
| 280 |
+
|
| 281 |
+
self.conv_in = CausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
|
| 282 |
+
self.up_blocks = nn.ModuleList([])
|
| 283 |
+
|
| 284 |
+
# mid
|
| 285 |
+
self.mid_block = UNetMidBlockCausal3D(
|
| 286 |
+
in_channels=block_out_channels[-1],
|
| 287 |
+
dropout=dropout,
|
| 288 |
+
eps=eps,
|
| 289 |
+
num_groups=num_groups,
|
| 290 |
+
attention_head_dim=block_out_channels[-1],
|
| 291 |
+
)
|
| 292 |
+
|
| 293 |
+
# up
|
| 294 |
+
reversed_block_out_channels = list(reversed(block_out_channels))
|
| 295 |
+
output_channel = reversed_block_out_channels[0]
|
| 296 |
+
for i in range(len(block_out_channels)):
|
| 297 |
+
prev_output_channel = output_channel
|
| 298 |
+
output_channel = reversed_block_out_channels[i]
|
| 299 |
+
is_final_block = i == len(block_out_channels) - 1
|
| 300 |
+
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
|
| 301 |
+
num_time_upsample_layers = int(np.log2(time_compression_ratio))
|
| 302 |
+
|
| 303 |
+
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
|
| 304 |
+
add_time_upsample = bool(i >= len(block_out_channels) - 1 - num_time_upsample_layers and not is_final_block)
|
| 305 |
+
|
| 306 |
+
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
|
| 307 |
+
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
|
| 308 |
+
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
|
| 309 |
+
|
| 310 |
+
up_block = UpDecoderBlockCausal3D(
|
| 311 |
+
in_channels=prev_output_channel,
|
| 312 |
+
out_channels=output_channel,
|
| 313 |
+
dropout=dropout,
|
| 314 |
+
num_layers=layers_per_block + 1,
|
| 315 |
+
eps=eps,
|
| 316 |
+
num_groups=num_groups,
|
| 317 |
+
add_upsample=bool(add_spatial_upsample or add_time_upsample),
|
| 318 |
+
upsample_scale_factor=upsample_scale_factor,
|
| 319 |
+
)
|
| 320 |
+
|
| 321 |
+
self.up_blocks.append(up_block)
|
| 322 |
+
prev_output_channel = output_channel
|
| 323 |
+
|
| 324 |
+
# out
|
| 325 |
+
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=num_groups, eps=eps)
|
| 326 |
+
self.conv_act = nn.SiLU()
|
| 327 |
+
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
|
| 328 |
+
|
| 329 |
+
self.gradient_checkpointing = gradient_checkpointing
|
| 330 |
+
|
| 331 |
+
def forward(self, hidden_states):
|
| 332 |
+
hidden_states = self.conv_in(hidden_states)
|
| 333 |
+
if self.training and self.gradient_checkpointing:
|
| 334 |
+
|
| 335 |
+
def create_custom_forward(module):
|
| 336 |
+
|
| 337 |
+
def custom_forward(*inputs):
|
| 338 |
+
return module(*inputs)
|
| 339 |
+
|
| 340 |
+
return custom_forward
|
| 341 |
+
|
| 342 |
+
# middle
|
| 343 |
+
hidden_states = torch.utils.checkpoint.checkpoint(
|
| 344 |
+
create_custom_forward(self.mid_block),
|
| 345 |
+
hidden_states,
|
| 346 |
+
use_reentrant=False,
|
| 347 |
+
)
|
| 348 |
+
# up
|
| 349 |
+
for up_block in self.up_blocks:
|
| 350 |
+
hidden_states = torch.utils.checkpoint.checkpoint(
|
| 351 |
+
create_custom_forward(up_block),
|
| 352 |
+
hidden_states,
|
| 353 |
+
use_reentrant=False,
|
| 354 |
+
)
|
| 355 |
+
else:
|
| 356 |
+
# middle
|
| 357 |
+
hidden_states = self.mid_block(hidden_states)
|
| 358 |
+
# up
|
| 359 |
+
for up_block in self.up_blocks:
|
| 360 |
+
hidden_states = up_block(hidden_states)
|
| 361 |
+
# post-process
|
| 362 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 363 |
+
hidden_states = self.conv_act(hidden_states)
|
| 364 |
+
hidden_states = self.conv_out(hidden_states)
|
| 365 |
+
|
| 366 |
+
return hidden_states
|
| 367 |
+
|
| 368 |
+
|
| 369 |
+
class HunyuanVideoVAEDecoder(nn.Module):
|
| 370 |
+
|
| 371 |
+
def __init__(
|
| 372 |
+
self,
|
| 373 |
+
in_channels=16,
|
| 374 |
+
out_channels=3,
|
| 375 |
+
eps=1e-6,
|
| 376 |
+
dropout=0.0,
|
| 377 |
+
block_out_channels=[128, 256, 512, 512],
|
| 378 |
+
layers_per_block=2,
|
| 379 |
+
num_groups=32,
|
| 380 |
+
time_compression_ratio=4,
|
| 381 |
+
spatial_compression_ratio=8,
|
| 382 |
+
gradient_checkpointing=False,
|
| 383 |
+
):
|
| 384 |
+
super().__init__()
|
| 385 |
+
self.decoder = DecoderCausal3D(
|
| 386 |
+
in_channels=in_channels,
|
| 387 |
+
out_channels=out_channels,
|
| 388 |
+
eps=eps,
|
| 389 |
+
dropout=dropout,
|
| 390 |
+
block_out_channels=block_out_channels,
|
| 391 |
+
layers_per_block=layers_per_block,
|
| 392 |
+
num_groups=num_groups,
|
| 393 |
+
time_compression_ratio=time_compression_ratio,
|
| 394 |
+
spatial_compression_ratio=spatial_compression_ratio,
|
| 395 |
+
gradient_checkpointing=gradient_checkpointing,
|
| 396 |
+
)
|
| 397 |
+
self.post_quant_conv = nn.Conv3d(in_channels, in_channels, kernel_size=1)
|
| 398 |
+
self.scaling_factor = 0.476986
|
| 399 |
+
|
| 400 |
+
|
| 401 |
+
def forward(self, latents):
|
| 402 |
+
latents = latents / self.scaling_factor
|
| 403 |
+
latents = self.post_quant_conv(latents)
|
| 404 |
+
dec = self.decoder(latents)
|
| 405 |
+
return dec
|
| 406 |
+
|
| 407 |
+
|
| 408 |
+
def build_1d_mask(self, length, left_bound, right_bound, border_width):
|
| 409 |
+
x = torch.ones((length,))
|
| 410 |
+
if not left_bound:
|
| 411 |
+
x[:border_width] = (torch.arange(border_width) + 1) / border_width
|
| 412 |
+
if not right_bound:
|
| 413 |
+
x[-border_width:] = torch.flip((torch.arange(border_width) + 1) / border_width, dims=(0,))
|
| 414 |
+
return x
|
| 415 |
+
|
| 416 |
+
|
| 417 |
+
def build_mask(self, data, is_bound, border_width):
|
| 418 |
+
_, _, T, H, W = data.shape
|
| 419 |
+
t = self.build_1d_mask(T, is_bound[0], is_bound[1], border_width[0])
|
| 420 |
+
h = self.build_1d_mask(H, is_bound[2], is_bound[3], border_width[1])
|
| 421 |
+
w = self.build_1d_mask(W, is_bound[4], is_bound[5], border_width[2])
|
| 422 |
+
|
| 423 |
+
t = repeat(t, "T -> T H W", T=T, H=H, W=W)
|
| 424 |
+
h = repeat(h, "H -> T H W", T=T, H=H, W=W)
|
| 425 |
+
w = repeat(w, "W -> T H W", T=T, H=H, W=W)
|
| 426 |
+
|
| 427 |
+
mask = torch.stack([t, h, w]).min(dim=0).values
|
| 428 |
+
mask = rearrange(mask, "T H W -> 1 1 T H W")
|
| 429 |
+
return mask
|
| 430 |
+
|
| 431 |
+
|
| 432 |
+
def tile_forward(self, hidden_states, tile_size, tile_stride):
|
| 433 |
+
B, C, T, H, W = hidden_states.shape
|
| 434 |
+
size_t, size_h, size_w = tile_size
|
| 435 |
+
stride_t, stride_h, stride_w = tile_stride
|
| 436 |
+
|
| 437 |
+
# Split tasks
|
| 438 |
+
tasks = []
|
| 439 |
+
for t in range(0, T, stride_t):
|
| 440 |
+
if (t-stride_t >= 0 and t-stride_t+size_t >= T): continue
|
| 441 |
+
for h in range(0, H, stride_h):
|
| 442 |
+
if (h-stride_h >= 0 and h-stride_h+size_h >= H): continue
|
| 443 |
+
for w in range(0, W, stride_w):
|
| 444 |
+
if (w-stride_w >= 0 and w-stride_w+size_w >= W): continue
|
| 445 |
+
t_, h_, w_ = t + size_t, h + size_h, w + size_w
|
| 446 |
+
tasks.append((t, t_, h, h_, w, w_))
|
| 447 |
+
|
| 448 |
+
# Run
|
| 449 |
+
torch_dtype = self.post_quant_conv.weight.dtype
|
| 450 |
+
data_device = hidden_states.device
|
| 451 |
+
computation_device = self.post_quant_conv.weight.device
|
| 452 |
+
|
| 453 |
+
weight = torch.zeros((1, 1, (T - 1) * 4 + 1, H * 8, W * 8), dtype=torch_dtype, device=data_device)
|
| 454 |
+
values = torch.zeros((B, 3, (T - 1) * 4 + 1, H * 8, W * 8), dtype=torch_dtype, device=data_device)
|
| 455 |
+
|
| 456 |
+
for t, t_, h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
|
| 457 |
+
hidden_states_batch = hidden_states[:, :, t:t_, h:h_, w:w_].to(computation_device)
|
| 458 |
+
hidden_states_batch = self.forward(hidden_states_batch).to(data_device)
|
| 459 |
+
if t > 0:
|
| 460 |
+
hidden_states_batch = hidden_states_batch[:, :, 1:]
|
| 461 |
+
|
| 462 |
+
mask = self.build_mask(
|
| 463 |
+
hidden_states_batch,
|
| 464 |
+
is_bound=(t==0, t_>=T, h==0, h_>=H, w==0, w_>=W),
|
| 465 |
+
border_width=((size_t - stride_t) * 4, (size_h - stride_h) * 8, (size_w - stride_w) * 8)
|
| 466 |
+
).to(dtype=torch_dtype, device=data_device)
|
| 467 |
+
|
| 468 |
+
target_t = 0 if t==0 else t * 4 + 1
|
| 469 |
+
target_h = h * 8
|
| 470 |
+
target_w = w * 8
|
| 471 |
+
values[
|
| 472 |
+
:,
|
| 473 |
+
:,
|
| 474 |
+
target_t: target_t + hidden_states_batch.shape[2],
|
| 475 |
+
target_h: target_h + hidden_states_batch.shape[3],
|
| 476 |
+
target_w: target_w + hidden_states_batch.shape[4],
|
| 477 |
+
] += hidden_states_batch * mask
|
| 478 |
+
weight[
|
| 479 |
+
:,
|
| 480 |
+
:,
|
| 481 |
+
target_t: target_t + hidden_states_batch.shape[2],
|
| 482 |
+
target_h: target_h + hidden_states_batch.shape[3],
|
| 483 |
+
target_w: target_w + hidden_states_batch.shape[4],
|
| 484 |
+
] += mask
|
| 485 |
+
return values / weight
|
| 486 |
+
|
| 487 |
+
|
| 488 |
+
def decode_video(self, latents, tile_size=(17, 32, 32), tile_stride=(12, 24, 24)):
|
| 489 |
+
latents = latents.to(self.post_quant_conv.weight.dtype)
|
| 490 |
+
return self.tile_forward(latents, tile_size=tile_size, tile_stride=tile_stride)
|
| 491 |
+
|
| 492 |
+
@staticmethod
|
| 493 |
+
def state_dict_converter():
|
| 494 |
+
return HunyuanVideoVAEDecoderStateDictConverter()
|
| 495 |
+
|
| 496 |
+
|
| 497 |
+
class HunyuanVideoVAEDecoderStateDictConverter:
|
| 498 |
+
|
| 499 |
+
def __init__(self):
|
| 500 |
+
pass
|
| 501 |
+
|
| 502 |
+
def from_diffusers(self, state_dict):
|
| 503 |
+
state_dict_ = {}
|
| 504 |
+
for name in state_dict:
|
| 505 |
+
if name.startswith('decoder.') or name.startswith('post_quant_conv.'):
|
| 506 |
+
state_dict_[name] = state_dict[name]
|
| 507 |
+
return state_dict_
|
diffsynth/models/hunyuan_video_vae_encoder.py
ADDED
|
@@ -0,0 +1,307 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from einops import rearrange, repeat
|
| 5 |
+
import numpy as np
|
| 6 |
+
from tqdm import tqdm
|
| 7 |
+
from .hunyuan_video_vae_decoder import CausalConv3d, ResnetBlockCausal3D, UNetMidBlockCausal3D
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class DownsampleCausal3D(nn.Module):
|
| 11 |
+
|
| 12 |
+
def __init__(self, channels, out_channels, kernel_size=3, bias=True, stride=2):
|
| 13 |
+
super().__init__()
|
| 14 |
+
self.conv = CausalConv3d(channels, out_channels, kernel_size, stride=stride, bias=bias)
|
| 15 |
+
|
| 16 |
+
def forward(self, hidden_states):
|
| 17 |
+
hidden_states = self.conv(hidden_states)
|
| 18 |
+
return hidden_states
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class DownEncoderBlockCausal3D(nn.Module):
|
| 22 |
+
|
| 23 |
+
def __init__(
|
| 24 |
+
self,
|
| 25 |
+
in_channels,
|
| 26 |
+
out_channels,
|
| 27 |
+
dropout=0.0,
|
| 28 |
+
num_layers=1,
|
| 29 |
+
eps=1e-6,
|
| 30 |
+
num_groups=32,
|
| 31 |
+
add_downsample=True,
|
| 32 |
+
downsample_stride=2,
|
| 33 |
+
):
|
| 34 |
+
|
| 35 |
+
super().__init__()
|
| 36 |
+
resnets = []
|
| 37 |
+
for i in range(num_layers):
|
| 38 |
+
cur_in_channel = in_channels if i == 0 else out_channels
|
| 39 |
+
resnets.append(
|
| 40 |
+
ResnetBlockCausal3D(
|
| 41 |
+
in_channels=cur_in_channel,
|
| 42 |
+
out_channels=out_channels,
|
| 43 |
+
groups=num_groups,
|
| 44 |
+
dropout=dropout,
|
| 45 |
+
eps=eps,
|
| 46 |
+
))
|
| 47 |
+
self.resnets = nn.ModuleList(resnets)
|
| 48 |
+
|
| 49 |
+
self.downsamplers = None
|
| 50 |
+
if add_downsample:
|
| 51 |
+
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
|
| 52 |
+
out_channels,
|
| 53 |
+
out_channels,
|
| 54 |
+
stride=downsample_stride,
|
| 55 |
+
)])
|
| 56 |
+
|
| 57 |
+
def forward(self, hidden_states):
|
| 58 |
+
for resnet in self.resnets:
|
| 59 |
+
hidden_states = resnet(hidden_states)
|
| 60 |
+
|
| 61 |
+
if self.downsamplers is not None:
|
| 62 |
+
for downsampler in self.downsamplers:
|
| 63 |
+
hidden_states = downsampler(hidden_states)
|
| 64 |
+
|
| 65 |
+
return hidden_states
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
class EncoderCausal3D(nn.Module):
|
| 69 |
+
|
| 70 |
+
def __init__(
|
| 71 |
+
self,
|
| 72 |
+
in_channels: int = 3,
|
| 73 |
+
out_channels: int = 16,
|
| 74 |
+
eps=1e-6,
|
| 75 |
+
dropout=0.0,
|
| 76 |
+
block_out_channels=[128, 256, 512, 512],
|
| 77 |
+
layers_per_block=2,
|
| 78 |
+
num_groups=32,
|
| 79 |
+
time_compression_ratio: int = 4,
|
| 80 |
+
spatial_compression_ratio: int = 8,
|
| 81 |
+
gradient_checkpointing=False,
|
| 82 |
+
):
|
| 83 |
+
super().__init__()
|
| 84 |
+
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
| 85 |
+
self.down_blocks = nn.ModuleList([])
|
| 86 |
+
|
| 87 |
+
# down
|
| 88 |
+
output_channel = block_out_channels[0]
|
| 89 |
+
for i in range(len(block_out_channels)):
|
| 90 |
+
input_channel = output_channel
|
| 91 |
+
output_channel = block_out_channels[i]
|
| 92 |
+
is_final_block = i == len(block_out_channels) - 1
|
| 93 |
+
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
|
| 94 |
+
num_time_downsample_layers = int(np.log2(time_compression_ratio))
|
| 95 |
+
|
| 96 |
+
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
|
| 97 |
+
add_time_downsample = bool(i >= (len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
|
| 98 |
+
|
| 99 |
+
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
| 100 |
+
downsample_stride_T = (2,) if add_time_downsample else (1,)
|
| 101 |
+
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
|
| 102 |
+
down_block = DownEncoderBlockCausal3D(
|
| 103 |
+
in_channels=input_channel,
|
| 104 |
+
out_channels=output_channel,
|
| 105 |
+
dropout=dropout,
|
| 106 |
+
num_layers=layers_per_block,
|
| 107 |
+
eps=eps,
|
| 108 |
+
num_groups=num_groups,
|
| 109 |
+
add_downsample=bool(add_spatial_downsample or add_time_downsample),
|
| 110 |
+
downsample_stride=downsample_stride,
|
| 111 |
+
)
|
| 112 |
+
self.down_blocks.append(down_block)
|
| 113 |
+
|
| 114 |
+
# mid
|
| 115 |
+
self.mid_block = UNetMidBlockCausal3D(
|
| 116 |
+
in_channels=block_out_channels[-1],
|
| 117 |
+
dropout=dropout,
|
| 118 |
+
eps=eps,
|
| 119 |
+
num_groups=num_groups,
|
| 120 |
+
attention_head_dim=block_out_channels[-1],
|
| 121 |
+
)
|
| 122 |
+
# out
|
| 123 |
+
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=num_groups, eps=eps)
|
| 124 |
+
self.conv_act = nn.SiLU()
|
| 125 |
+
self.conv_out = CausalConv3d(block_out_channels[-1], 2 * out_channels, kernel_size=3)
|
| 126 |
+
|
| 127 |
+
self.gradient_checkpointing = gradient_checkpointing
|
| 128 |
+
|
| 129 |
+
def forward(self, hidden_states):
|
| 130 |
+
hidden_states = self.conv_in(hidden_states)
|
| 131 |
+
if self.training and self.gradient_checkpointing:
|
| 132 |
+
|
| 133 |
+
def create_custom_forward(module):
|
| 134 |
+
|
| 135 |
+
def custom_forward(*inputs):
|
| 136 |
+
return module(*inputs)
|
| 137 |
+
|
| 138 |
+
return custom_forward
|
| 139 |
+
|
| 140 |
+
# down
|
| 141 |
+
for down_block in self.down_blocks:
|
| 142 |
+
torch.utils.checkpoint.checkpoint(
|
| 143 |
+
create_custom_forward(down_block),
|
| 144 |
+
hidden_states,
|
| 145 |
+
use_reentrant=False,
|
| 146 |
+
)
|
| 147 |
+
# middle
|
| 148 |
+
hidden_states = torch.utils.checkpoint.checkpoint(
|
| 149 |
+
create_custom_forward(self.mid_block),
|
| 150 |
+
hidden_states,
|
| 151 |
+
use_reentrant=False,
|
| 152 |
+
)
|
| 153 |
+
else:
|
| 154 |
+
# down
|
| 155 |
+
for down_block in self.down_blocks:
|
| 156 |
+
hidden_states = down_block(hidden_states)
|
| 157 |
+
# middle
|
| 158 |
+
hidden_states = self.mid_block(hidden_states)
|
| 159 |
+
# post-process
|
| 160 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 161 |
+
hidden_states = self.conv_act(hidden_states)
|
| 162 |
+
hidden_states = self.conv_out(hidden_states)
|
| 163 |
+
|
| 164 |
+
return hidden_states
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class HunyuanVideoVAEEncoder(nn.Module):
|
| 168 |
+
|
| 169 |
+
def __init__(
|
| 170 |
+
self,
|
| 171 |
+
in_channels=3,
|
| 172 |
+
out_channels=16,
|
| 173 |
+
eps=1e-6,
|
| 174 |
+
dropout=0.0,
|
| 175 |
+
block_out_channels=[128, 256, 512, 512],
|
| 176 |
+
layers_per_block=2,
|
| 177 |
+
num_groups=32,
|
| 178 |
+
time_compression_ratio=4,
|
| 179 |
+
spatial_compression_ratio=8,
|
| 180 |
+
gradient_checkpointing=False,
|
| 181 |
+
):
|
| 182 |
+
super().__init__()
|
| 183 |
+
self.encoder = EncoderCausal3D(
|
| 184 |
+
in_channels=in_channels,
|
| 185 |
+
out_channels=out_channels,
|
| 186 |
+
eps=eps,
|
| 187 |
+
dropout=dropout,
|
| 188 |
+
block_out_channels=block_out_channels,
|
| 189 |
+
layers_per_block=layers_per_block,
|
| 190 |
+
num_groups=num_groups,
|
| 191 |
+
time_compression_ratio=time_compression_ratio,
|
| 192 |
+
spatial_compression_ratio=spatial_compression_ratio,
|
| 193 |
+
gradient_checkpointing=gradient_checkpointing,
|
| 194 |
+
)
|
| 195 |
+
self.quant_conv = nn.Conv3d(2 * out_channels, 2 * out_channels, kernel_size=1)
|
| 196 |
+
self.scaling_factor = 0.476986
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
def forward(self, images):
|
| 200 |
+
latents = self.encoder(images)
|
| 201 |
+
latents = self.quant_conv(latents)
|
| 202 |
+
latents = latents[:, :16]
|
| 203 |
+
latents = latents * self.scaling_factor
|
| 204 |
+
return latents
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def build_1d_mask(self, length, left_bound, right_bound, border_width):
|
| 208 |
+
x = torch.ones((length,))
|
| 209 |
+
if not left_bound:
|
| 210 |
+
x[:border_width] = (torch.arange(border_width) + 1) / border_width
|
| 211 |
+
if not right_bound:
|
| 212 |
+
x[-border_width:] = torch.flip((torch.arange(border_width) + 1) / border_width, dims=(0,))
|
| 213 |
+
return x
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def build_mask(self, data, is_bound, border_width):
|
| 217 |
+
_, _, T, H, W = data.shape
|
| 218 |
+
t = self.build_1d_mask(T, is_bound[0], is_bound[1], border_width[0])
|
| 219 |
+
h = self.build_1d_mask(H, is_bound[2], is_bound[3], border_width[1])
|
| 220 |
+
w = self.build_1d_mask(W, is_bound[4], is_bound[5], border_width[2])
|
| 221 |
+
|
| 222 |
+
t = repeat(t, "T -> T H W", T=T, H=H, W=W)
|
| 223 |
+
h = repeat(h, "H -> T H W", T=T, H=H, W=W)
|
| 224 |
+
w = repeat(w, "W -> T H W", T=T, H=H, W=W)
|
| 225 |
+
|
| 226 |
+
mask = torch.stack([t, h, w]).min(dim=0).values
|
| 227 |
+
mask = rearrange(mask, "T H W -> 1 1 T H W")
|
| 228 |
+
return mask
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
def tile_forward(self, hidden_states, tile_size, tile_stride):
|
| 232 |
+
B, C, T, H, W = hidden_states.shape
|
| 233 |
+
size_t, size_h, size_w = tile_size
|
| 234 |
+
stride_t, stride_h, stride_w = tile_stride
|
| 235 |
+
|
| 236 |
+
# Split tasks
|
| 237 |
+
tasks = []
|
| 238 |
+
for t in range(0, T, stride_t):
|
| 239 |
+
if (t-stride_t >= 0 and t-stride_t+size_t >= T): continue
|
| 240 |
+
for h in range(0, H, stride_h):
|
| 241 |
+
if (h-stride_h >= 0 and h-stride_h+size_h >= H): continue
|
| 242 |
+
for w in range(0, W, stride_w):
|
| 243 |
+
if (w-stride_w >= 0 and w-stride_w+size_w >= W): continue
|
| 244 |
+
t_, h_, w_ = t + size_t, h + size_h, w + size_w
|
| 245 |
+
tasks.append((t, t_, h, h_, w, w_))
|
| 246 |
+
|
| 247 |
+
# Run
|
| 248 |
+
torch_dtype = self.quant_conv.weight.dtype
|
| 249 |
+
data_device = hidden_states.device
|
| 250 |
+
computation_device = self.quant_conv.weight.device
|
| 251 |
+
|
| 252 |
+
weight = torch.zeros((1, 1, (T - 1) // 4 + 1, H // 8, W // 8), dtype=torch_dtype, device=data_device)
|
| 253 |
+
values = torch.zeros((B, 16, (T - 1) // 4 + 1, H // 8, W // 8), dtype=torch_dtype, device=data_device)
|
| 254 |
+
|
| 255 |
+
for t, t_, h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
|
| 256 |
+
hidden_states_batch = hidden_states[:, :, t:t_, h:h_, w:w_].to(computation_device)
|
| 257 |
+
hidden_states_batch = self.forward(hidden_states_batch).to(data_device)
|
| 258 |
+
if t > 0:
|
| 259 |
+
hidden_states_batch = hidden_states_batch[:, :, 1:]
|
| 260 |
+
|
| 261 |
+
mask = self.build_mask(
|
| 262 |
+
hidden_states_batch,
|
| 263 |
+
is_bound=(t==0, t_>=T, h==0, h_>=H, w==0, w_>=W),
|
| 264 |
+
border_width=((size_t - stride_t) // 4, (size_h - stride_h) // 8, (size_w - stride_w) // 8)
|
| 265 |
+
).to(dtype=torch_dtype, device=data_device)
|
| 266 |
+
|
| 267 |
+
target_t = 0 if t==0 else t // 4 + 1
|
| 268 |
+
target_h = h // 8
|
| 269 |
+
target_w = w // 8
|
| 270 |
+
values[
|
| 271 |
+
:,
|
| 272 |
+
:,
|
| 273 |
+
target_t: target_t + hidden_states_batch.shape[2],
|
| 274 |
+
target_h: target_h + hidden_states_batch.shape[3],
|
| 275 |
+
target_w: target_w + hidden_states_batch.shape[4],
|
| 276 |
+
] += hidden_states_batch * mask
|
| 277 |
+
weight[
|
| 278 |
+
:,
|
| 279 |
+
:,
|
| 280 |
+
target_t: target_t + hidden_states_batch.shape[2],
|
| 281 |
+
target_h: target_h + hidden_states_batch.shape[3],
|
| 282 |
+
target_w: target_w + hidden_states_batch.shape[4],
|
| 283 |
+
] += mask
|
| 284 |
+
return values / weight
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
def encode_video(self, latents, tile_size=(65, 256, 256), tile_stride=(48, 192, 192)):
|
| 288 |
+
latents = latents.to(self.quant_conv.weight.dtype)
|
| 289 |
+
return self.tile_forward(latents, tile_size=tile_size, tile_stride=tile_stride)
|
| 290 |
+
|
| 291 |
+
|
| 292 |
+
@staticmethod
|
| 293 |
+
def state_dict_converter():
|
| 294 |
+
return HunyuanVideoVAEEncoderStateDictConverter()
|
| 295 |
+
|
| 296 |
+
|
| 297 |
+
class HunyuanVideoVAEEncoderStateDictConverter:
|
| 298 |
+
|
| 299 |
+
def __init__(self):
|
| 300 |
+
pass
|
| 301 |
+
|
| 302 |
+
def from_diffusers(self, state_dict):
|
| 303 |
+
state_dict_ = {}
|
| 304 |
+
for name in state_dict:
|
| 305 |
+
if name.startswith('encoder.') or name.startswith('quant_conv.'):
|
| 306 |
+
state_dict_[name] = state_dict[name]
|
| 307 |
+
return state_dict_
|
diffsynth/models/kolors_text_encoder.py
ADDED
|
@@ -0,0 +1,1551 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
This model is copied from https://github.com/Kwai-Kolors/Kolors/tree/master/kolors/models.
|
| 3 |
+
We didn't modify this model.
|
| 4 |
+
The tensor operation is performed in the prompter.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
""" PyTorch ChatGLM model. """
|
| 9 |
+
|
| 10 |
+
import math
|
| 11 |
+
import copy
|
| 12 |
+
import warnings
|
| 13 |
+
import re
|
| 14 |
+
import sys
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
import torch.utils.checkpoint
|
| 18 |
+
import torch.nn.functional as F
|
| 19 |
+
from torch import nn
|
| 20 |
+
from torch.nn import CrossEntropyLoss, LayerNorm
|
| 21 |
+
from torch.nn import CrossEntropyLoss, LayerNorm, MSELoss, BCEWithLogitsLoss
|
| 22 |
+
from torch.nn.utils import skip_init
|
| 23 |
+
from typing import Optional, Tuple, Union, List, Callable, Dict, Any
|
| 24 |
+
from copy import deepcopy
|
| 25 |
+
|
| 26 |
+
from transformers.modeling_outputs import (
|
| 27 |
+
BaseModelOutputWithPast,
|
| 28 |
+
CausalLMOutputWithPast,
|
| 29 |
+
SequenceClassifierOutputWithPast,
|
| 30 |
+
)
|
| 31 |
+
from transformers.modeling_utils import PreTrainedModel
|
| 32 |
+
from transformers.utils import logging
|
| 33 |
+
from transformers.generation.logits_process import LogitsProcessor
|
| 34 |
+
from transformers.generation.utils import LogitsProcessorList, StoppingCriteriaList, GenerationConfig, ModelOutput
|
| 35 |
+
from transformers import PretrainedConfig
|
| 36 |
+
from torch.nn.parameter import Parameter
|
| 37 |
+
import bz2
|
| 38 |
+
import torch
|
| 39 |
+
import base64
|
| 40 |
+
import ctypes
|
| 41 |
+
from transformers.utils import logging
|
| 42 |
+
from typing import List
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
logger = logging.get_logger(__name__)
|
| 47 |
+
|
| 48 |
+
try:
|
| 49 |
+
from cpm_kernels.kernels.base import LazyKernelCModule, KernelFunction, round_up
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class Kernel:
|
| 53 |
+
def __init__(self, code: bytes, function_names: List[str]):
|
| 54 |
+
self.code = code
|
| 55 |
+
self._function_names = function_names
|
| 56 |
+
self._cmodule = LazyKernelCModule(self.code)
|
| 57 |
+
|
| 58 |
+
for name in self._function_names:
|
| 59 |
+
setattr(self, name, KernelFunction(self._cmodule, name))
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
quantization_code = "$QlpoOTFBWSZTWU9yuJUAQHN//////////f/n/8/n///n//bt4dTidcVx8X3V9FV/92/v4B7/AD5FBQFAAAChSgKpFCFAFVSigUAAAEKhSgUUqgFBKigqVREQAABQBQIANDTTIGI00BkZBkNGE0A0BkBkGQGRkaNAaAGQNBoGgDIAAYIGTI0DQAQAaGmmQMRpoDIyDIaMJoBoDIDIMgMjI0aA0AMgaDQNAGQAAwQMmRoGgAgA0NNMgYjTQGRkGQ0YTQDQGQGQZAZGRo0BoAZA0GgaAMgABggZMjQNABABoaaZAxGmgMjIMhowmgGgMgMgyAyMjRoDQAyBoNA0AZAADBAyZGgaAAmqU1NEgJqnptU/Sn4jRR6J6epk2pqb1Q/SgAPUGgyNNGjQ2SBpoAZAAGg0NB6mgDIAAAAA2oaApSREBNAARhGiYEaEwU8pvImlP0k2aam1GaGqbFNM1MHpTwmkepmyU9R6nqPKekHqNNPUxNGhp6n6p6QaZ6o9TG1GMqcoV9ly6nRanHlq6zPNbnGZNi6HSug+2nPiZ13XcnFYZW+45W11CumhzYhchOJ2GLLV1OBjBjGf4TptOddTSOcVxhqYZMYwZXZZY00zI1paX5X9J+b+f4e+x43RXSxXPOdquiGpduatGyXneN696M9t4HU2eR5XX/kPhP261NTx3JO1Ow7LyuDmeo9a7d351T1ZxnvnrvYnrXv/hXxPCeuYx2XsNmO003eg9J3Z6U7b23meJ4ri01OdzTk9BNO96brz+qT5nuvvH3ds/G+m/JcG/F2XYuhXlvO+jP7U3XgrzPN/lr8Sf1n6j4j7jZs+s/T0tNaNNYzTs12rxjwztHlnire3Nzc3N1wuBwOBwXBvZfoHpD7rFmR99V5vj3aXza3xdBbXMalubTg/jIv5dfAi54Pdc75j4z412n3Npj3Ld/ENm7a3b/Cod6h/ret1/5vn/C+l+gdslMvgPSLJ8d8q+U66fevYn/tW1chleEtNTGlcHCbLRlq0tHzF5tsbbZZfHjjLgZu42XCuC3NrdjTasZGNzgxPIrGqp7r3p7L2p5XjnpPSmTd5XtzqnB6U87zzg1Ol0zd0zsLszxR6lkxp35u6/teL0L0W922cR7Lu1lpL9CsHirzuM2T+BgsyViT6LHcm0/Vr6U/7LGGyJeqTEjt0PHWhF5mCT7R9mtlDwriYv0Tyr/OxYt6qp5r0mPVT0608TqnqMZaarU2nFwrTzzlrs1ed7z1ux60wyr4ydCaTi3enW8x68x0zU7tXSlcmPSW1mGpWJMg4zmPC2lK96tp0OE80y4MfEvnZj8zGluR6b22ki1Ou9V2nCd9xovcPvcYMZYy0lvN60ScZ45vN6yeCeeXFb1lVjnnCar5fwXwE2bzJ4HI1XVPXfXZMm44GUsMpYsmLB65TuVdm0cl0b+i/wGNN66XjeV7zuPpHcnK/juhhjdfId5jMdE5nN0dGmmm2zZs2cexD5n9p/dY352XsvXHaZNWWsmmS1atjR452nYudzvqv2HMRyvNNnlMcDl3R2+yx2uVrBubTW9icHDVtbNXlZm7jma1rM4VurZZd2y6nUau7ZXZ7bVU+mnoOVxZGMrVmvX60605JwmzGZhhhjTWtaaaMaaGTGmNMZasY0iX8VMUl8eepaIrzGSpemWOQyZORk2bNpjUybMmxqYmknCGCFynutfksaZpjTNMaaatM0xsxcGR0sociNqxNSmhhR1ZJPbsn8qyF0t2qH6iYBclclalbtTTcHTDsPaX6rlnElph2Jyumumtynv2Kk8GI7rsvXbIcJgHJOSaSXnnGaI3m87RtVXJOZ/YtgdTE6Wpha6ZlE8ayXkef1fh602r2WwvfMXtMdLlkfnLFdYYwYso+bWqm7yJqHXZGw2nrS5ZanSYnWlxBxMF1V940K2wdrI7R6OYf7DGGamMmTSbRhlS45xmVOumF1EyPCmHrrN8wwZOOrdNtLeMtzFzDlWnfTBxMk2NaXIZHBYxYLD4w8yju0ao65Vz1OIXoS9dLanwCe1PWrYuWMqf1if1z2k2yYfKJ741PDgno1ZQ8DRqvUny3mNoWTzGO6m1DkrJI8JiR5cSd+vZdGOO8nrMoc5+NDUFsMSXaZJeNlMmGLtJsovOsUp7I9S5VojKxF6bTVEelXqlfJobQr3LozSh2Jk7VcrVMfhXqszGWMzNqGhqZY0OadxkyyMssKugZR0KNFXBHlqwmJgTE/BNVMk6ItJXZMR0H47GpXv/DMOvNkmVuaV1PRfEdxuqc7Hcd+ZV/zTLaRxWk0nl9CdCeM6mn5rstHIBcpiuwmUZXeq81DacHI2rmrZ5SuE5mOZd6LQrZg9mx32TprA8BMo5jKN6yLTCi3WzQaZSuhzTtM1fUTGVpG8Tw+KXI0tjEpiWxtLYynOlktSbVlaI5kxP8TDH8kx50xoxi5KcA4pcja8KWLRlO/Ks6q06ergnvm1ca3Tq8Uw7LTUsmWyctXPWmpitl/uvGcWTGXGuAXDfhqazGmjkxcJW5hMMMMpYsXl2TZYtVOddG3XCarUt6Ptq9CZXSNzyuRzqRZOjsxdBbFVz6OA5HI43r1jityVlVpVkxmOsyaYWE1NTGq1sOVh36mHMcxtSvcy70edG0ZGR3I1Go1GRlV7mWWo1G0ZGRqlvH40l7o4m5xMWLLLYyNjnqc8556mdPqLJ31n/1nWOncxzG1tizrHs/Z+d2vP/B/l8wdJ6rHUn2nbbDq4p6htFtYzMMMTaZis1K5GKzGNmxhmUx2DDlZ/qNnIx41xnaMfCZWYaZWtNLTNW8ND4Fw1MyZOCdM428suKG1ehW8TesOydg7J+YYcD4cYR+8dFK6M4E3HM9ZfRNNL+Sn6rsl4DsrDl2HpPCnfxjGXtbZtYys1ttlyJ4T+BvexjGWRjMszK4Jpc77D3GyuVD7q0+G8m9G+2+rGm7cOR2y7FdtY2XUYx/oNlfRYxhMYyYZkyyg55enna9Kt/FFi6GMMwYwdwxWgxGMLKYmUyGExTKMZkMFhkymKuh0NOBNnBu+23LdwDoZYYzGGMxtORaTU1pjTGWTTGGtMrNWUsyyTTLLG1qy2ZjbK2DBllWqxMtBMaYZQmcE7zvvRcTkclUwdkxTaSdyySt/7fpL+T1v516Ji97fwr5JbLu305zMn5+GMTTZ9F+y7ExwmGVfG44yxn3dLv6l5i+Wth1jCrDq21nW9LqvvDzz3Vf3LLH/O/32TJ/erx3bXftO4eF+G956D952K/An4NfvOpjFjExjevP/UmE0fIoZXx6/w6lX/no3D0bLt+ixjieBM6ksRd0yB4Lt2SwYNE+gd1detlZWUnpiZfGfFaK+4PyCa/v18V8X75pe9fLXzp7l3VjF76vWZmHwGz1IZNWT7b8yddJ4q5kyrVdfru6atWc7bVYztL9Jf4GXvT+Y8m9/YsXP6H018a8D4XVOqvfzqeR+6yZOD8dPv0+U7/q5Pl+2dNb0MjzGVH5p6MNQ7cOWvw62U9aHE8DprDek+McLyvDz+te+9Zhq5+YTruufMcWMabqysTmZVWjKPfnK0wyVcrsuhjZRdLkHNvD72b9abriOSGIxiLixMOoalNPXzy+wT/tf+U6HHONfsz+xe8ufHBdQWWGWLA9if0rsnmrxK5LvRZQeWsTCsrmOYy8VteVfuRfcVTtDLItLIsMYxZLdU/DbtSemxF6Z6Zo5WBXE4tFdCyVMMXMTEMZXVlS6Xec2T4e0tHsRcEuWshcJ2YsNF5rUx1E8ifCq6Z+ZP7qdCeu/aTwFd53l16/o0NOw6O3dLavP4Hbi4RdmuDk6DoYaninC0+o4uZjbJ7Rxeu0/FbuFg+q7DVS6fQe0rZ6NDGUNNU6DEqOaLTicKnYZMnBWruljQxoaS3dZhocDge0bSTyOvdAbG5hxe2xji7E/L55xX13wWNDi6HCekcFxfCPGxY0MXC+s7afWaMdDyjyr+o8Rudm/NabOZvdl274zH4f5XK9z6On1Pe/K5TdPAslg77BjuO6Y3eO7GqvOPG/stknp1leyvLL0Z7bl9I4noMvLkzytLhWYzrOZzLXCORe028rORzOg4N/L0HlMOQ3Pgmnbb6KczlabORpu980q37TBqRu0/p3PO6234Bl03Ynuz+9W7gnsEcmvYaYY3aMYY0wx3pYd+ujsXauWdaY5Xkbtl23fPzFHiDB/QMo0yFjBllYxTQYYyxkrwn7JufwJ/PfgJ+C83X69ni6zvXcnyXabv0ncbLwsceS+RNlyN2mnneJtX0ngYO0+e+0+UnA+Wch3ji8hj5an4h+i6XBySU4n+R0roVcbw5yvHrmr4Yw8Y7x6c+9POPYHI5HI5HI5HI5HGXGww4nE4nrVyOR8XeqPEO7PLOiukYa3Novk5hV4cdtYZLI93e+uxff2jRo0aNGjRo0aNG1bVtW1dy3m83m8+tQ5ZzHw3nObwOu8La9Rc1dtkdS8A3eTk823tnktXWlxN6Oixe06zrN70Isd9jiOgZFq9yfkPqP/SLhN2Myl8jDM43bl1nbcb4cO57jlh8Jow6pzXZdL4dyODTuuhu77FyO27DdwdRxmvO+O+3N2+BdqyTwLHVczDVY4UPE4O66/ZO2cx1LFzVdSXtF7G4HMbrauOHRw6c8FdZ5m9fHZHYZXfTlZquyynSyTTKke6vcffSD9pzPA/G7n7jxPmuhc1DHMynPMrGL6AdewYmwu5ko+UUyTwrMv27rPH1v1nGqd87+p6N6LU8k3NEng53xXyHS97+44OSg/sy/hn+Se6yfYNjW0/uTgP+PvWYzLMmjhcLB/gGpri6H83/84eUXWT6T9Hsv7785z/7z4icpW+zfXypuR7rx/gMdZb1/wC678pcs8/2a3mDitGHxl9mfPlll5MafWWqxk/eYuTDgcNMzDGWLWvsuglNxs53GtN6uWpktlW1tZZYcuinMMWmnNnJydze3b2Y1McBxrBkXw799izLMZZYyy0TkbsGM4p03S2uVu5s/XXUdSdec6smVxZYYGpVmT8A+8ajuEyV5FatkvVru2x6uxGXXbH4A+jvgP4GMYy3iPLXzq/6z65+E005ey+cwMZD3fZcqc6xpjTFjQ0P3U+e++cPYmTIwj0nrK5NPTfl3WvpfLtXDcb2HQMudYOxFXQBor4L4T6vrOauFctYXJQ++NUWmJe5bmx1jDiZS1dTqWxo4GR8jm3fttpmPHppk9PEyv4/y8/sO07XacOmcqc0x2Vi9BvNJvN5oW8x4mOsydpidRxMYJPx06m1bqPzq9KtK8sxXNXFodD/+MYYaJTLwOhc9brCsV18oOR1i4tXChyTkq4lf4y1Ke+9axjDHqs1mfBbMXuP4Hzi+X7t8vzv7bHerrUPgPCxhjre4fXdfLNtNM+Jd+Zdh8xd8wP87uNPoPgv4W7/5P2BuxfsMabNnMnza+54Pdi5U671GPZY8CehX8Voeoo7FHpkeEc6715FwHZrIrUrHaviPUbPZHND+IhczrP6FcYvhOZ0Di/ETt0OI+YwNWR9r7tpf6WDeZKZDB1+z2IthOl1mPyb5FluvEx9h9d0NnM0Y1XPFkWIsk1WotJ0PBMmkvjvQTd0e71tfeV+8r8lQ/tpzpsmxJ+InrI/dj2UajUajVTUajatRqNRtGo1Go1Go4wjeMpZFMVV9CHbofPraLsJ3JpWV2XOoanCuFky4y3PPNxucK2uKC1Lbdb1eo+m5XomN6HfeZsabHLHRX/K+offtNGGmHWctcVcG44MdSqsOLY9VzX+Zxfxn2HPdWTpzWvkrtJ8M5zorrKcquRytJ5N5DZmcaW02l76nWO+BqPXm1A2Ry/0q71dH/mqrqeFjkYxjEXtsX8qubTk67rGycyqsdm4tZx5D6D5hhi0waaWmiaMP81Yjii5qxPlPuU/GfTL1Y5E6Jyfiq63qTa39A4J0sOGDgO9WF9bOXl0XfPRbsY2bPNKPy1YrFYrFYmRhhlTIyMjJWJYZHXuCXI8OoXsvfljGLFicNifpp2XunoPiG1wtx3p1Tah+/DD66OnVtVXP9rKbVxOnL0tR/rHtqB5UDErUVcl11D4qqvjpOcxX7armUNJB3LpW6bxVvD08e8h3odKKvyCFZBdSh2FVcST9xV3n3T8t1j7Kr9qgrqXg+13Pt5U7JCvFXVIV1YG5lRhkVYZJYYDDD4KOIMoHCp26WS8GB7uBh2zIdgq/PKyInjV2STShuoapUdCpX1yTwqq/z1VvET7Kh5nVPkO8YyxjLt2MaaMmWTLQvx3qnzltnXW0p2jxgbEtSny/Osv8Y9pLMXYoHVPAhkVdWVeODhR6q9/Sxe2liwwZWMVvFXfRkeIDxAePUPIrdJ4ey6yquzH+PD/bUOWAu05qVHtFd8rrKHSoeNIOUqrYr3FXyToqfYJgwmJdKpXXOwYYegNNGMzfZPp/t3t/DVs4zjNTN61rRqaWaa4NYbRjTa0tWwy2Y2tGN8ZO8ofNKq4j9SL7I+cSm4/6ovLV5HNXLI0jJidwrtk6ynCaP6Z++GjRlWS3tLeW129Mi9evxU9mtz6s5J3Z7M2ngTgnKvmpomxpaLCzPfmx0JWE+m3NLDDGOX47RctdYYNK5jakdqLkRlI39n590T5zctGSwwZZDJj6kW8XSi6ot2MmWWJ0DUT3nuvebBudScjZ79g8cWJ8av0k+/bE5WKd5MdbFpbDVMxu1DVMmtNZGJvq1mtRbn6M+g/kP0FwDwr7quZs7xosNGpbscyxhhd9TyJyFwbLcxlTasg75vW7TsV5K7ji44XPMMrdoj+Y3rT0Hie62nlYV/pwczzOmdLqLhYkzGMzCZWGMQzGMSsZYY6Di1t4nlJ+Em63mJxrVLxPbYxNEdgc1dU2iOKyoYYWjNrEeHTYybVk0atSa7ehuwsWMWTqn1TrnS6hYsi71d1+s+k+ic70e20fzE/VaTdxT9ZtU4GIXdeNx3X77guYYfpHeTQjaMX6brOu4OY4K7Y2d9mbHarI5ox3p4GpJ2Vd/Tst60f7j999pppjR+Q/Qf8J/VaORs3cji7FfFuN61+ui9s8hix1OCh5KGVV23BPXvZfz3CLyHpix+exi8z/KnCnosY2eunor+cxyPO/xJ0vKey9OvE9VjqaYu0x3Z3jd6o2b1T12D+F8l232lwaaacD5LE8LBxu7WTlbWraWpew8Xexjel3E+wWD4APITdNqR8F3R3T0lunCQ4GaE9R37DxeCYfcHi4xci5ovKfxVs55y2hf+65E/Xdp6jR5nrebTmi5incpkyOjs50JvrZwstbbW6kfuuQw+2mykf/EXNFzxfKTrxew929TR6bWnGL//F3JFOFCQT3K4lQ"
|
| 63 |
+
|
| 64 |
+
kernels = Kernel(
|
| 65 |
+
bz2.decompress(base64.b64decode(quantization_code)),
|
| 66 |
+
[
|
| 67 |
+
"int4WeightCompression",
|
| 68 |
+
"int4WeightExtractionFloat",
|
| 69 |
+
"int4WeightExtractionHalf",
|
| 70 |
+
"int8WeightExtractionFloat",
|
| 71 |
+
"int8WeightExtractionHalf",
|
| 72 |
+
],
|
| 73 |
+
)
|
| 74 |
+
except Exception as exception:
|
| 75 |
+
kernels = None
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
class W8A16Linear(torch.autograd.Function):
|
| 79 |
+
@staticmethod
|
| 80 |
+
def forward(ctx, inp: torch.Tensor, quant_w: torch.Tensor, scale_w: torch.Tensor, weight_bit_width):
|
| 81 |
+
ctx.inp_shape = inp.size()
|
| 82 |
+
ctx.weight_bit_width = weight_bit_width
|
| 83 |
+
out_features = quant_w.size(0)
|
| 84 |
+
inp = inp.contiguous().view(-1, inp.size(-1))
|
| 85 |
+
weight = extract_weight_to_half(quant_w, scale_w, weight_bit_width)
|
| 86 |
+
ctx.weight_shape = weight.size()
|
| 87 |
+
output = inp.mm(weight.t())
|
| 88 |
+
ctx.save_for_backward(inp, quant_w, scale_w)
|
| 89 |
+
return output.view(*(ctx.inp_shape[:-1] + (out_features,)))
|
| 90 |
+
|
| 91 |
+
@staticmethod
|
| 92 |
+
def backward(ctx, grad_output: torch.Tensor):
|
| 93 |
+
inp, quant_w, scale_w = ctx.saved_tensors
|
| 94 |
+
weight = extract_weight_to_half(quant_w, scale_w, ctx.weight_bit_width)
|
| 95 |
+
grad_output = grad_output.contiguous().view(-1, weight.size(0))
|
| 96 |
+
grad_input = grad_output.mm(weight)
|
| 97 |
+
grad_weight = grad_output.t().mm(inp)
|
| 98 |
+
return grad_input.view(ctx.inp_shape), grad_weight.view(ctx.weight_shape), None, None
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def compress_int4_weight(weight: torch.Tensor): # (n, m)
|
| 102 |
+
with torch.cuda.device(weight.device):
|
| 103 |
+
n, m = weight.size(0), weight.size(1)
|
| 104 |
+
assert m % 2 == 0
|
| 105 |
+
m = m // 2
|
| 106 |
+
out = torch.empty(n, m, dtype=torch.int8, device="cuda")
|
| 107 |
+
stream = torch.cuda.current_stream()
|
| 108 |
+
|
| 109 |
+
gridDim = (n, 1, 1)
|
| 110 |
+
blockDim = (min(round_up(m, 32), 1024), 1, 1)
|
| 111 |
+
|
| 112 |
+
kernels.int4WeightCompression(
|
| 113 |
+
gridDim,
|
| 114 |
+
blockDim,
|
| 115 |
+
0,
|
| 116 |
+
stream,
|
| 117 |
+
[ctypes.c_void_p(weight.data_ptr()), ctypes.c_void_p(out.data_ptr()), ctypes.c_int32(n), ctypes.c_int32(m)],
|
| 118 |
+
)
|
| 119 |
+
return out
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
def extract_weight_to_half(weight: torch.Tensor, scale_list: torch.Tensor, source_bit_width: int):
|
| 123 |
+
assert scale_list.dtype in [torch.half, torch.bfloat16]
|
| 124 |
+
assert weight.dtype in [torch.int8]
|
| 125 |
+
if source_bit_width == 8:
|
| 126 |
+
return weight.to(scale_list.dtype) * scale_list[:, None]
|
| 127 |
+
elif source_bit_width == 4:
|
| 128 |
+
func = (
|
| 129 |
+
kernels.int4WeightExtractionHalf if scale_list.dtype == torch.half else kernels.int4WeightExtractionBFloat16
|
| 130 |
+
)
|
| 131 |
+
else:
|
| 132 |
+
assert False, "Unsupported bit-width"
|
| 133 |
+
|
| 134 |
+
with torch.cuda.device(weight.device):
|
| 135 |
+
n, m = weight.size(0), weight.size(1)
|
| 136 |
+
out = torch.empty(n, m * (8 // source_bit_width), dtype=scale_list.dtype, device="cuda")
|
| 137 |
+
stream = torch.cuda.current_stream()
|
| 138 |
+
|
| 139 |
+
gridDim = (n, 1, 1)
|
| 140 |
+
blockDim = (min(round_up(m, 32), 1024), 1, 1)
|
| 141 |
+
|
| 142 |
+
func(
|
| 143 |
+
gridDim,
|
| 144 |
+
blockDim,
|
| 145 |
+
0,
|
| 146 |
+
stream,
|
| 147 |
+
[
|
| 148 |
+
ctypes.c_void_p(weight.data_ptr()),
|
| 149 |
+
ctypes.c_void_p(scale_list.data_ptr()),
|
| 150 |
+
ctypes.c_void_p(out.data_ptr()),
|
| 151 |
+
ctypes.c_int32(n),
|
| 152 |
+
ctypes.c_int32(m),
|
| 153 |
+
],
|
| 154 |
+
)
|
| 155 |
+
return out
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
class QuantizedLinear(torch.nn.Module):
|
| 159 |
+
def __init__(self, weight_bit_width: int, weight, bias=None, device="cuda", dtype=None, empty_init=False):
|
| 160 |
+
super().__init__()
|
| 161 |
+
weight = weight.to(device) # ensure the weight is on the cuda device
|
| 162 |
+
assert str(weight.device).startswith(
|
| 163 |
+
'cuda'), 'The weights that need to be quantified should be on the CUDA device'
|
| 164 |
+
self.weight_bit_width = weight_bit_width
|
| 165 |
+
shape = weight.shape
|
| 166 |
+
|
| 167 |
+
if weight is None or empty_init:
|
| 168 |
+
self.weight = torch.empty(shape[0], shape[1] * weight_bit_width // 8, dtype=torch.int8, device=device)
|
| 169 |
+
self.weight_scale = torch.empty(shape[0], dtype=dtype, device=device)
|
| 170 |
+
else:
|
| 171 |
+
self.weight_scale = weight.abs().max(dim=-1).values / ((2 ** (weight_bit_width - 1)) - 1)
|
| 172 |
+
self.weight = torch.round(weight / self.weight_scale[:, None]).to(torch.int8)
|
| 173 |
+
if weight_bit_width == 4:
|
| 174 |
+
self.weight = compress_int4_weight(self.weight)
|
| 175 |
+
|
| 176 |
+
self.weight = Parameter(self.weight.to(device), requires_grad=False)
|
| 177 |
+
self.weight_scale = Parameter(self.weight_scale.to(device), requires_grad=False)
|
| 178 |
+
self.bias = Parameter(bias.to(device), requires_grad=False) if bias is not None else None
|
| 179 |
+
|
| 180 |
+
def forward(self, input):
|
| 181 |
+
output = W8A16Linear.apply(input, self.weight, self.weight_scale, self.weight_bit_width)
|
| 182 |
+
if self.bias is not None:
|
| 183 |
+
output = output + self.bias
|
| 184 |
+
return output
|
| 185 |
+
|
| 186 |
+
|
| 187 |
+
def quantize(model, weight_bit_width, empty_init=False, device=None):
|
| 188 |
+
"""Replace fp16 linear with quantized linear"""
|
| 189 |
+
for layer in model.layers:
|
| 190 |
+
layer.self_attention.query_key_value = QuantizedLinear(
|
| 191 |
+
weight_bit_width=weight_bit_width,
|
| 192 |
+
weight=layer.self_attention.query_key_value.weight,
|
| 193 |
+
bias=layer.self_attention.query_key_value.bias,
|
| 194 |
+
dtype=layer.self_attention.query_key_value.weight.dtype,
|
| 195 |
+
device=layer.self_attention.query_key_value.weight.device if device is None else device,
|
| 196 |
+
empty_init=empty_init
|
| 197 |
+
)
|
| 198 |
+
layer.self_attention.dense = QuantizedLinear(
|
| 199 |
+
weight_bit_width=weight_bit_width,
|
| 200 |
+
weight=layer.self_attention.dense.weight,
|
| 201 |
+
bias=layer.self_attention.dense.bias,
|
| 202 |
+
dtype=layer.self_attention.dense.weight.dtype,
|
| 203 |
+
device=layer.self_attention.dense.weight.device if device is None else device,
|
| 204 |
+
empty_init=empty_init
|
| 205 |
+
)
|
| 206 |
+
layer.mlp.dense_h_to_4h = QuantizedLinear(
|
| 207 |
+
weight_bit_width=weight_bit_width,
|
| 208 |
+
weight=layer.mlp.dense_h_to_4h.weight,
|
| 209 |
+
bias=layer.mlp.dense_h_to_4h.bias,
|
| 210 |
+
dtype=layer.mlp.dense_h_to_4h.weight.dtype,
|
| 211 |
+
device=layer.mlp.dense_h_to_4h.weight.device if device is None else device,
|
| 212 |
+
empty_init=empty_init
|
| 213 |
+
)
|
| 214 |
+
layer.mlp.dense_4h_to_h = QuantizedLinear(
|
| 215 |
+
weight_bit_width=weight_bit_width,
|
| 216 |
+
weight=layer.mlp.dense_4h_to_h.weight,
|
| 217 |
+
bias=layer.mlp.dense_4h_to_h.bias,
|
| 218 |
+
dtype=layer.mlp.dense_4h_to_h.weight.dtype,
|
| 219 |
+
device=layer.mlp.dense_4h_to_h.weight.device if device is None else device,
|
| 220 |
+
empty_init=empty_init
|
| 221 |
+
)
|
| 222 |
+
|
| 223 |
+
return model
|
| 224 |
+
|
| 225 |
+
|
| 226 |
+
|
| 227 |
+
class ChatGLMConfig(PretrainedConfig):
|
| 228 |
+
model_type = "chatglm"
|
| 229 |
+
def __init__(
|
| 230 |
+
self,
|
| 231 |
+
num_layers=28,
|
| 232 |
+
padded_vocab_size=65024,
|
| 233 |
+
hidden_size=4096,
|
| 234 |
+
ffn_hidden_size=13696,
|
| 235 |
+
kv_channels=128,
|
| 236 |
+
num_attention_heads=32,
|
| 237 |
+
seq_length=2048,
|
| 238 |
+
hidden_dropout=0.0,
|
| 239 |
+
classifier_dropout=None,
|
| 240 |
+
attention_dropout=0.0,
|
| 241 |
+
layernorm_epsilon=1e-5,
|
| 242 |
+
rmsnorm=True,
|
| 243 |
+
apply_residual_connection_post_layernorm=False,
|
| 244 |
+
post_layer_norm=True,
|
| 245 |
+
add_bias_linear=False,
|
| 246 |
+
add_qkv_bias=False,
|
| 247 |
+
bias_dropout_fusion=True,
|
| 248 |
+
multi_query_attention=False,
|
| 249 |
+
multi_query_group_num=1,
|
| 250 |
+
apply_query_key_layer_scaling=True,
|
| 251 |
+
attention_softmax_in_fp32=True,
|
| 252 |
+
fp32_residual_connection=False,
|
| 253 |
+
quantization_bit=0,
|
| 254 |
+
pre_seq_len=None,
|
| 255 |
+
prefix_projection=False,
|
| 256 |
+
**kwargs
|
| 257 |
+
):
|
| 258 |
+
self.num_layers = num_layers
|
| 259 |
+
self.vocab_size = padded_vocab_size
|
| 260 |
+
self.padded_vocab_size = padded_vocab_size
|
| 261 |
+
self.hidden_size = hidden_size
|
| 262 |
+
self.ffn_hidden_size = ffn_hidden_size
|
| 263 |
+
self.kv_channels = kv_channels
|
| 264 |
+
self.num_attention_heads = num_attention_heads
|
| 265 |
+
self.seq_length = seq_length
|
| 266 |
+
self.hidden_dropout = hidden_dropout
|
| 267 |
+
self.classifier_dropout = classifier_dropout
|
| 268 |
+
self.attention_dropout = attention_dropout
|
| 269 |
+
self.layernorm_epsilon = layernorm_epsilon
|
| 270 |
+
self.rmsnorm = rmsnorm
|
| 271 |
+
self.apply_residual_connection_post_layernorm = apply_residual_connection_post_layernorm
|
| 272 |
+
self.post_layer_norm = post_layer_norm
|
| 273 |
+
self.add_bias_linear = add_bias_linear
|
| 274 |
+
self.add_qkv_bias = add_qkv_bias
|
| 275 |
+
self.bias_dropout_fusion = bias_dropout_fusion
|
| 276 |
+
self.multi_query_attention = multi_query_attention
|
| 277 |
+
self.multi_query_group_num = multi_query_group_num
|
| 278 |
+
self.apply_query_key_layer_scaling = apply_query_key_layer_scaling
|
| 279 |
+
self.attention_softmax_in_fp32 = attention_softmax_in_fp32
|
| 280 |
+
self.fp32_residual_connection = fp32_residual_connection
|
| 281 |
+
self.quantization_bit = quantization_bit
|
| 282 |
+
self.pre_seq_len = pre_seq_len
|
| 283 |
+
self.prefix_projection = prefix_projection
|
| 284 |
+
super().__init__(**kwargs)
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
|
| 288 |
+
# flags required to enable jit fusion kernels
|
| 289 |
+
|
| 290 |
+
if sys.platform != 'darwin':
|
| 291 |
+
torch._C._jit_set_profiling_mode(False)
|
| 292 |
+
torch._C._jit_set_profiling_executor(False)
|
| 293 |
+
torch._C._jit_override_can_fuse_on_cpu(True)
|
| 294 |
+
torch._C._jit_override_can_fuse_on_gpu(True)
|
| 295 |
+
|
| 296 |
+
logger = logging.get_logger(__name__)
|
| 297 |
+
|
| 298 |
+
_CHECKPOINT_FOR_DOC = "THUDM/ChatGLM"
|
| 299 |
+
_CONFIG_FOR_DOC = "ChatGLM6BConfig"
|
| 300 |
+
|
| 301 |
+
CHATGLM_6B_PRETRAINED_MODEL_ARCHIVE_LIST = [
|
| 302 |
+
"THUDM/chatglm3-6b-base",
|
| 303 |
+
# See all ChatGLM models at https://huggingface.co/models?filter=chatglm
|
| 304 |
+
]
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
def default_init(cls, *args, **kwargs):
|
| 308 |
+
return cls(*args, **kwargs)
|
| 309 |
+
|
| 310 |
+
|
| 311 |
+
class InvalidScoreLogitsProcessor(LogitsProcessor):
|
| 312 |
+
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:
|
| 313 |
+
if torch.isnan(scores).any() or torch.isinf(scores).any():
|
| 314 |
+
scores.zero_()
|
| 315 |
+
scores[..., 5] = 5e4
|
| 316 |
+
return scores
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
class PrefixEncoder(torch.nn.Module):
|
| 320 |
+
"""
|
| 321 |
+
The torch.nn model to encode the prefix
|
| 322 |
+
Input shape: (batch-size, prefix-length)
|
| 323 |
+
Output shape: (batch-size, prefix-length, 2*layers*hidden)
|
| 324 |
+
"""
|
| 325 |
+
|
| 326 |
+
def __init__(self, config: ChatGLMConfig):
|
| 327 |
+
super().__init__()
|
| 328 |
+
self.prefix_projection = config.prefix_projection
|
| 329 |
+
if self.prefix_projection:
|
| 330 |
+
# Use a two-layer MLP to encode the prefix
|
| 331 |
+
kv_size = config.num_layers * config.kv_channels * config.multi_query_group_num * 2
|
| 332 |
+
self.embedding = torch.nn.Embedding(config.pre_seq_len, kv_size)
|
| 333 |
+
self.trans = torch.nn.Sequential(
|
| 334 |
+
torch.nn.Linear(kv_size, config.hidden_size),
|
| 335 |
+
torch.nn.Tanh(),
|
| 336 |
+
torch.nn.Linear(config.hidden_size, kv_size)
|
| 337 |
+
)
|
| 338 |
+
else:
|
| 339 |
+
self.embedding = torch.nn.Embedding(config.pre_seq_len,
|
| 340 |
+
config.num_layers * config.kv_channels * config.multi_query_group_num * 2)
|
| 341 |
+
|
| 342 |
+
def forward(self, prefix: torch.Tensor):
|
| 343 |
+
if self.prefix_projection:
|
| 344 |
+
prefix_tokens = self.embedding(prefix)
|
| 345 |
+
past_key_values = self.trans(prefix_tokens)
|
| 346 |
+
else:
|
| 347 |
+
past_key_values = self.embedding(prefix)
|
| 348 |
+
return past_key_values
|
| 349 |
+
|
| 350 |
+
|
| 351 |
+
def split_tensor_along_last_dim(
|
| 352 |
+
tensor: torch.Tensor,
|
| 353 |
+
num_partitions: int,
|
| 354 |
+
contiguous_split_chunks: bool = False,
|
| 355 |
+
) -> List[torch.Tensor]:
|
| 356 |
+
"""Split a tensor along its last dimension.
|
| 357 |
+
|
| 358 |
+
Arguments:
|
| 359 |
+
tensor: input tensor.
|
| 360 |
+
num_partitions: number of partitions to split the tensor
|
| 361 |
+
contiguous_split_chunks: If True, make each chunk contiguous
|
| 362 |
+
in memory.
|
| 363 |
+
|
| 364 |
+
Returns:
|
| 365 |
+
A list of Tensors
|
| 366 |
+
"""
|
| 367 |
+
# Get the size and dimension.
|
| 368 |
+
last_dim = tensor.dim() - 1
|
| 369 |
+
last_dim_size = tensor.size()[last_dim] // num_partitions
|
| 370 |
+
# Split.
|
| 371 |
+
tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)
|
| 372 |
+
# Note: torch.split does not create contiguous tensors by default.
|
| 373 |
+
if contiguous_split_chunks:
|
| 374 |
+
return tuple(chunk.contiguous() for chunk in tensor_list)
|
| 375 |
+
|
| 376 |
+
return tensor_list
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
class RotaryEmbedding(nn.Module):
|
| 380 |
+
def __init__(self, dim, original_impl=False, device=None, dtype=None):
|
| 381 |
+
super().__init__()
|
| 382 |
+
inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2, device=device).to(dtype=dtype) / dim))
|
| 383 |
+
self.register_buffer("inv_freq", inv_freq)
|
| 384 |
+
self.dim = dim
|
| 385 |
+
self.original_impl = original_impl
|
| 386 |
+
|
| 387 |
+
def forward_impl(
|
| 388 |
+
self, seq_len: int, n_elem: int, dtype: torch.dtype, device: torch.device, base: int = 10000
|
| 389 |
+
):
|
| 390 |
+
"""Enhanced Transformer with Rotary Position Embedding.
|
| 391 |
+
|
| 392 |
+
Derived from: https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/labml_nn/
|
| 393 |
+
transformers/rope/__init__.py. MIT License:
|
| 394 |
+
https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/license.
|
| 395 |
+
"""
|
| 396 |
+
# $\Theta = {\theta_i = 10000^{\frac{2(i-1)}{d}}, i \in [1, 2, ..., \frac{d}{2}]}$
|
| 397 |
+
theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, dtype=torch.float, device=device) / n_elem))
|
| 398 |
+
|
| 399 |
+
# Create position indexes `[0, 1, ..., seq_len - 1]`
|
| 400 |
+
seq_idx = torch.arange(seq_len, dtype=torch.float, device=device)
|
| 401 |
+
|
| 402 |
+
# Calculate the product of position index and $\theta_i$
|
| 403 |
+
idx_theta = torch.outer(seq_idx, theta).float()
|
| 404 |
+
|
| 405 |
+
cache = torch.stack([torch.cos(idx_theta), torch.sin(idx_theta)], dim=-1)
|
| 406 |
+
|
| 407 |
+
# this is to mimic the behaviour of complex32, else we will get different results
|
| 408 |
+
if dtype in (torch.float16, torch.bfloat16, torch.int8):
|
| 409 |
+
cache = cache.bfloat16() if dtype == torch.bfloat16 else cache.half()
|
| 410 |
+
return cache
|
| 411 |
+
|
| 412 |
+
def forward(self, max_seq_len, offset=0):
|
| 413 |
+
return self.forward_impl(
|
| 414 |
+
max_seq_len, self.dim, dtype=self.inv_freq.dtype, device=self.inv_freq.device
|
| 415 |
+
)
|
| 416 |
+
|
| 417 |
+
|
| 418 |
+
@torch.jit.script
|
| 419 |
+
def apply_rotary_pos_emb(x: torch.Tensor, rope_cache: torch.Tensor) -> torch.Tensor:
|
| 420 |
+
# x: [sq, b, np, hn]
|
| 421 |
+
sq, b, np, hn = x.size(0), x.size(1), x.size(2), x.size(3)
|
| 422 |
+
rot_dim = rope_cache.shape[-2] * 2
|
| 423 |
+
x, x_pass = x[..., :rot_dim], x[..., rot_dim:]
|
| 424 |
+
# truncate to support variable sizes
|
| 425 |
+
rope_cache = rope_cache[:sq]
|
| 426 |
+
xshaped = x.reshape(sq, -1, np, rot_dim // 2, 2)
|
| 427 |
+
rope_cache = rope_cache.view(sq, -1, 1, xshaped.size(3), 2)
|
| 428 |
+
x_out2 = torch.stack(
|
| 429 |
+
[
|
| 430 |
+
xshaped[..., 0] * rope_cache[..., 0] - xshaped[..., 1] * rope_cache[..., 1],
|
| 431 |
+
xshaped[..., 1] * rope_cache[..., 0] + xshaped[..., 0] * rope_cache[..., 1],
|
| 432 |
+
],
|
| 433 |
+
-1,
|
| 434 |
+
)
|
| 435 |
+
x_out2 = x_out2.flatten(3)
|
| 436 |
+
return torch.cat((x_out2, x_pass), dim=-1)
|
| 437 |
+
|
| 438 |
+
|
| 439 |
+
class RMSNorm(torch.nn.Module):
|
| 440 |
+
def __init__(self, normalized_shape, eps=1e-5, device=None, dtype=None, **kwargs):
|
| 441 |
+
super().__init__()
|
| 442 |
+
self.weight = torch.nn.Parameter(torch.empty(normalized_shape, device=device, dtype=dtype))
|
| 443 |
+
self.eps = eps
|
| 444 |
+
|
| 445 |
+
def forward(self, hidden_states: torch.Tensor):
|
| 446 |
+
input_dtype = hidden_states.dtype
|
| 447 |
+
variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
| 448 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
|
| 449 |
+
|
| 450 |
+
return (self.weight * hidden_states).to(input_dtype)
|
| 451 |
+
|
| 452 |
+
|
| 453 |
+
class CoreAttention(torch.nn.Module):
|
| 454 |
+
def __init__(self, config: ChatGLMConfig, layer_number):
|
| 455 |
+
super(CoreAttention, self).__init__()
|
| 456 |
+
|
| 457 |
+
self.apply_query_key_layer_scaling = config.apply_query_key_layer_scaling
|
| 458 |
+
self.attention_softmax_in_fp32 = config.attention_softmax_in_fp32
|
| 459 |
+
if self.apply_query_key_layer_scaling:
|
| 460 |
+
self.attention_softmax_in_fp32 = True
|
| 461 |
+
self.layer_number = max(1, layer_number)
|
| 462 |
+
|
| 463 |
+
projection_size = config.kv_channels * config.num_attention_heads
|
| 464 |
+
|
| 465 |
+
# Per attention head and per partition values.
|
| 466 |
+
self.hidden_size_per_partition = projection_size
|
| 467 |
+
self.hidden_size_per_attention_head = projection_size // config.num_attention_heads
|
| 468 |
+
self.num_attention_heads_per_partition = config.num_attention_heads
|
| 469 |
+
|
| 470 |
+
coeff = None
|
| 471 |
+
self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)
|
| 472 |
+
if self.apply_query_key_layer_scaling:
|
| 473 |
+
coeff = self.layer_number
|
| 474 |
+
self.norm_factor *= coeff
|
| 475 |
+
self.coeff = coeff
|
| 476 |
+
|
| 477 |
+
self.attention_dropout = torch.nn.Dropout(config.attention_dropout)
|
| 478 |
+
|
| 479 |
+
def forward(self, query_layer, key_layer, value_layer, attention_mask):
|
| 480 |
+
pytorch_major_version = int(torch.__version__.split('.')[0])
|
| 481 |
+
if pytorch_major_version >= 2:
|
| 482 |
+
query_layer, key_layer, value_layer = [k.permute(1, 2, 0, 3) for k in [query_layer, key_layer, value_layer]]
|
| 483 |
+
if attention_mask is None and query_layer.shape[2] == key_layer.shape[2]:
|
| 484 |
+
context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,
|
| 485 |
+
is_causal=True)
|
| 486 |
+
else:
|
| 487 |
+
if attention_mask is not None:
|
| 488 |
+
attention_mask = ~attention_mask
|
| 489 |
+
context_layer = torch.nn.functional.scaled_dot_product_attention(query_layer, key_layer, value_layer,
|
| 490 |
+
attention_mask)
|
| 491 |
+
context_layer = context_layer.permute(2, 0, 1, 3)
|
| 492 |
+
new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)
|
| 493 |
+
context_layer = context_layer.reshape(*new_context_layer_shape)
|
| 494 |
+
else:
|
| 495 |
+
# Raw attention scores
|
| 496 |
+
|
| 497 |
+
# [b, np, sq, sk]
|
| 498 |
+
output_size = (query_layer.size(1), query_layer.size(2), query_layer.size(0), key_layer.size(0))
|
| 499 |
+
|
| 500 |
+
# [sq, b, np, hn] -> [sq, b * np, hn]
|
| 501 |
+
query_layer = query_layer.view(output_size[2], output_size[0] * output_size[1], -1)
|
| 502 |
+
# [sk, b, np, hn] -> [sk, b * np, hn]
|
| 503 |
+
key_layer = key_layer.view(output_size[3], output_size[0] * output_size[1], -1)
|
| 504 |
+
|
| 505 |
+
# preallocting input tensor: [b * np, sq, sk]
|
| 506 |
+
matmul_input_buffer = torch.empty(
|
| 507 |
+
output_size[0] * output_size[1], output_size[2], output_size[3], dtype=query_layer.dtype,
|
| 508 |
+
device=query_layer.device
|
| 509 |
+
)
|
| 510 |
+
|
| 511 |
+
# Raw attention scores. [b * np, sq, sk]
|
| 512 |
+
matmul_result = torch.baddbmm(
|
| 513 |
+
matmul_input_buffer,
|
| 514 |
+
query_layer.transpose(0, 1), # [b * np, sq, hn]
|
| 515 |
+
key_layer.transpose(0, 1).transpose(1, 2), # [b * np, hn, sk]
|
| 516 |
+
beta=0.0,
|
| 517 |
+
alpha=(1.0 / self.norm_factor),
|
| 518 |
+
)
|
| 519 |
+
|
| 520 |
+
# change view to [b, np, sq, sk]
|
| 521 |
+
attention_scores = matmul_result.view(*output_size)
|
| 522 |
+
|
| 523 |
+
# ===========================
|
| 524 |
+
# Attention probs and dropout
|
| 525 |
+
# ===========================
|
| 526 |
+
|
| 527 |
+
# attention scores and attention mask [b, np, sq, sk]
|
| 528 |
+
if self.attention_softmax_in_fp32:
|
| 529 |
+
attention_scores = attention_scores.float()
|
| 530 |
+
if self.coeff is not None:
|
| 531 |
+
attention_scores = attention_scores * self.coeff
|
| 532 |
+
if attention_mask is None and attention_scores.shape[2] == attention_scores.shape[3]:
|
| 533 |
+
attention_mask = torch.ones(output_size[0], 1, output_size[2], output_size[3],
|
| 534 |
+
device=attention_scores.device, dtype=torch.bool)
|
| 535 |
+
attention_mask.tril_()
|
| 536 |
+
attention_mask = ~attention_mask
|
| 537 |
+
if attention_mask is not None:
|
| 538 |
+
attention_scores = attention_scores.masked_fill(attention_mask, float("-inf"))
|
| 539 |
+
attention_probs = F.softmax(attention_scores, dim=-1)
|
| 540 |
+
attention_probs = attention_probs.type_as(value_layer)
|
| 541 |
+
|
| 542 |
+
# This is actually dropping out entire tokens to attend to, which might
|
| 543 |
+
# seem a bit unusual, but is taken from the original Transformer paper.
|
| 544 |
+
attention_probs = self.attention_dropout(attention_probs)
|
| 545 |
+
# =========================
|
| 546 |
+
# Context layer. [sq, b, hp]
|
| 547 |
+
# =========================
|
| 548 |
+
|
| 549 |
+
# value_layer -> context layer.
|
| 550 |
+
# [sk, b, np, hn] --> [b, np, sq, hn]
|
| 551 |
+
|
| 552 |
+
# context layer shape: [b, np, sq, hn]
|
| 553 |
+
output_size = (value_layer.size(1), value_layer.size(2), query_layer.size(0), value_layer.size(3))
|
| 554 |
+
# change view [sk, b * np, hn]
|
| 555 |
+
value_layer = value_layer.view(value_layer.size(0), output_size[0] * output_size[1], -1)
|
| 556 |
+
# change view [b * np, sq, sk]
|
| 557 |
+
attention_probs = attention_probs.view(output_size[0] * output_size[1], output_size[2], -1)
|
| 558 |
+
# matmul: [b * np, sq, hn]
|
| 559 |
+
context_layer = torch.bmm(attention_probs, value_layer.transpose(0, 1))
|
| 560 |
+
# change view [b, np, sq, hn]
|
| 561 |
+
context_layer = context_layer.view(*output_size)
|
| 562 |
+
# [b, np, sq, hn] --> [sq, b, np, hn]
|
| 563 |
+
context_layer = context_layer.permute(2, 0, 1, 3).contiguous()
|
| 564 |
+
# [sq, b, np, hn] --> [sq, b, hp]
|
| 565 |
+
new_context_layer_shape = context_layer.size()[:-2] + (self.hidden_size_per_partition,)
|
| 566 |
+
context_layer = context_layer.view(*new_context_layer_shape)
|
| 567 |
+
|
| 568 |
+
return context_layer
|
| 569 |
+
|
| 570 |
+
|
| 571 |
+
class SelfAttention(torch.nn.Module):
|
| 572 |
+
"""Parallel self-attention layer abstract class.
|
| 573 |
+
|
| 574 |
+
Self-attention layer takes input with size [s, b, h]
|
| 575 |
+
and returns output of the same size.
|
| 576 |
+
"""
|
| 577 |
+
|
| 578 |
+
def __init__(self, config: ChatGLMConfig, layer_number, device=None):
|
| 579 |
+
super(SelfAttention, self).__init__()
|
| 580 |
+
self.layer_number = max(1, layer_number)
|
| 581 |
+
|
| 582 |
+
self.projection_size = config.kv_channels * config.num_attention_heads
|
| 583 |
+
|
| 584 |
+
# Per attention head and per partition values.
|
| 585 |
+
self.hidden_size_per_attention_head = self.projection_size // config.num_attention_heads
|
| 586 |
+
self.num_attention_heads_per_partition = config.num_attention_heads
|
| 587 |
+
|
| 588 |
+
self.multi_query_attention = config.multi_query_attention
|
| 589 |
+
self.qkv_hidden_size = 3 * self.projection_size
|
| 590 |
+
if self.multi_query_attention:
|
| 591 |
+
self.num_multi_query_groups_per_partition = config.multi_query_group_num
|
| 592 |
+
self.qkv_hidden_size = (
|
| 593 |
+
self.projection_size + 2 * self.hidden_size_per_attention_head * config.multi_query_group_num
|
| 594 |
+
)
|
| 595 |
+
self.query_key_value = nn.Linear(config.hidden_size, self.qkv_hidden_size,
|
| 596 |
+
bias=config.add_bias_linear or config.add_qkv_bias,
|
| 597 |
+
device=device, **_config_to_kwargs(config)
|
| 598 |
+
)
|
| 599 |
+
|
| 600 |
+
self.core_attention = CoreAttention(config, self.layer_number)
|
| 601 |
+
|
| 602 |
+
# Output.
|
| 603 |
+
self.dense = nn.Linear(self.projection_size, config.hidden_size, bias=config.add_bias_linear,
|
| 604 |
+
device=device, **_config_to_kwargs(config)
|
| 605 |
+
)
|
| 606 |
+
|
| 607 |
+
def _allocate_memory(self, inference_max_sequence_len, batch_size, device=None, dtype=None):
|
| 608 |
+
if self.multi_query_attention:
|
| 609 |
+
num_attention_heads = self.num_multi_query_groups_per_partition
|
| 610 |
+
else:
|
| 611 |
+
num_attention_heads = self.num_attention_heads_per_partition
|
| 612 |
+
return torch.empty(
|
| 613 |
+
inference_max_sequence_len,
|
| 614 |
+
batch_size,
|
| 615 |
+
num_attention_heads,
|
| 616 |
+
self.hidden_size_per_attention_head,
|
| 617 |
+
dtype=dtype,
|
| 618 |
+
device=device,
|
| 619 |
+
)
|
| 620 |
+
|
| 621 |
+
def forward(
|
| 622 |
+
self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True
|
| 623 |
+
):
|
| 624 |
+
# hidden_states: [sq, b, h]
|
| 625 |
+
|
| 626 |
+
# =================================================
|
| 627 |
+
# Pre-allocate memory for key-values for inference.
|
| 628 |
+
# =================================================
|
| 629 |
+
# =====================
|
| 630 |
+
# Query, Key, and Value
|
| 631 |
+
# =====================
|
| 632 |
+
|
| 633 |
+
# Attention heads [sq, b, h] --> [sq, b, (np * 3 * hn)]
|
| 634 |
+
mixed_x_layer = self.query_key_value(hidden_states)
|
| 635 |
+
|
| 636 |
+
if self.multi_query_attention:
|
| 637 |
+
(query_layer, key_layer, value_layer) = mixed_x_layer.split(
|
| 638 |
+
[
|
| 639 |
+
self.num_attention_heads_per_partition * self.hidden_size_per_attention_head,
|
| 640 |
+
self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,
|
| 641 |
+
self.num_multi_query_groups_per_partition * self.hidden_size_per_attention_head,
|
| 642 |
+
],
|
| 643 |
+
dim=-1,
|
| 644 |
+
)
|
| 645 |
+
query_layer = query_layer.view(
|
| 646 |
+
query_layer.size()[:-1] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)
|
| 647 |
+
)
|
| 648 |
+
key_layer = key_layer.view(
|
| 649 |
+
key_layer.size()[:-1] + (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)
|
| 650 |
+
)
|
| 651 |
+
value_layer = value_layer.view(
|
| 652 |
+
value_layer.size()[:-1]
|
| 653 |
+
+ (self.num_multi_query_groups_per_partition, self.hidden_size_per_attention_head)
|
| 654 |
+
)
|
| 655 |
+
else:
|
| 656 |
+
new_tensor_shape = mixed_x_layer.size()[:-1] + \
|
| 657 |
+
(self.num_attention_heads_per_partition,
|
| 658 |
+
3 * self.hidden_size_per_attention_head)
|
| 659 |
+
mixed_x_layer = mixed_x_layer.view(*new_tensor_shape)
|
| 660 |
+
|
| 661 |
+
# [sq, b, np, 3 * hn] --> 3 [sq, b, np, hn]
|
| 662 |
+
(query_layer, key_layer, value_layer) = split_tensor_along_last_dim(mixed_x_layer, 3)
|
| 663 |
+
|
| 664 |
+
# apply relative positional encoding (rotary embedding)
|
| 665 |
+
if rotary_pos_emb is not None:
|
| 666 |
+
query_layer = apply_rotary_pos_emb(query_layer, rotary_pos_emb)
|
| 667 |
+
key_layer = apply_rotary_pos_emb(key_layer, rotary_pos_emb)
|
| 668 |
+
|
| 669 |
+
# adjust key and value for inference
|
| 670 |
+
if kv_cache is not None:
|
| 671 |
+
cache_k, cache_v = kv_cache
|
| 672 |
+
key_layer = torch.cat((cache_k, key_layer), dim=0)
|
| 673 |
+
value_layer = torch.cat((cache_v, value_layer), dim=0)
|
| 674 |
+
if use_cache:
|
| 675 |
+
kv_cache = (key_layer, value_layer)
|
| 676 |
+
else:
|
| 677 |
+
kv_cache = None
|
| 678 |
+
|
| 679 |
+
if self.multi_query_attention:
|
| 680 |
+
key_layer = key_layer.unsqueeze(-2)
|
| 681 |
+
key_layer = key_layer.expand(
|
| 682 |
+
-1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1
|
| 683 |
+
)
|
| 684 |
+
key_layer = key_layer.contiguous().view(
|
| 685 |
+
key_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)
|
| 686 |
+
)
|
| 687 |
+
value_layer = value_layer.unsqueeze(-2)
|
| 688 |
+
value_layer = value_layer.expand(
|
| 689 |
+
-1, -1, -1, self.num_attention_heads_per_partition // self.num_multi_query_groups_per_partition, -1
|
| 690 |
+
)
|
| 691 |
+
value_layer = value_layer.contiguous().view(
|
| 692 |
+
value_layer.size()[:2] + (self.num_attention_heads_per_partition, self.hidden_size_per_attention_head)
|
| 693 |
+
)
|
| 694 |
+
|
| 695 |
+
# ==================================
|
| 696 |
+
# core attention computation
|
| 697 |
+
# ==================================
|
| 698 |
+
|
| 699 |
+
context_layer = self.core_attention(query_layer, key_layer, value_layer, attention_mask)
|
| 700 |
+
|
| 701 |
+
# =================
|
| 702 |
+
# Output. [sq, b, h]
|
| 703 |
+
# =================
|
| 704 |
+
|
| 705 |
+
output = self.dense(context_layer)
|
| 706 |
+
|
| 707 |
+
return output, kv_cache
|
| 708 |
+
|
| 709 |
+
|
| 710 |
+
def _config_to_kwargs(args):
|
| 711 |
+
common_kwargs = {
|
| 712 |
+
"dtype": args.torch_dtype,
|
| 713 |
+
}
|
| 714 |
+
return common_kwargs
|
| 715 |
+
|
| 716 |
+
|
| 717 |
+
class MLP(torch.nn.Module):
|
| 718 |
+
"""MLP.
|
| 719 |
+
|
| 720 |
+
MLP will take the input with h hidden state, project it to 4*h
|
| 721 |
+
hidden dimension, perform nonlinear transformation, and project the
|
| 722 |
+
state back into h hidden dimension.
|
| 723 |
+
"""
|
| 724 |
+
|
| 725 |
+
def __init__(self, config: ChatGLMConfig, device=None):
|
| 726 |
+
super(MLP, self).__init__()
|
| 727 |
+
|
| 728 |
+
self.add_bias = config.add_bias_linear
|
| 729 |
+
|
| 730 |
+
# Project to 4h. If using swiglu double the output width, see https://arxiv.org/pdf/2002.05202.pdf
|
| 731 |
+
self.dense_h_to_4h = nn.Linear(
|
| 732 |
+
config.hidden_size,
|
| 733 |
+
config.ffn_hidden_size * 2,
|
| 734 |
+
bias=self.add_bias,
|
| 735 |
+
device=device,
|
| 736 |
+
**_config_to_kwargs(config)
|
| 737 |
+
)
|
| 738 |
+
|
| 739 |
+
def swiglu(x):
|
| 740 |
+
x = torch.chunk(x, 2, dim=-1)
|
| 741 |
+
return F.silu(x[0]) * x[1]
|
| 742 |
+
|
| 743 |
+
self.activation_func = swiglu
|
| 744 |
+
|
| 745 |
+
# Project back to h.
|
| 746 |
+
self.dense_4h_to_h = nn.Linear(
|
| 747 |
+
config.ffn_hidden_size,
|
| 748 |
+
config.hidden_size,
|
| 749 |
+
bias=self.add_bias,
|
| 750 |
+
device=device,
|
| 751 |
+
**_config_to_kwargs(config)
|
| 752 |
+
)
|
| 753 |
+
|
| 754 |
+
def forward(self, hidden_states):
|
| 755 |
+
# [s, b, 4hp]
|
| 756 |
+
intermediate_parallel = self.dense_h_to_4h(hidden_states)
|
| 757 |
+
intermediate_parallel = self.activation_func(intermediate_parallel)
|
| 758 |
+
# [s, b, h]
|
| 759 |
+
output = self.dense_4h_to_h(intermediate_parallel)
|
| 760 |
+
return output
|
| 761 |
+
|
| 762 |
+
|
| 763 |
+
class GLMBlock(torch.nn.Module):
|
| 764 |
+
"""A single transformer layer.
|
| 765 |
+
|
| 766 |
+
Transformer layer takes input with size [s, b, h] and returns an
|
| 767 |
+
output of the same size.
|
| 768 |
+
"""
|
| 769 |
+
|
| 770 |
+
def __init__(self, config: ChatGLMConfig, layer_number, device=None):
|
| 771 |
+
super(GLMBlock, self).__init__()
|
| 772 |
+
self.layer_number = layer_number
|
| 773 |
+
|
| 774 |
+
self.apply_residual_connection_post_layernorm = config.apply_residual_connection_post_layernorm
|
| 775 |
+
|
| 776 |
+
self.fp32_residual_connection = config.fp32_residual_connection
|
| 777 |
+
|
| 778 |
+
LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm
|
| 779 |
+
# Layernorm on the input data.
|
| 780 |
+
self.input_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,
|
| 781 |
+
dtype=config.torch_dtype)
|
| 782 |
+
|
| 783 |
+
# Self attention.
|
| 784 |
+
self.self_attention = SelfAttention(config, layer_number, device=device)
|
| 785 |
+
self.hidden_dropout = config.hidden_dropout
|
| 786 |
+
|
| 787 |
+
# Layernorm on the attention output
|
| 788 |
+
self.post_attention_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,
|
| 789 |
+
dtype=config.torch_dtype)
|
| 790 |
+
|
| 791 |
+
# MLP
|
| 792 |
+
self.mlp = MLP(config, device=device)
|
| 793 |
+
|
| 794 |
+
def forward(
|
| 795 |
+
self, hidden_states, attention_mask, rotary_pos_emb, kv_cache=None, use_cache=True,
|
| 796 |
+
):
|
| 797 |
+
# hidden_states: [s, b, h]
|
| 798 |
+
|
| 799 |
+
# Layer norm at the beginning of the transformer layer.
|
| 800 |
+
layernorm_output = self.input_layernorm(hidden_states)
|
| 801 |
+
# Self attention.
|
| 802 |
+
attention_output, kv_cache = self.self_attention(
|
| 803 |
+
layernorm_output,
|
| 804 |
+
attention_mask,
|
| 805 |
+
rotary_pos_emb,
|
| 806 |
+
kv_cache=kv_cache,
|
| 807 |
+
use_cache=use_cache
|
| 808 |
+
)
|
| 809 |
+
|
| 810 |
+
# Residual connection.
|
| 811 |
+
if self.apply_residual_connection_post_layernorm:
|
| 812 |
+
residual = layernorm_output
|
| 813 |
+
else:
|
| 814 |
+
residual = hidden_states
|
| 815 |
+
|
| 816 |
+
layernorm_input = torch.nn.functional.dropout(attention_output, p=self.hidden_dropout, training=self.training)
|
| 817 |
+
layernorm_input = residual + layernorm_input
|
| 818 |
+
|
| 819 |
+
# Layer norm post the self attention.
|
| 820 |
+
layernorm_output = self.post_attention_layernorm(layernorm_input)
|
| 821 |
+
|
| 822 |
+
# MLP.
|
| 823 |
+
mlp_output = self.mlp(layernorm_output)
|
| 824 |
+
|
| 825 |
+
# Second residual connection.
|
| 826 |
+
if self.apply_residual_connection_post_layernorm:
|
| 827 |
+
residual = layernorm_output
|
| 828 |
+
else:
|
| 829 |
+
residual = layernorm_input
|
| 830 |
+
|
| 831 |
+
output = torch.nn.functional.dropout(mlp_output, p=self.hidden_dropout, training=self.training)
|
| 832 |
+
output = residual + output
|
| 833 |
+
|
| 834 |
+
return output, kv_cache
|
| 835 |
+
|
| 836 |
+
|
| 837 |
+
class GLMTransformer(torch.nn.Module):
|
| 838 |
+
"""Transformer class."""
|
| 839 |
+
|
| 840 |
+
def __init__(self, config: ChatGLMConfig, device=None):
|
| 841 |
+
super(GLMTransformer, self).__init__()
|
| 842 |
+
|
| 843 |
+
self.fp32_residual_connection = config.fp32_residual_connection
|
| 844 |
+
self.post_layer_norm = config.post_layer_norm
|
| 845 |
+
|
| 846 |
+
# Number of layers.
|
| 847 |
+
self.num_layers = config.num_layers
|
| 848 |
+
|
| 849 |
+
# Transformer layers.
|
| 850 |
+
def build_layer(layer_number):
|
| 851 |
+
return GLMBlock(config, layer_number, device=device)
|
| 852 |
+
|
| 853 |
+
self.layers = torch.nn.ModuleList([build_layer(i + 1) for i in range(self.num_layers)])
|
| 854 |
+
|
| 855 |
+
if self.post_layer_norm:
|
| 856 |
+
LayerNormFunc = RMSNorm if config.rmsnorm else LayerNorm
|
| 857 |
+
# Final layer norm before output.
|
| 858 |
+
self.final_layernorm = LayerNormFunc(config.hidden_size, eps=config.layernorm_epsilon, device=device,
|
| 859 |
+
dtype=config.torch_dtype)
|
| 860 |
+
|
| 861 |
+
self.gradient_checkpointing = False
|
| 862 |
+
|
| 863 |
+
def _get_layer(self, layer_number):
|
| 864 |
+
return self.layers[layer_number]
|
| 865 |
+
|
| 866 |
+
def forward(
|
| 867 |
+
self, hidden_states, attention_mask, rotary_pos_emb, kv_caches=None,
|
| 868 |
+
use_cache: Optional[bool] = True,
|
| 869 |
+
output_hidden_states: Optional[bool] = False,
|
| 870 |
+
):
|
| 871 |
+
if not kv_caches:
|
| 872 |
+
kv_caches = [None for _ in range(self.num_layers)]
|
| 873 |
+
presents = () if use_cache else None
|
| 874 |
+
if self.gradient_checkpointing and self.training:
|
| 875 |
+
if use_cache:
|
| 876 |
+
logger.warning_once(
|
| 877 |
+
"`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
|
| 878 |
+
)
|
| 879 |
+
use_cache = False
|
| 880 |
+
|
| 881 |
+
all_self_attentions = None
|
| 882 |
+
all_hidden_states = () if output_hidden_states else None
|
| 883 |
+
for index in range(self.num_layers):
|
| 884 |
+
if output_hidden_states:
|
| 885 |
+
all_hidden_states = all_hidden_states + (hidden_states,)
|
| 886 |
+
|
| 887 |
+
layer = self._get_layer(index)
|
| 888 |
+
if self.gradient_checkpointing and self.training:
|
| 889 |
+
layer_ret = torch.utils.checkpoint.checkpoint(
|
| 890 |
+
layer,
|
| 891 |
+
hidden_states,
|
| 892 |
+
attention_mask,
|
| 893 |
+
rotary_pos_emb,
|
| 894 |
+
kv_caches[index],
|
| 895 |
+
use_cache
|
| 896 |
+
)
|
| 897 |
+
else:
|
| 898 |
+
layer_ret = layer(
|
| 899 |
+
hidden_states,
|
| 900 |
+
attention_mask,
|
| 901 |
+
rotary_pos_emb,
|
| 902 |
+
kv_cache=kv_caches[index],
|
| 903 |
+
use_cache=use_cache
|
| 904 |
+
)
|
| 905 |
+
hidden_states, kv_cache = layer_ret
|
| 906 |
+
if use_cache:
|
| 907 |
+
presents = presents + (kv_cache,)
|
| 908 |
+
|
| 909 |
+
if output_hidden_states:
|
| 910 |
+
all_hidden_states = all_hidden_states + (hidden_states,)
|
| 911 |
+
|
| 912 |
+
# Final layer norm.
|
| 913 |
+
if self.post_layer_norm:
|
| 914 |
+
hidden_states = self.final_layernorm(hidden_states)
|
| 915 |
+
|
| 916 |
+
return hidden_states, presents, all_hidden_states, all_self_attentions
|
| 917 |
+
|
| 918 |
+
|
| 919 |
+
class ChatGLMPreTrainedModel(PreTrainedModel):
|
| 920 |
+
"""
|
| 921 |
+
An abstract class to handle weights initialization and
|
| 922 |
+
a simple interface for downloading and loading pretrained models.
|
| 923 |
+
"""
|
| 924 |
+
|
| 925 |
+
is_parallelizable = False
|
| 926 |
+
supports_gradient_checkpointing = True
|
| 927 |
+
config_class = ChatGLMConfig
|
| 928 |
+
base_model_prefix = "transformer"
|
| 929 |
+
_no_split_modules = ["GLMBlock"]
|
| 930 |
+
|
| 931 |
+
def _init_weights(self, module: nn.Module):
|
| 932 |
+
"""Initialize the weights."""
|
| 933 |
+
return
|
| 934 |
+
|
| 935 |
+
def get_masks(self, input_ids, past_key_values, padding_mask=None):
|
| 936 |
+
batch_size, seq_length = input_ids.shape
|
| 937 |
+
full_attention_mask = torch.ones(batch_size, seq_length, seq_length, device=input_ids.device)
|
| 938 |
+
full_attention_mask.tril_()
|
| 939 |
+
past_length = 0
|
| 940 |
+
if past_key_values:
|
| 941 |
+
past_length = past_key_values[0][0].shape[0]
|
| 942 |
+
if past_length:
|
| 943 |
+
full_attention_mask = torch.cat((torch.ones(batch_size, seq_length, past_length,
|
| 944 |
+
device=input_ids.device), full_attention_mask), dim=-1)
|
| 945 |
+
if padding_mask is not None:
|
| 946 |
+
full_attention_mask = full_attention_mask * padding_mask.unsqueeze(1)
|
| 947 |
+
if not past_length and padding_mask is not None:
|
| 948 |
+
full_attention_mask -= padding_mask.unsqueeze(-1) - 1
|
| 949 |
+
full_attention_mask = (full_attention_mask < 0.5).bool()
|
| 950 |
+
full_attention_mask.unsqueeze_(1)
|
| 951 |
+
return full_attention_mask
|
| 952 |
+
|
| 953 |
+
def get_position_ids(self, input_ids, device):
|
| 954 |
+
batch_size, seq_length = input_ids.shape
|
| 955 |
+
position_ids = torch.arange(seq_length, dtype=torch.long, device=device).unsqueeze(0).repeat(batch_size, 1)
|
| 956 |
+
return position_ids
|
| 957 |
+
|
| 958 |
+
def _set_gradient_checkpointing(self, module, value=False):
|
| 959 |
+
if isinstance(module, GLMTransformer):
|
| 960 |
+
module.gradient_checkpointing = value
|
| 961 |
+
|
| 962 |
+
|
| 963 |
+
class Embedding(torch.nn.Module):
|
| 964 |
+
"""Language model embeddings."""
|
| 965 |
+
|
| 966 |
+
def __init__(self, config: ChatGLMConfig, device=None):
|
| 967 |
+
super(Embedding, self).__init__()
|
| 968 |
+
|
| 969 |
+
self.hidden_size = config.hidden_size
|
| 970 |
+
# Word embeddings (parallel).
|
| 971 |
+
self.word_embeddings = nn.Embedding(
|
| 972 |
+
config.padded_vocab_size,
|
| 973 |
+
self.hidden_size,
|
| 974 |
+
dtype=config.torch_dtype,
|
| 975 |
+
device=device
|
| 976 |
+
)
|
| 977 |
+
self.fp32_residual_connection = config.fp32_residual_connection
|
| 978 |
+
|
| 979 |
+
def forward(self, input_ids):
|
| 980 |
+
# Embeddings.
|
| 981 |
+
words_embeddings = self.word_embeddings(input_ids)
|
| 982 |
+
embeddings = words_embeddings
|
| 983 |
+
# Data format change to avoid explicit transposes : [b s h] --> [s b h].
|
| 984 |
+
embeddings = embeddings.transpose(0, 1).contiguous()
|
| 985 |
+
# If the input flag for fp32 residual connection is set, convert for float.
|
| 986 |
+
if self.fp32_residual_connection:
|
| 987 |
+
embeddings = embeddings.float()
|
| 988 |
+
return embeddings
|
| 989 |
+
|
| 990 |
+
|
| 991 |
+
class ChatGLMModel(ChatGLMPreTrainedModel):
|
| 992 |
+
def __init__(self, config: ChatGLMConfig, device=None, empty_init=True):
|
| 993 |
+
super().__init__(config)
|
| 994 |
+
if empty_init:
|
| 995 |
+
init_method = skip_init
|
| 996 |
+
else:
|
| 997 |
+
init_method = default_init
|
| 998 |
+
init_kwargs = {}
|
| 999 |
+
if device is not None:
|
| 1000 |
+
init_kwargs["device"] = device
|
| 1001 |
+
self.embedding = init_method(Embedding, config, **init_kwargs)
|
| 1002 |
+
self.num_layers = config.num_layers
|
| 1003 |
+
self.multi_query_group_num = config.multi_query_group_num
|
| 1004 |
+
self.kv_channels = config.kv_channels
|
| 1005 |
+
|
| 1006 |
+
# Rotary positional embeddings
|
| 1007 |
+
self.seq_length = config.seq_length
|
| 1008 |
+
rotary_dim = (
|
| 1009 |
+
config.hidden_size // config.num_attention_heads if config.kv_channels is None else config.kv_channels
|
| 1010 |
+
)
|
| 1011 |
+
|
| 1012 |
+
self.rotary_pos_emb = RotaryEmbedding(rotary_dim // 2, original_impl=config.original_rope, device=device,
|
| 1013 |
+
dtype=config.torch_dtype)
|
| 1014 |
+
self.encoder = init_method(GLMTransformer, config, **init_kwargs)
|
| 1015 |
+
self.output_layer = init_method(nn.Linear, config.hidden_size, config.padded_vocab_size, bias=False,
|
| 1016 |
+
dtype=config.torch_dtype, **init_kwargs)
|
| 1017 |
+
self.pre_seq_len = config.pre_seq_len
|
| 1018 |
+
self.prefix_projection = config.prefix_projection
|
| 1019 |
+
if self.pre_seq_len is not None:
|
| 1020 |
+
for param in self.parameters():
|
| 1021 |
+
param.requires_grad = False
|
| 1022 |
+
self.prefix_tokens = torch.arange(self.pre_seq_len).long()
|
| 1023 |
+
self.prefix_encoder = PrefixEncoder(config)
|
| 1024 |
+
self.dropout = torch.nn.Dropout(0.1)
|
| 1025 |
+
|
| 1026 |
+
def get_input_embeddings(self):
|
| 1027 |
+
return self.embedding.word_embeddings
|
| 1028 |
+
|
| 1029 |
+
def get_prompt(self, batch_size, device, dtype=torch.half):
|
| 1030 |
+
prefix_tokens = self.prefix_tokens.unsqueeze(0).expand(batch_size, -1).to(device)
|
| 1031 |
+
past_key_values = self.prefix_encoder(prefix_tokens).type(dtype)
|
| 1032 |
+
past_key_values = past_key_values.view(
|
| 1033 |
+
batch_size,
|
| 1034 |
+
self.pre_seq_len,
|
| 1035 |
+
self.num_layers * 2,
|
| 1036 |
+
self.multi_query_group_num,
|
| 1037 |
+
self.kv_channels
|
| 1038 |
+
)
|
| 1039 |
+
# seq_len, b, nh, hidden_size
|
| 1040 |
+
past_key_values = self.dropout(past_key_values)
|
| 1041 |
+
past_key_values = past_key_values.permute([2, 1, 0, 3, 4]).split(2)
|
| 1042 |
+
return past_key_values
|
| 1043 |
+
|
| 1044 |
+
def forward(
|
| 1045 |
+
self,
|
| 1046 |
+
input_ids,
|
| 1047 |
+
position_ids: Optional[torch.Tensor] = None,
|
| 1048 |
+
attention_mask: Optional[torch.BoolTensor] = None,
|
| 1049 |
+
full_attention_mask: Optional[torch.BoolTensor] = None,
|
| 1050 |
+
past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,
|
| 1051 |
+
inputs_embeds: Optional[torch.Tensor] = None,
|
| 1052 |
+
use_cache: Optional[bool] = None,
|
| 1053 |
+
output_hidden_states: Optional[bool] = None,
|
| 1054 |
+
return_dict: Optional[bool] = None,
|
| 1055 |
+
):
|
| 1056 |
+
output_hidden_states = (
|
| 1057 |
+
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
| 1058 |
+
)
|
| 1059 |
+
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
| 1060 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 1061 |
+
|
| 1062 |
+
batch_size, seq_length = input_ids.shape
|
| 1063 |
+
|
| 1064 |
+
if inputs_embeds is None:
|
| 1065 |
+
inputs_embeds = self.embedding(input_ids)
|
| 1066 |
+
|
| 1067 |
+
if self.pre_seq_len is not None:
|
| 1068 |
+
if past_key_values is None:
|
| 1069 |
+
past_key_values = self.get_prompt(batch_size=batch_size, device=input_ids.device,
|
| 1070 |
+
dtype=inputs_embeds.dtype)
|
| 1071 |
+
if attention_mask is not None:
|
| 1072 |
+
attention_mask = torch.cat([attention_mask.new_ones((batch_size, self.pre_seq_len)),
|
| 1073 |
+
attention_mask], dim=-1)
|
| 1074 |
+
|
| 1075 |
+
if full_attention_mask is None:
|
| 1076 |
+
if (attention_mask is not None and not attention_mask.all()) or (past_key_values and seq_length != 1):
|
| 1077 |
+
full_attention_mask = self.get_masks(input_ids, past_key_values, padding_mask=attention_mask)
|
| 1078 |
+
|
| 1079 |
+
# Rotary positional embeddings
|
| 1080 |
+
rotary_pos_emb = self.rotary_pos_emb(self.seq_length)
|
| 1081 |
+
if position_ids is not None:
|
| 1082 |
+
rotary_pos_emb = rotary_pos_emb[position_ids]
|
| 1083 |
+
else:
|
| 1084 |
+
rotary_pos_emb = rotary_pos_emb[None, :seq_length]
|
| 1085 |
+
rotary_pos_emb = rotary_pos_emb.transpose(0, 1).contiguous()
|
| 1086 |
+
|
| 1087 |
+
# Run encoder.
|
| 1088 |
+
hidden_states, presents, all_hidden_states, all_self_attentions = self.encoder(
|
| 1089 |
+
inputs_embeds, full_attention_mask, rotary_pos_emb=rotary_pos_emb,
|
| 1090 |
+
kv_caches=past_key_values, use_cache=use_cache, output_hidden_states=output_hidden_states
|
| 1091 |
+
)
|
| 1092 |
+
|
| 1093 |
+
if not return_dict:
|
| 1094 |
+
return tuple(v for v in [hidden_states, presents, all_hidden_states, all_self_attentions] if v is not None)
|
| 1095 |
+
|
| 1096 |
+
return BaseModelOutputWithPast(
|
| 1097 |
+
last_hidden_state=hidden_states,
|
| 1098 |
+
past_key_values=presents,
|
| 1099 |
+
hidden_states=all_hidden_states,
|
| 1100 |
+
attentions=all_self_attentions,
|
| 1101 |
+
)
|
| 1102 |
+
|
| 1103 |
+
def quantize(self, weight_bit_width: int):
|
| 1104 |
+
# from .quantization import quantize
|
| 1105 |
+
quantize(self.encoder, weight_bit_width)
|
| 1106 |
+
return self
|
| 1107 |
+
|
| 1108 |
+
|
| 1109 |
+
class ChatGLMForConditionalGeneration(ChatGLMPreTrainedModel):
|
| 1110 |
+
def __init__(self, config: ChatGLMConfig, empty_init=True, device=None):
|
| 1111 |
+
super().__init__(config)
|
| 1112 |
+
|
| 1113 |
+
self.max_sequence_length = config.max_length
|
| 1114 |
+
self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device)
|
| 1115 |
+
self.config = config
|
| 1116 |
+
self.quantized = False
|
| 1117 |
+
|
| 1118 |
+
if self.config.quantization_bit:
|
| 1119 |
+
self.quantize(self.config.quantization_bit, empty_init=True)
|
| 1120 |
+
|
| 1121 |
+
def _update_model_kwargs_for_generation(
|
| 1122 |
+
self,
|
| 1123 |
+
outputs: ModelOutput,
|
| 1124 |
+
model_kwargs: Dict[str, Any],
|
| 1125 |
+
is_encoder_decoder: bool = False,
|
| 1126 |
+
standardize_cache_format: bool = False,
|
| 1127 |
+
) -> Dict[str, Any]:
|
| 1128 |
+
# update past_key_values
|
| 1129 |
+
model_kwargs["past_key_values"] = self._extract_past_from_model_output(
|
| 1130 |
+
outputs, standardize_cache_format=standardize_cache_format
|
| 1131 |
+
)
|
| 1132 |
+
|
| 1133 |
+
# update attention mask
|
| 1134 |
+
if "attention_mask" in model_kwargs:
|
| 1135 |
+
attention_mask = model_kwargs["attention_mask"]
|
| 1136 |
+
model_kwargs["attention_mask"] = torch.cat(
|
| 1137 |
+
[attention_mask, attention_mask.new_ones((attention_mask.shape[0], 1))], dim=-1
|
| 1138 |
+
)
|
| 1139 |
+
|
| 1140 |
+
# update position ids
|
| 1141 |
+
if "position_ids" in model_kwargs:
|
| 1142 |
+
position_ids = model_kwargs["position_ids"]
|
| 1143 |
+
new_position_id = position_ids[..., -1:].clone()
|
| 1144 |
+
new_position_id += 1
|
| 1145 |
+
model_kwargs["position_ids"] = torch.cat(
|
| 1146 |
+
[position_ids, new_position_id], dim=-1
|
| 1147 |
+
)
|
| 1148 |
+
|
| 1149 |
+
model_kwargs["is_first_forward"] = False
|
| 1150 |
+
return model_kwargs
|
| 1151 |
+
|
| 1152 |
+
def prepare_inputs_for_generation(
|
| 1153 |
+
self,
|
| 1154 |
+
input_ids: torch.LongTensor,
|
| 1155 |
+
past_key_values: Optional[torch.Tensor] = None,
|
| 1156 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 1157 |
+
position_ids: Optional[torch.Tensor] = None,
|
| 1158 |
+
use_cache: Optional[bool] = None,
|
| 1159 |
+
is_first_forward: bool = True,
|
| 1160 |
+
**kwargs
|
| 1161 |
+
) -> dict:
|
| 1162 |
+
# only last token for input_ids if past is not None
|
| 1163 |
+
if position_ids is None:
|
| 1164 |
+
position_ids = self.get_position_ids(input_ids, device=input_ids.device)
|
| 1165 |
+
if not is_first_forward:
|
| 1166 |
+
if past_key_values is not None:
|
| 1167 |
+
position_ids = position_ids[..., -1:]
|
| 1168 |
+
input_ids = input_ids[:, -1:]
|
| 1169 |
+
return {
|
| 1170 |
+
"input_ids": input_ids,
|
| 1171 |
+
"past_key_values": past_key_values,
|
| 1172 |
+
"position_ids": position_ids,
|
| 1173 |
+
"attention_mask": attention_mask,
|
| 1174 |
+
"return_last_logit": True,
|
| 1175 |
+
"use_cache": use_cache
|
| 1176 |
+
}
|
| 1177 |
+
|
| 1178 |
+
def forward(
|
| 1179 |
+
self,
|
| 1180 |
+
input_ids: Optional[torch.Tensor] = None,
|
| 1181 |
+
position_ids: Optional[torch.Tensor] = None,
|
| 1182 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 1183 |
+
past_key_values: Optional[Tuple[torch.FloatTensor]] = None,
|
| 1184 |
+
inputs_embeds: Optional[torch.Tensor] = None,
|
| 1185 |
+
labels: Optional[torch.Tensor] = None,
|
| 1186 |
+
use_cache: Optional[bool] = None,
|
| 1187 |
+
output_attentions: Optional[bool] = None,
|
| 1188 |
+
output_hidden_states: Optional[bool] = None,
|
| 1189 |
+
return_dict: Optional[bool] = None,
|
| 1190 |
+
return_last_logit: Optional[bool] = False,
|
| 1191 |
+
):
|
| 1192 |
+
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
| 1193 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 1194 |
+
|
| 1195 |
+
transformer_outputs = self.transformer(
|
| 1196 |
+
input_ids=input_ids,
|
| 1197 |
+
position_ids=position_ids,
|
| 1198 |
+
attention_mask=attention_mask,
|
| 1199 |
+
past_key_values=past_key_values,
|
| 1200 |
+
inputs_embeds=inputs_embeds,
|
| 1201 |
+
use_cache=use_cache,
|
| 1202 |
+
output_hidden_states=output_hidden_states,
|
| 1203 |
+
return_dict=return_dict,
|
| 1204 |
+
)
|
| 1205 |
+
|
| 1206 |
+
hidden_states = transformer_outputs[0]
|
| 1207 |
+
if return_last_logit:
|
| 1208 |
+
hidden_states = hidden_states[-1:]
|
| 1209 |
+
lm_logits = self.transformer.output_layer(hidden_states)
|
| 1210 |
+
lm_logits = lm_logits.transpose(0, 1).contiguous()
|
| 1211 |
+
|
| 1212 |
+
loss = None
|
| 1213 |
+
if labels is not None:
|
| 1214 |
+
lm_logits = lm_logits.to(torch.float32)
|
| 1215 |
+
|
| 1216 |
+
# Shift so that tokens < n predict n
|
| 1217 |
+
shift_logits = lm_logits[..., :-1, :].contiguous()
|
| 1218 |
+
shift_labels = labels[..., 1:].contiguous()
|
| 1219 |
+
# Flatten the tokens
|
| 1220 |
+
loss_fct = CrossEntropyLoss(ignore_index=-100)
|
| 1221 |
+
loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
|
| 1222 |
+
|
| 1223 |
+
lm_logits = lm_logits.to(hidden_states.dtype)
|
| 1224 |
+
loss = loss.to(hidden_states.dtype)
|
| 1225 |
+
|
| 1226 |
+
if not return_dict:
|
| 1227 |
+
output = (lm_logits,) + transformer_outputs[1:]
|
| 1228 |
+
return ((loss,) + output) if loss is not None else output
|
| 1229 |
+
|
| 1230 |
+
return CausalLMOutputWithPast(
|
| 1231 |
+
loss=loss,
|
| 1232 |
+
logits=lm_logits,
|
| 1233 |
+
past_key_values=transformer_outputs.past_key_values,
|
| 1234 |
+
hidden_states=transformer_outputs.hidden_states,
|
| 1235 |
+
attentions=transformer_outputs.attentions,
|
| 1236 |
+
)
|
| 1237 |
+
|
| 1238 |
+
@staticmethod
|
| 1239 |
+
def _reorder_cache(
|
| 1240 |
+
past: Tuple[Tuple[torch.Tensor, torch.Tensor], ...], beam_idx: torch.LongTensor
|
| 1241 |
+
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], ...]:
|
| 1242 |
+
"""
|
| 1243 |
+
This function is used to re-order the `past_key_values` cache if [`~PreTrainedModel.beam_search`] or
|
| 1244 |
+
[`~PreTrainedModel.beam_sample`] is called. This is required to match `past_key_values` with the correct
|
| 1245 |
+
beam_idx at every generation step.
|
| 1246 |
+
|
| 1247 |
+
Output shares the same memory storage as `past`.
|
| 1248 |
+
"""
|
| 1249 |
+
return tuple(
|
| 1250 |
+
(
|
| 1251 |
+
layer_past[0].index_select(1, beam_idx.to(layer_past[0].device)),
|
| 1252 |
+
layer_past[1].index_select(1, beam_idx.to(layer_past[1].device)),
|
| 1253 |
+
)
|
| 1254 |
+
for layer_past in past
|
| 1255 |
+
)
|
| 1256 |
+
|
| 1257 |
+
def process_response(self, output, history):
|
| 1258 |
+
content = ""
|
| 1259 |
+
history = deepcopy(history)
|
| 1260 |
+
for response in output.split("<|assistant|>"):
|
| 1261 |
+
metadata, content = response.split("\n", maxsplit=1)
|
| 1262 |
+
if not metadata.strip():
|
| 1263 |
+
content = content.strip()
|
| 1264 |
+
history.append({"role": "assistant", "metadata": metadata, "content": content})
|
| 1265 |
+
content = content.replace("[[训练时间]]", "2023年")
|
| 1266 |
+
else:
|
| 1267 |
+
history.append({"role": "assistant", "metadata": metadata, "content": content})
|
| 1268 |
+
if history[0]["role"] == "system" and "tools" in history[0]:
|
| 1269 |
+
content = "\n".join(content.split("\n")[1:-1])
|
| 1270 |
+
def tool_call(**kwargs):
|
| 1271 |
+
return kwargs
|
| 1272 |
+
parameters = eval(content)
|
| 1273 |
+
content = {"name": metadata.strip(), "parameters": parameters}
|
| 1274 |
+
else:
|
| 1275 |
+
content = {"name": metadata.strip(), "content": content}
|
| 1276 |
+
return content, history
|
| 1277 |
+
|
| 1278 |
+
@torch.inference_mode()
|
| 1279 |
+
def chat(self, tokenizer, query: str, history: List[Tuple[str, str]] = None, role: str = "user",
|
| 1280 |
+
max_length: int = 8192, num_beams=1, do_sample=True, top_p=0.8, temperature=0.8, logits_processor=None,
|
| 1281 |
+
**kwargs):
|
| 1282 |
+
if history is None:
|
| 1283 |
+
history = []
|
| 1284 |
+
if logits_processor is None:
|
| 1285 |
+
logits_processor = LogitsProcessorList()
|
| 1286 |
+
logits_processor.append(InvalidScoreLogitsProcessor())
|
| 1287 |
+
gen_kwargs = {"max_length": max_length, "num_beams": num_beams, "do_sample": do_sample, "top_p": top_p,
|
| 1288 |
+
"temperature": temperature, "logits_processor": logits_processor, **kwargs}
|
| 1289 |
+
inputs = tokenizer.build_chat_input(query, history=history, role=role)
|
| 1290 |
+
inputs = inputs.to(self.device)
|
| 1291 |
+
eos_token_id = [tokenizer.eos_token_id, tokenizer.get_command("<|user|>"),
|
| 1292 |
+
tokenizer.get_command("<|observation|>")]
|
| 1293 |
+
outputs = self.generate(**inputs, **gen_kwargs, eos_token_id=eos_token_id)
|
| 1294 |
+
outputs = outputs.tolist()[0][len(inputs["input_ids"][0]):-1]
|
| 1295 |
+
response = tokenizer.decode(outputs)
|
| 1296 |
+
history.append({"role": role, "content": query})
|
| 1297 |
+
response, history = self.process_response(response, history)
|
| 1298 |
+
return response, history
|
| 1299 |
+
|
| 1300 |
+
@torch.inference_mode()
|
| 1301 |
+
def stream_chat(self, tokenizer, query: str, history: List[Tuple[str, str]] = None, role: str = "user",
|
| 1302 |
+
past_key_values=None,max_length: int = 8192, do_sample=True, top_p=0.8, temperature=0.8,
|
| 1303 |
+
logits_processor=None, return_past_key_values=False, **kwargs):
|
| 1304 |
+
if history is None:
|
| 1305 |
+
history = []
|
| 1306 |
+
if logits_processor is None:
|
| 1307 |
+
logits_processor = LogitsProcessorList()
|
| 1308 |
+
logits_processor.append(InvalidScoreLogitsProcessor())
|
| 1309 |
+
eos_token_id = [tokenizer.eos_token_id, tokenizer.get_command("<|user|>"),
|
| 1310 |
+
tokenizer.get_command("<|observation|>")]
|
| 1311 |
+
gen_kwargs = {"max_length": max_length, "do_sample": do_sample, "top_p": top_p,
|
| 1312 |
+
"temperature": temperature, "logits_processor": logits_processor, **kwargs}
|
| 1313 |
+
if past_key_values is None:
|
| 1314 |
+
inputs = tokenizer.build_chat_input(query, history=history, role=role)
|
| 1315 |
+
else:
|
| 1316 |
+
inputs = tokenizer.build_chat_input(query, role=role)
|
| 1317 |
+
inputs = inputs.to(self.device)
|
| 1318 |
+
if past_key_values is not None:
|
| 1319 |
+
past_length = past_key_values[0][0].shape[0]
|
| 1320 |
+
if self.transformer.pre_seq_len is not None:
|
| 1321 |
+
past_length -= self.transformer.pre_seq_len
|
| 1322 |
+
inputs.position_ids += past_length
|
| 1323 |
+
attention_mask = inputs.attention_mask
|
| 1324 |
+
attention_mask = torch.cat((attention_mask.new_ones(1, past_length), attention_mask), dim=1)
|
| 1325 |
+
inputs['attention_mask'] = attention_mask
|
| 1326 |
+
history.append({"role": role, "content": query})
|
| 1327 |
+
for outputs in self.stream_generate(**inputs, past_key_values=past_key_values,
|
| 1328 |
+
eos_token_id=eos_token_id, return_past_key_values=return_past_key_values,
|
| 1329 |
+
**gen_kwargs):
|
| 1330 |
+
if return_past_key_values:
|
| 1331 |
+
outputs, past_key_values = outputs
|
| 1332 |
+
outputs = outputs.tolist()[0][len(inputs["input_ids"][0]):-1]
|
| 1333 |
+
response = tokenizer.decode(outputs)
|
| 1334 |
+
if response and response[-1] != "�":
|
| 1335 |
+
response, new_history = self.process_response(response, history)
|
| 1336 |
+
if return_past_key_values:
|
| 1337 |
+
yield response, new_history, past_key_values
|
| 1338 |
+
else:
|
| 1339 |
+
yield response, new_history
|
| 1340 |
+
|
| 1341 |
+
@torch.inference_mode()
|
| 1342 |
+
def stream_generate(
|
| 1343 |
+
self,
|
| 1344 |
+
input_ids,
|
| 1345 |
+
generation_config: Optional[GenerationConfig] = None,
|
| 1346 |
+
logits_processor: Optional[LogitsProcessorList] = None,
|
| 1347 |
+
stopping_criteria: Optional[StoppingCriteriaList] = None,
|
| 1348 |
+
prefix_allowed_tokens_fn: Optional[Callable[[int, torch.Tensor], List[int]]] = None,
|
| 1349 |
+
return_past_key_values=False,
|
| 1350 |
+
**kwargs,
|
| 1351 |
+
):
|
| 1352 |
+
batch_size, input_ids_seq_length = input_ids.shape[0], input_ids.shape[-1]
|
| 1353 |
+
|
| 1354 |
+
if generation_config is None:
|
| 1355 |
+
generation_config = self.generation_config
|
| 1356 |
+
generation_config = copy.deepcopy(generation_config)
|
| 1357 |
+
model_kwargs = generation_config.update(**kwargs)
|
| 1358 |
+
model_kwargs["use_cache"] = generation_config.use_cache
|
| 1359 |
+
bos_token_id, eos_token_id = generation_config.bos_token_id, generation_config.eos_token_id
|
| 1360 |
+
|
| 1361 |
+
if isinstance(eos_token_id, int):
|
| 1362 |
+
eos_token_id = [eos_token_id]
|
| 1363 |
+
eos_token_id_tensor = torch.tensor(eos_token_id).to(input_ids.device) if eos_token_id is not None else None
|
| 1364 |
+
|
| 1365 |
+
has_default_max_length = kwargs.get("max_length") is None and generation_config.max_length is not None
|
| 1366 |
+
if has_default_max_length and generation_config.max_new_tokens is None:
|
| 1367 |
+
warnings.warn(
|
| 1368 |
+
f"Using `max_length`'s default ({generation_config.max_length}) to control the generation length. "
|
| 1369 |
+
"This behaviour is deprecated and will be removed from the config in v5 of Transformers -- we"
|
| 1370 |
+
" recommend using `max_new_tokens` to control the maximum length of the generation.",
|
| 1371 |
+
UserWarning,
|
| 1372 |
+
)
|
| 1373 |
+
elif generation_config.max_new_tokens is not None:
|
| 1374 |
+
generation_config.max_length = generation_config.max_new_tokens + input_ids_seq_length
|
| 1375 |
+
if not has_default_max_length:
|
| 1376 |
+
logger.warning(
|
| 1377 |
+
f"Both `max_new_tokens` (={generation_config.max_new_tokens}) and `max_length`(="
|
| 1378 |
+
f"{generation_config.max_length}) seem to have been set. `max_new_tokens` will take precedence. "
|
| 1379 |
+
"Please refer to the documentation for more information. "
|
| 1380 |
+
"(https://huggingface.co/docs/transformers/main/en/main_classes/text_generation)",
|
| 1381 |
+
UserWarning,
|
| 1382 |
+
)
|
| 1383 |
+
|
| 1384 |
+
if input_ids_seq_length >= generation_config.max_length:
|
| 1385 |
+
input_ids_string = "decoder_input_ids" if self.config.is_encoder_decoder else "input_ids"
|
| 1386 |
+
logger.warning(
|
| 1387 |
+
f"Input length of {input_ids_string} is {input_ids_seq_length}, but `max_length` is set to"
|
| 1388 |
+
f" {generation_config.max_length}. This can lead to unexpected behavior. You should consider"
|
| 1389 |
+
" increasing `max_new_tokens`."
|
| 1390 |
+
)
|
| 1391 |
+
|
| 1392 |
+
# 2. Set generation parameters if not already defined
|
| 1393 |
+
logits_processor = logits_processor if logits_processor is not None else LogitsProcessorList()
|
| 1394 |
+
stopping_criteria = stopping_criteria if stopping_criteria is not None else StoppingCriteriaList()
|
| 1395 |
+
|
| 1396 |
+
logits_processor = self._get_logits_processor(
|
| 1397 |
+
generation_config=generation_config,
|
| 1398 |
+
input_ids_seq_length=input_ids_seq_length,
|
| 1399 |
+
encoder_input_ids=input_ids,
|
| 1400 |
+
prefix_allowed_tokens_fn=prefix_allowed_tokens_fn,
|
| 1401 |
+
logits_processor=logits_processor,
|
| 1402 |
+
)
|
| 1403 |
+
|
| 1404 |
+
stopping_criteria = self._get_stopping_criteria(
|
| 1405 |
+
generation_config=generation_config, stopping_criteria=stopping_criteria
|
| 1406 |
+
)
|
| 1407 |
+
logits_warper = self._get_logits_warper(generation_config)
|
| 1408 |
+
|
| 1409 |
+
unfinished_sequences = input_ids.new(input_ids.shape[0]).fill_(1)
|
| 1410 |
+
scores = None
|
| 1411 |
+
while True:
|
| 1412 |
+
model_inputs = self.prepare_inputs_for_generation(input_ids, **model_kwargs)
|
| 1413 |
+
# forward pass to get next token
|
| 1414 |
+
outputs = self(
|
| 1415 |
+
**model_inputs,
|
| 1416 |
+
return_dict=True,
|
| 1417 |
+
output_attentions=False,
|
| 1418 |
+
output_hidden_states=False,
|
| 1419 |
+
)
|
| 1420 |
+
|
| 1421 |
+
next_token_logits = outputs.logits[:, -1, :]
|
| 1422 |
+
|
| 1423 |
+
# pre-process distribution
|
| 1424 |
+
next_token_scores = logits_processor(input_ids, next_token_logits)
|
| 1425 |
+
next_token_scores = logits_warper(input_ids, next_token_scores)
|
| 1426 |
+
|
| 1427 |
+
# sample
|
| 1428 |
+
probs = nn.functional.softmax(next_token_scores, dim=-1)
|
| 1429 |
+
if generation_config.do_sample:
|
| 1430 |
+
next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)
|
| 1431 |
+
else:
|
| 1432 |
+
next_tokens = torch.argmax(probs, dim=-1)
|
| 1433 |
+
# update generated ids, model inputs, and length for next step
|
| 1434 |
+
input_ids = torch.cat([input_ids, next_tokens[:, None]], dim=-1)
|
| 1435 |
+
model_kwargs = self._update_model_kwargs_for_generation(
|
| 1436 |
+
outputs, model_kwargs, is_encoder_decoder=self.config.is_encoder_decoder
|
| 1437 |
+
)
|
| 1438 |
+
unfinished_sequences = unfinished_sequences.mul(
|
| 1439 |
+
next_tokens.tile(eos_token_id_tensor.shape[0], 1).ne(eos_token_id_tensor.unsqueeze(1)).prod(dim=0)
|
| 1440 |
+
)
|
| 1441 |
+
if return_past_key_values:
|
| 1442 |
+
yield input_ids, outputs.past_key_values
|
| 1443 |
+
else:
|
| 1444 |
+
yield input_ids
|
| 1445 |
+
# stop when each sentence is finished, or if we exceed the maximum length
|
| 1446 |
+
if unfinished_sequences.max() == 0 or stopping_criteria(input_ids, scores):
|
| 1447 |
+
break
|
| 1448 |
+
|
| 1449 |
+
def quantize(self, bits: int, empty_init=False, device=None, **kwargs):
|
| 1450 |
+
if bits == 0:
|
| 1451 |
+
return
|
| 1452 |
+
|
| 1453 |
+
# from .quantization import quantize
|
| 1454 |
+
|
| 1455 |
+
if self.quantized:
|
| 1456 |
+
logger.info("Already quantized.")
|
| 1457 |
+
return self
|
| 1458 |
+
|
| 1459 |
+
self.quantized = True
|
| 1460 |
+
|
| 1461 |
+
self.config.quantization_bit = bits
|
| 1462 |
+
|
| 1463 |
+
self.transformer.encoder = quantize(self.transformer.encoder, bits, empty_init=empty_init, device=device,
|
| 1464 |
+
**kwargs)
|
| 1465 |
+
return self
|
| 1466 |
+
|
| 1467 |
+
|
| 1468 |
+
class ChatGLMForSequenceClassification(ChatGLMPreTrainedModel):
|
| 1469 |
+
def __init__(self, config: ChatGLMConfig, empty_init=True, device=None):
|
| 1470 |
+
super().__init__(config)
|
| 1471 |
+
|
| 1472 |
+
self.num_labels = config.num_labels
|
| 1473 |
+
self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device)
|
| 1474 |
+
|
| 1475 |
+
self.classifier_head = nn.Linear(config.hidden_size, config.num_labels, bias=True, dtype=torch.half)
|
| 1476 |
+
if config.classifier_dropout is not None:
|
| 1477 |
+
self.dropout = nn.Dropout(config.classifier_dropout)
|
| 1478 |
+
else:
|
| 1479 |
+
self.dropout = None
|
| 1480 |
+
self.config = config
|
| 1481 |
+
|
| 1482 |
+
if self.config.quantization_bit:
|
| 1483 |
+
self.quantize(self.config.quantization_bit, empty_init=True)
|
| 1484 |
+
|
| 1485 |
+
def forward(
|
| 1486 |
+
self,
|
| 1487 |
+
input_ids: Optional[torch.LongTensor] = None,
|
| 1488 |
+
position_ids: Optional[torch.LongTensor] = None,
|
| 1489 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 1490 |
+
full_attention_mask: Optional[torch.Tensor] = None,
|
| 1491 |
+
past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None,
|
| 1492 |
+
inputs_embeds: Optional[torch.LongTensor] = None,
|
| 1493 |
+
labels: Optional[torch.LongTensor] = None,
|
| 1494 |
+
use_cache: Optional[bool] = None,
|
| 1495 |
+
output_hidden_states: Optional[bool] = None,
|
| 1496 |
+
return_dict: Optional[bool] = None,
|
| 1497 |
+
) -> Union[Tuple[torch.Tensor, ...], SequenceClassifierOutputWithPast]:
|
| 1498 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 1499 |
+
|
| 1500 |
+
transformer_outputs = self.transformer(
|
| 1501 |
+
input_ids=input_ids,
|
| 1502 |
+
position_ids=position_ids,
|
| 1503 |
+
attention_mask=attention_mask,
|
| 1504 |
+
full_attention_mask=full_attention_mask,
|
| 1505 |
+
past_key_values=past_key_values,
|
| 1506 |
+
inputs_embeds=inputs_embeds,
|
| 1507 |
+
use_cache=use_cache,
|
| 1508 |
+
output_hidden_states=output_hidden_states,
|
| 1509 |
+
return_dict=return_dict,
|
| 1510 |
+
)
|
| 1511 |
+
|
| 1512 |
+
hidden_states = transformer_outputs[0]
|
| 1513 |
+
pooled_hidden_states = hidden_states[-1]
|
| 1514 |
+
if self.dropout is not None:
|
| 1515 |
+
pooled_hidden_states = self.dropout(pooled_hidden_states)
|
| 1516 |
+
logits = self.classifier_head(pooled_hidden_states)
|
| 1517 |
+
|
| 1518 |
+
loss = None
|
| 1519 |
+
if labels is not None:
|
| 1520 |
+
if self.config.problem_type is None:
|
| 1521 |
+
if self.num_labels == 1:
|
| 1522 |
+
self.config.problem_type = "regression"
|
| 1523 |
+
elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):
|
| 1524 |
+
self.config.problem_type = "single_label_classification"
|
| 1525 |
+
else:
|
| 1526 |
+
self.config.problem_type = "multi_label_classification"
|
| 1527 |
+
|
| 1528 |
+
if self.config.problem_type == "regression":
|
| 1529 |
+
loss_fct = MSELoss()
|
| 1530 |
+
if self.num_labels == 1:
|
| 1531 |
+
loss = loss_fct(logits.squeeze().float(), labels.squeeze())
|
| 1532 |
+
else:
|
| 1533 |
+
loss = loss_fct(logits.float(), labels)
|
| 1534 |
+
elif self.config.problem_type == "single_label_classification":
|
| 1535 |
+
loss_fct = CrossEntropyLoss()
|
| 1536 |
+
loss = loss_fct(logits.view(-1, self.num_labels).float(), labels.view(-1))
|
| 1537 |
+
elif self.config.problem_type == "multi_label_classification":
|
| 1538 |
+
loss_fct = BCEWithLogitsLoss()
|
| 1539 |
+
loss = loss_fct(logits.float(), labels.view(-1, self.num_labels))
|
| 1540 |
+
|
| 1541 |
+
if not return_dict:
|
| 1542 |
+
output = (logits,) + transformer_outputs[1:]
|
| 1543 |
+
return ((loss,) + output) if loss is not None else output
|
| 1544 |
+
|
| 1545 |
+
return SequenceClassifierOutputWithPast(
|
| 1546 |
+
loss=loss,
|
| 1547 |
+
logits=logits,
|
| 1548 |
+
past_key_values=transformer_outputs.past_key_values,
|
| 1549 |
+
hidden_states=transformer_outputs.hidden_states,
|
| 1550 |
+
attentions=transformer_outputs.attentions,
|
| 1551 |
+
)
|
diffsynth/models/lora.py
ADDED
|
@@ -0,0 +1,387 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_unet import SDUNet
|
| 3 |
+
from .sdxl_unet import SDXLUNet
|
| 4 |
+
from .sd_text_encoder import SDTextEncoder
|
| 5 |
+
from .sdxl_text_encoder import SDXLTextEncoder, SDXLTextEncoder2
|
| 6 |
+
from .sd3_dit import SD3DiT
|
| 7 |
+
from .flux_dit import FluxDiT
|
| 8 |
+
from .hunyuan_dit import HunyuanDiT
|
| 9 |
+
from .cog_dit import CogDiT
|
| 10 |
+
from .hunyuan_video_dit import HunyuanVideoDiT
|
| 11 |
+
from .wan_video_dit import WanModel
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class LoRAFromCivitai:
|
| 16 |
+
def __init__(self):
|
| 17 |
+
self.supported_model_classes = []
|
| 18 |
+
self.lora_prefix = []
|
| 19 |
+
self.renamed_lora_prefix = {}
|
| 20 |
+
self.special_keys = {}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def convert_state_dict(self, state_dict, lora_prefix="lora_unet_", alpha=1.0):
|
| 24 |
+
for key in state_dict:
|
| 25 |
+
if ".lora_up" in key:
|
| 26 |
+
return self.convert_state_dict_up_down(state_dict, lora_prefix, alpha)
|
| 27 |
+
return self.convert_state_dict_AB(state_dict, lora_prefix, alpha)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def convert_state_dict_up_down(self, state_dict, lora_prefix="lora_unet_", alpha=1.0):
|
| 31 |
+
renamed_lora_prefix = self.renamed_lora_prefix.get(lora_prefix, "")
|
| 32 |
+
state_dict_ = {}
|
| 33 |
+
for key in state_dict:
|
| 34 |
+
if ".lora_up" not in key:
|
| 35 |
+
continue
|
| 36 |
+
if not key.startswith(lora_prefix):
|
| 37 |
+
continue
|
| 38 |
+
weight_up = state_dict[key].to(device="cuda", dtype=torch.float16)
|
| 39 |
+
weight_down = state_dict[key.replace(".lora_up", ".lora_down")].to(device="cuda", dtype=torch.float16)
|
| 40 |
+
if len(weight_up.shape) == 4:
|
| 41 |
+
weight_up = weight_up.squeeze(3).squeeze(2).to(torch.float32)
|
| 42 |
+
weight_down = weight_down.squeeze(3).squeeze(2).to(torch.float32)
|
| 43 |
+
lora_weight = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)
|
| 44 |
+
else:
|
| 45 |
+
lora_weight = alpha * torch.mm(weight_up, weight_down)
|
| 46 |
+
target_name = key.split(".")[0].replace(lora_prefix, renamed_lora_prefix).replace("_", ".") + ".weight"
|
| 47 |
+
for special_key in self.special_keys:
|
| 48 |
+
target_name = target_name.replace(special_key, self.special_keys[special_key])
|
| 49 |
+
state_dict_[target_name] = lora_weight.cpu()
|
| 50 |
+
return state_dict_
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def convert_state_dict_AB(self, state_dict, lora_prefix="", alpha=1.0, device="cuda", torch_dtype=torch.float16):
|
| 54 |
+
state_dict_ = {}
|
| 55 |
+
for key in state_dict:
|
| 56 |
+
if ".lora_B." not in key:
|
| 57 |
+
continue
|
| 58 |
+
if not key.startswith(lora_prefix):
|
| 59 |
+
continue
|
| 60 |
+
weight_up = state_dict[key].to(device=device, dtype=torch_dtype)
|
| 61 |
+
weight_down = state_dict[key.replace(".lora_B.", ".lora_A.")].to(device=device, dtype=torch_dtype)
|
| 62 |
+
if len(weight_up.shape) == 4:
|
| 63 |
+
weight_up = weight_up.squeeze(3).squeeze(2)
|
| 64 |
+
weight_down = weight_down.squeeze(3).squeeze(2)
|
| 65 |
+
lora_weight = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)
|
| 66 |
+
else:
|
| 67 |
+
lora_weight = alpha * torch.mm(weight_up, weight_down)
|
| 68 |
+
keys = key.split(".")
|
| 69 |
+
keys.pop(keys.index("lora_B"))
|
| 70 |
+
target_name = ".".join(keys)
|
| 71 |
+
target_name = target_name[len(lora_prefix):]
|
| 72 |
+
state_dict_[target_name] = lora_weight.cpu()
|
| 73 |
+
return state_dict_
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def load(self, model, state_dict_lora, lora_prefix, alpha=1.0, model_resource=None):
|
| 77 |
+
state_dict_model = model.state_dict()
|
| 78 |
+
state_dict_lora = self.convert_state_dict(state_dict_lora, lora_prefix=lora_prefix, alpha=alpha)
|
| 79 |
+
if model_resource == "diffusers":
|
| 80 |
+
state_dict_lora = model.__class__.state_dict_converter().from_diffusers(state_dict_lora)
|
| 81 |
+
elif model_resource == "civitai":
|
| 82 |
+
state_dict_lora = model.__class__.state_dict_converter().from_civitai(state_dict_lora)
|
| 83 |
+
if isinstance(state_dict_lora, tuple):
|
| 84 |
+
state_dict_lora = state_dict_lora[0]
|
| 85 |
+
if len(state_dict_lora) > 0:
|
| 86 |
+
print(f" {len(state_dict_lora)} tensors are updated.")
|
| 87 |
+
for name in state_dict_lora:
|
| 88 |
+
fp8=False
|
| 89 |
+
if state_dict_model[name].dtype == torch.float8_e4m3fn:
|
| 90 |
+
state_dict_model[name]= state_dict_model[name].to(state_dict_lora[name].dtype)
|
| 91 |
+
fp8=True
|
| 92 |
+
state_dict_model[name] += state_dict_lora[name].to(
|
| 93 |
+
dtype=state_dict_model[name].dtype, device=state_dict_model[name].device)
|
| 94 |
+
if fp8:
|
| 95 |
+
state_dict_model[name] = state_dict_model[name].to(torch.float8_e4m3fn)
|
| 96 |
+
model.load_state_dict(state_dict_model)
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def match(self, model, state_dict_lora):
|
| 100 |
+
for lora_prefix, model_class in zip(self.lora_prefix, self.supported_model_classes):
|
| 101 |
+
if not isinstance(model, model_class):
|
| 102 |
+
continue
|
| 103 |
+
state_dict_model = model.state_dict()
|
| 104 |
+
for model_resource in ["diffusers", "civitai"]:
|
| 105 |
+
try:
|
| 106 |
+
state_dict_lora_ = self.convert_state_dict(state_dict_lora, lora_prefix=lora_prefix, alpha=1.0)
|
| 107 |
+
converter_fn = model.__class__.state_dict_converter().from_diffusers if model_resource == "diffusers" \
|
| 108 |
+
else model.__class__.state_dict_converter().from_civitai
|
| 109 |
+
state_dict_lora_ = converter_fn(state_dict_lora_)
|
| 110 |
+
if isinstance(state_dict_lora_, tuple):
|
| 111 |
+
state_dict_lora_ = state_dict_lora_[0]
|
| 112 |
+
if len(state_dict_lora_) == 0:
|
| 113 |
+
continue
|
| 114 |
+
for name in state_dict_lora_:
|
| 115 |
+
if name not in state_dict_model:
|
| 116 |
+
break
|
| 117 |
+
else:
|
| 118 |
+
return lora_prefix, model_resource
|
| 119 |
+
except:
|
| 120 |
+
pass
|
| 121 |
+
return None
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
class SDLoRAFromCivitai(LoRAFromCivitai):
|
| 126 |
+
def __init__(self):
|
| 127 |
+
super().__init__()
|
| 128 |
+
self.supported_model_classes = [SDUNet, SDTextEncoder]
|
| 129 |
+
self.lora_prefix = ["lora_unet_", "lora_te_"]
|
| 130 |
+
self.special_keys = {
|
| 131 |
+
"down.blocks": "down_blocks",
|
| 132 |
+
"up.blocks": "up_blocks",
|
| 133 |
+
"mid.block": "mid_block",
|
| 134 |
+
"proj.in": "proj_in",
|
| 135 |
+
"proj.out": "proj_out",
|
| 136 |
+
"transformer.blocks": "transformer_blocks",
|
| 137 |
+
"to.q": "to_q",
|
| 138 |
+
"to.k": "to_k",
|
| 139 |
+
"to.v": "to_v",
|
| 140 |
+
"to.out": "to_out",
|
| 141 |
+
"text.model": "text_model",
|
| 142 |
+
"self.attn.q.proj": "self_attn.q_proj",
|
| 143 |
+
"self.attn.k.proj": "self_attn.k_proj",
|
| 144 |
+
"self.attn.v.proj": "self_attn.v_proj",
|
| 145 |
+
"self.attn.out.proj": "self_attn.out_proj",
|
| 146 |
+
"input.blocks": "model.diffusion_model.input_blocks",
|
| 147 |
+
"middle.block": "model.diffusion_model.middle_block",
|
| 148 |
+
"output.blocks": "model.diffusion_model.output_blocks",
|
| 149 |
+
}
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
class SDXLLoRAFromCivitai(LoRAFromCivitai):
|
| 153 |
+
def __init__(self):
|
| 154 |
+
super().__init__()
|
| 155 |
+
self.supported_model_classes = [SDXLUNet, SDXLTextEncoder, SDXLTextEncoder2]
|
| 156 |
+
self.lora_prefix = ["lora_unet_", "lora_te1_", "lora_te2_"]
|
| 157 |
+
self.renamed_lora_prefix = {"lora_te2_": "2"}
|
| 158 |
+
self.special_keys = {
|
| 159 |
+
"down.blocks": "down_blocks",
|
| 160 |
+
"up.blocks": "up_blocks",
|
| 161 |
+
"mid.block": "mid_block",
|
| 162 |
+
"proj.in": "proj_in",
|
| 163 |
+
"proj.out": "proj_out",
|
| 164 |
+
"transformer.blocks": "transformer_blocks",
|
| 165 |
+
"to.q": "to_q",
|
| 166 |
+
"to.k": "to_k",
|
| 167 |
+
"to.v": "to_v",
|
| 168 |
+
"to.out": "to_out",
|
| 169 |
+
"text.model": "conditioner.embedders.0.transformer.text_model",
|
| 170 |
+
"self.attn.q.proj": "self_attn.q_proj",
|
| 171 |
+
"self.attn.k.proj": "self_attn.k_proj",
|
| 172 |
+
"self.attn.v.proj": "self_attn.v_proj",
|
| 173 |
+
"self.attn.out.proj": "self_attn.out_proj",
|
| 174 |
+
"input.blocks": "model.diffusion_model.input_blocks",
|
| 175 |
+
"middle.block": "model.diffusion_model.middle_block",
|
| 176 |
+
"output.blocks": "model.diffusion_model.output_blocks",
|
| 177 |
+
"2conditioner.embedders.0.transformer.text_model.encoder.layers": "text_model.encoder.layers"
|
| 178 |
+
}
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
class FluxLoRAFromCivitai(LoRAFromCivitai):
|
| 182 |
+
def __init__(self):
|
| 183 |
+
super().__init__()
|
| 184 |
+
self.supported_model_classes = [FluxDiT, FluxDiT]
|
| 185 |
+
self.lora_prefix = ["lora_unet_", "transformer."]
|
| 186 |
+
self.renamed_lora_prefix = {}
|
| 187 |
+
self.special_keys = {
|
| 188 |
+
"single.blocks": "single_blocks",
|
| 189 |
+
"double.blocks": "double_blocks",
|
| 190 |
+
"img.attn": "img_attn",
|
| 191 |
+
"img.mlp": "img_mlp",
|
| 192 |
+
"img.mod": "img_mod",
|
| 193 |
+
"txt.attn": "txt_attn",
|
| 194 |
+
"txt.mlp": "txt_mlp",
|
| 195 |
+
"txt.mod": "txt_mod",
|
| 196 |
+
}
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
class GeneralLoRAFromPeft:
|
| 201 |
+
def __init__(self):
|
| 202 |
+
self.supported_model_classes = [SDUNet, SDXLUNet, SD3DiT, HunyuanDiT, FluxDiT, CogDiT, WanModel]
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
def get_name_dict(self, lora_state_dict):
|
| 206 |
+
lora_name_dict = {}
|
| 207 |
+
for key in lora_state_dict:
|
| 208 |
+
if ".lora_B." not in key:
|
| 209 |
+
continue
|
| 210 |
+
keys = key.split(".")
|
| 211 |
+
if len(keys) > keys.index("lora_B") + 2:
|
| 212 |
+
keys.pop(keys.index("lora_B") + 1)
|
| 213 |
+
keys.pop(keys.index("lora_B"))
|
| 214 |
+
if keys[0] == "diffusion_model":
|
| 215 |
+
keys.pop(0)
|
| 216 |
+
target_name = ".".join(keys)
|
| 217 |
+
lora_name_dict[target_name] = (key, key.replace(".lora_B.", ".lora_A."))
|
| 218 |
+
return lora_name_dict
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def match(self, model: torch.nn.Module, state_dict_lora):
|
| 222 |
+
lora_name_dict = self.get_name_dict(state_dict_lora)
|
| 223 |
+
model_name_dict = {name: None for name, _ in model.named_parameters()}
|
| 224 |
+
matched_num = sum([i in model_name_dict for i in lora_name_dict])
|
| 225 |
+
if matched_num == len(lora_name_dict):
|
| 226 |
+
return "", ""
|
| 227 |
+
else:
|
| 228 |
+
return None
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
def fetch_device_and_dtype(self, state_dict):
|
| 232 |
+
device, dtype = None, None
|
| 233 |
+
for name, param in state_dict.items():
|
| 234 |
+
device, dtype = param.device, param.dtype
|
| 235 |
+
break
|
| 236 |
+
computation_device = device
|
| 237 |
+
computation_dtype = dtype
|
| 238 |
+
if computation_device == torch.device("cpu"):
|
| 239 |
+
if torch.cuda.is_available():
|
| 240 |
+
computation_device = torch.device("cuda")
|
| 241 |
+
if computation_dtype == torch.float8_e4m3fn:
|
| 242 |
+
computation_dtype = torch.float32
|
| 243 |
+
return device, dtype, computation_device, computation_dtype
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def load(self, model, state_dict_lora, lora_prefix="", alpha=1.0, model_resource=""):
|
| 247 |
+
state_dict_model = model.state_dict()
|
| 248 |
+
device, dtype, computation_device, computation_dtype = self.fetch_device_and_dtype(state_dict_model)
|
| 249 |
+
lora_name_dict = self.get_name_dict(state_dict_lora)
|
| 250 |
+
for name in lora_name_dict:
|
| 251 |
+
weight_up = state_dict_lora[lora_name_dict[name][0]].to(device=computation_device, dtype=computation_dtype)
|
| 252 |
+
weight_down = state_dict_lora[lora_name_dict[name][1]].to(device=computation_device, dtype=computation_dtype)
|
| 253 |
+
if len(weight_up.shape) == 4:
|
| 254 |
+
weight_up = weight_up.squeeze(3).squeeze(2)
|
| 255 |
+
weight_down = weight_down.squeeze(3).squeeze(2)
|
| 256 |
+
weight_lora = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)
|
| 257 |
+
else:
|
| 258 |
+
weight_lora = alpha * torch.mm(weight_up, weight_down)
|
| 259 |
+
weight_model = state_dict_model[name].to(device=computation_device, dtype=computation_dtype)
|
| 260 |
+
weight_patched = weight_model + weight_lora
|
| 261 |
+
state_dict_model[name] = weight_patched.to(device=device, dtype=dtype)
|
| 262 |
+
print(f" {len(lora_name_dict)} tensors are updated.")
|
| 263 |
+
model.load_state_dict(state_dict_model)
|
| 264 |
+
|
| 265 |
+
|
| 266 |
+
|
| 267 |
+
class HunyuanVideoLoRAFromCivitai(LoRAFromCivitai):
|
| 268 |
+
def __init__(self):
|
| 269 |
+
super().__init__()
|
| 270 |
+
self.supported_model_classes = [HunyuanVideoDiT, HunyuanVideoDiT]
|
| 271 |
+
self.lora_prefix = ["diffusion_model.", "transformer."]
|
| 272 |
+
self.special_keys = {}
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
class FluxLoRAConverter:
|
| 276 |
+
def __init__(self):
|
| 277 |
+
pass
|
| 278 |
+
|
| 279 |
+
@staticmethod
|
| 280 |
+
def align_to_opensource_format(state_dict, alpha=None):
|
| 281 |
+
prefix_rename_dict = {
|
| 282 |
+
"single_blocks": "lora_unet_single_blocks",
|
| 283 |
+
"blocks": "lora_unet_double_blocks",
|
| 284 |
+
}
|
| 285 |
+
middle_rename_dict = {
|
| 286 |
+
"norm.linear": "modulation_lin",
|
| 287 |
+
"to_qkv_mlp": "linear1",
|
| 288 |
+
"proj_out": "linear2",
|
| 289 |
+
|
| 290 |
+
"norm1_a.linear": "img_mod_lin",
|
| 291 |
+
"norm1_b.linear": "txt_mod_lin",
|
| 292 |
+
"attn.a_to_qkv": "img_attn_qkv",
|
| 293 |
+
"attn.b_to_qkv": "txt_attn_qkv",
|
| 294 |
+
"attn.a_to_out": "img_attn_proj",
|
| 295 |
+
"attn.b_to_out": "txt_attn_proj",
|
| 296 |
+
"ff_a.0": "img_mlp_0",
|
| 297 |
+
"ff_a.2": "img_mlp_2",
|
| 298 |
+
"ff_b.0": "txt_mlp_0",
|
| 299 |
+
"ff_b.2": "txt_mlp_2",
|
| 300 |
+
}
|
| 301 |
+
suffix_rename_dict = {
|
| 302 |
+
"lora_B.weight": "lora_up.weight",
|
| 303 |
+
"lora_A.weight": "lora_down.weight",
|
| 304 |
+
}
|
| 305 |
+
state_dict_ = {}
|
| 306 |
+
for name, param in state_dict.items():
|
| 307 |
+
names = name.split(".")
|
| 308 |
+
if names[-2] != "lora_A" and names[-2] != "lora_B":
|
| 309 |
+
names.pop(-2)
|
| 310 |
+
prefix = names[0]
|
| 311 |
+
middle = ".".join(names[2:-2])
|
| 312 |
+
suffix = ".".join(names[-2:])
|
| 313 |
+
block_id = names[1]
|
| 314 |
+
if middle not in middle_rename_dict:
|
| 315 |
+
continue
|
| 316 |
+
rename = prefix_rename_dict[prefix] + "_" + block_id + "_" + middle_rename_dict[middle] + "." + suffix_rename_dict[suffix]
|
| 317 |
+
state_dict_[rename] = param
|
| 318 |
+
if rename.endswith("lora_up.weight"):
|
| 319 |
+
lora_alpha = alpha if alpha is not None else param.shape[-1]
|
| 320 |
+
state_dict_[rename.replace("lora_up.weight", "alpha")] = torch.tensor((lora_alpha,))[0]
|
| 321 |
+
return state_dict_
|
| 322 |
+
|
| 323 |
+
@staticmethod
|
| 324 |
+
def align_to_diffsynth_format(state_dict):
|
| 325 |
+
rename_dict = {
|
| 326 |
+
"lora_unet_double_blocks_blockid_img_mod_lin.lora_down.weight": "blocks.blockid.norm1_a.linear.lora_A.default.weight",
|
| 327 |
+
"lora_unet_double_blocks_blockid_img_mod_lin.lora_up.weight": "blocks.blockid.norm1_a.linear.lora_B.default.weight",
|
| 328 |
+
"lora_unet_double_blocks_blockid_txt_mod_lin.lora_down.weight": "blocks.blockid.norm1_b.linear.lora_A.default.weight",
|
| 329 |
+
"lora_unet_double_blocks_blockid_txt_mod_lin.lora_up.weight": "blocks.blockid.norm1_b.linear.lora_B.default.weight",
|
| 330 |
+
"lora_unet_double_blocks_blockid_img_attn_qkv.lora_down.weight": "blocks.blockid.attn.a_to_qkv.lora_A.default.weight",
|
| 331 |
+
"lora_unet_double_blocks_blockid_img_attn_qkv.lora_up.weight": "blocks.blockid.attn.a_to_qkv.lora_B.default.weight",
|
| 332 |
+
"lora_unet_double_blocks_blockid_txt_attn_qkv.lora_down.weight": "blocks.blockid.attn.b_to_qkv.lora_A.default.weight",
|
| 333 |
+
"lora_unet_double_blocks_blockid_txt_attn_qkv.lora_up.weight": "blocks.blockid.attn.b_to_qkv.lora_B.default.weight",
|
| 334 |
+
"lora_unet_double_blocks_blockid_img_attn_proj.lora_down.weight": "blocks.blockid.attn.a_to_out.lora_A.default.weight",
|
| 335 |
+
"lora_unet_double_blocks_blockid_img_attn_proj.lora_up.weight": "blocks.blockid.attn.a_to_out.lora_B.default.weight",
|
| 336 |
+
"lora_unet_double_blocks_blockid_txt_attn_proj.lora_down.weight": "blocks.blockid.attn.b_to_out.lora_A.default.weight",
|
| 337 |
+
"lora_unet_double_blocks_blockid_txt_attn_proj.lora_up.weight": "blocks.blockid.attn.b_to_out.lora_B.default.weight",
|
| 338 |
+
"lora_unet_double_blocks_blockid_img_mlp_0.lora_down.weight": "blocks.blockid.ff_a.0.lora_A.default.weight",
|
| 339 |
+
"lora_unet_double_blocks_blockid_img_mlp_0.lora_up.weight": "blocks.blockid.ff_a.0.lora_B.default.weight",
|
| 340 |
+
"lora_unet_double_blocks_blockid_img_mlp_2.lora_down.weight": "blocks.blockid.ff_a.2.lora_A.default.weight",
|
| 341 |
+
"lora_unet_double_blocks_blockid_img_mlp_2.lora_up.weight": "blocks.blockid.ff_a.2.lora_B.default.weight",
|
| 342 |
+
"lora_unet_double_blocks_blockid_txt_mlp_0.lora_down.weight": "blocks.blockid.ff_b.0.lora_A.default.weight",
|
| 343 |
+
"lora_unet_double_blocks_blockid_txt_mlp_0.lora_up.weight": "blocks.blockid.ff_b.0.lora_B.default.weight",
|
| 344 |
+
"lora_unet_double_blocks_blockid_txt_mlp_2.lora_down.weight": "blocks.blockid.ff_b.2.lora_A.default.weight",
|
| 345 |
+
"lora_unet_double_blocks_blockid_txt_mlp_2.lora_up.weight": "blocks.blockid.ff_b.2.lora_B.default.weight",
|
| 346 |
+
"lora_unet_single_blocks_blockid_modulation_lin.lora_down.weight": "single_blocks.blockid.norm.linear.lora_A.default.weight",
|
| 347 |
+
"lora_unet_single_blocks_blockid_modulation_lin.lora_up.weight": "single_blocks.blockid.norm.linear.lora_B.default.weight",
|
| 348 |
+
"lora_unet_single_blocks_blockid_linear1.lora_down.weight": "single_blocks.blockid.to_qkv_mlp.lora_A.default.weight",
|
| 349 |
+
"lora_unet_single_blocks_blockid_linear1.lora_up.weight": "single_blocks.blockid.to_qkv_mlp.lora_B.default.weight",
|
| 350 |
+
"lora_unet_single_blocks_blockid_linear2.lora_down.weight": "single_blocks.blockid.proj_out.lora_A.default.weight",
|
| 351 |
+
"lora_unet_single_blocks_blockid_linear2.lora_up.weight": "single_blocks.blockid.proj_out.lora_B.default.weight",
|
| 352 |
+
}
|
| 353 |
+
def guess_block_id(name):
|
| 354 |
+
names = name.split("_")
|
| 355 |
+
for i in names:
|
| 356 |
+
if i.isdigit():
|
| 357 |
+
return i, name.replace(f"_{i}_", "_blockid_")
|
| 358 |
+
return None, None
|
| 359 |
+
state_dict_ = {}
|
| 360 |
+
for name, param in state_dict.items():
|
| 361 |
+
block_id, source_name = guess_block_id(name)
|
| 362 |
+
if source_name in rename_dict:
|
| 363 |
+
target_name = rename_dict[source_name]
|
| 364 |
+
target_name = target_name.replace(".blockid.", f".{block_id}.")
|
| 365 |
+
state_dict_[target_name] = param
|
| 366 |
+
else:
|
| 367 |
+
state_dict_[name] = param
|
| 368 |
+
return state_dict_
|
| 369 |
+
|
| 370 |
+
|
| 371 |
+
class WanLoRAConverter:
|
| 372 |
+
def __init__(self):
|
| 373 |
+
pass
|
| 374 |
+
|
| 375 |
+
@staticmethod
|
| 376 |
+
def align_to_opensource_format(state_dict, **kwargs):
|
| 377 |
+
state_dict = {"diffusion_model." + name.replace(".default.", "."): param for name, param in state_dict.items()}
|
| 378 |
+
return state_dict
|
| 379 |
+
|
| 380 |
+
@staticmethod
|
| 381 |
+
def align_to_diffsynth_format(state_dict, **kwargs):
|
| 382 |
+
state_dict = {name.replace("diffusion_model.", "").replace(".lora_A.weight", ".lora_A.default.weight").replace(".lora_B.weight", ".lora_B.default.weight"): param for name, param in state_dict.items()}
|
| 383 |
+
return state_dict
|
| 384 |
+
|
| 385 |
+
|
| 386 |
+
def get_lora_loaders():
|
| 387 |
+
return [SDLoRAFromCivitai(), SDXLLoRAFromCivitai(), FluxLoRAFromCivitai(), HunyuanVideoLoRAFromCivitai(), GeneralLoRAFromPeft()]
|
diffsynth/models/model_manager.py
ADDED
|
@@ -0,0 +1,467 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os, torch, json, importlib
|
| 2 |
+
from typing import List
|
| 3 |
+
|
| 4 |
+
from .downloader import download_models, download_customized_models, Preset_model_id, Preset_model_website
|
| 5 |
+
|
| 6 |
+
from .sd_text_encoder import SDTextEncoder
|
| 7 |
+
from .sd_unet import SDUNet
|
| 8 |
+
from .sd_vae_encoder import SDVAEEncoder
|
| 9 |
+
from .sd_vae_decoder import SDVAEDecoder
|
| 10 |
+
from .lora import get_lora_loaders
|
| 11 |
+
|
| 12 |
+
from .sdxl_text_encoder import SDXLTextEncoder, SDXLTextEncoder2
|
| 13 |
+
from .sdxl_unet import SDXLUNet
|
| 14 |
+
from .sdxl_vae_decoder import SDXLVAEDecoder
|
| 15 |
+
from .sdxl_vae_encoder import SDXLVAEEncoder
|
| 16 |
+
|
| 17 |
+
from .sd3_text_encoder import SD3TextEncoder1, SD3TextEncoder2, SD3TextEncoder3
|
| 18 |
+
from .sd3_dit import SD3DiT
|
| 19 |
+
from .sd3_vae_decoder import SD3VAEDecoder
|
| 20 |
+
from .sd3_vae_encoder import SD3VAEEncoder
|
| 21 |
+
|
| 22 |
+
from .sd_controlnet import SDControlNet
|
| 23 |
+
from .sdxl_controlnet import SDXLControlNetUnion
|
| 24 |
+
|
| 25 |
+
from .sd_motion import SDMotionModel
|
| 26 |
+
from .sdxl_motion import SDXLMotionModel
|
| 27 |
+
|
| 28 |
+
from .svd_image_encoder import SVDImageEncoder
|
| 29 |
+
from .svd_unet import SVDUNet
|
| 30 |
+
from .svd_vae_decoder import SVDVAEDecoder
|
| 31 |
+
from .svd_vae_encoder import SVDVAEEncoder
|
| 32 |
+
|
| 33 |
+
from .sd_ipadapter import SDIpAdapter, IpAdapterCLIPImageEmbedder
|
| 34 |
+
from .sdxl_ipadapter import SDXLIpAdapter, IpAdapterXLCLIPImageEmbedder
|
| 35 |
+
|
| 36 |
+
from .hunyuan_dit_text_encoder import HunyuanDiTCLIPTextEncoder, HunyuanDiTT5TextEncoder
|
| 37 |
+
from .hunyuan_dit import HunyuanDiT
|
| 38 |
+
from .hunyuan_video_vae_decoder import HunyuanVideoVAEDecoder
|
| 39 |
+
from .hunyuan_video_vae_encoder import HunyuanVideoVAEEncoder
|
| 40 |
+
|
| 41 |
+
from .flux_dit import FluxDiT
|
| 42 |
+
from .flux_text_encoder import FluxTextEncoder2
|
| 43 |
+
from .flux_vae import FluxVAEEncoder, FluxVAEDecoder
|
| 44 |
+
from .flux_ipadapter import FluxIpAdapter
|
| 45 |
+
|
| 46 |
+
from .cog_vae import CogVAEEncoder, CogVAEDecoder
|
| 47 |
+
from .cog_dit import CogDiT
|
| 48 |
+
|
| 49 |
+
from ..extensions.RIFE import IFNet
|
| 50 |
+
from ..extensions.ESRGAN import RRDBNet
|
| 51 |
+
|
| 52 |
+
from ..configs.model_config import model_loader_configs, huggingface_model_loader_configs, patch_model_loader_configs
|
| 53 |
+
from .utils import load_state_dict, init_weights_on_device, hash_state_dict_keys, split_state_dict_with_prefix
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device):
|
| 57 |
+
loaded_model_names, loaded_models = [], []
|
| 58 |
+
for model_name, model_class in zip(model_names, model_classes):
|
| 59 |
+
print(f" model_name: {model_name} model_class: {model_class.__name__}")
|
| 60 |
+
state_dict_converter = model_class.state_dict_converter()
|
| 61 |
+
if model_resource == "civitai":
|
| 62 |
+
state_dict_results = state_dict_converter.from_civitai(state_dict)
|
| 63 |
+
elif model_resource == "diffusers":
|
| 64 |
+
state_dict_results = state_dict_converter.from_diffusers(state_dict)
|
| 65 |
+
if isinstance(state_dict_results, tuple):
|
| 66 |
+
model_state_dict, extra_kwargs = state_dict_results
|
| 67 |
+
print(f" This model is initialized with extra kwargs: {extra_kwargs}")
|
| 68 |
+
else:
|
| 69 |
+
model_state_dict, extra_kwargs = state_dict_results, {}
|
| 70 |
+
torch_dtype = torch.float32 if extra_kwargs.get("upcast_to_float32", False) else torch_dtype
|
| 71 |
+
with init_weights_on_device():
|
| 72 |
+
model = model_class(**extra_kwargs)
|
| 73 |
+
if hasattr(model, "eval"):
|
| 74 |
+
model = model.eval()
|
| 75 |
+
model.load_state_dict(model_state_dict, assign=True)
|
| 76 |
+
model = model.to(dtype=torch_dtype, device=device)
|
| 77 |
+
loaded_model_names.append(model_name)
|
| 78 |
+
loaded_models.append(model)
|
| 79 |
+
return loaded_model_names, loaded_models
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def load_model_from_huggingface_folder(file_path, model_names, model_classes, torch_dtype, device):
|
| 83 |
+
loaded_model_names, loaded_models = [], []
|
| 84 |
+
for model_name, model_class in zip(model_names, model_classes):
|
| 85 |
+
if torch_dtype in [torch.float32, torch.float16, torch.bfloat16]:
|
| 86 |
+
model = model_class.from_pretrained(file_path, torch_dtype=torch_dtype).eval()
|
| 87 |
+
else:
|
| 88 |
+
model = model_class.from_pretrained(file_path).eval().to(dtype=torch_dtype)
|
| 89 |
+
if torch_dtype == torch.float16 and hasattr(model, "half"):
|
| 90 |
+
model = model.half()
|
| 91 |
+
try:
|
| 92 |
+
model = model.to(device=device)
|
| 93 |
+
except:
|
| 94 |
+
pass
|
| 95 |
+
loaded_model_names.append(model_name)
|
| 96 |
+
loaded_models.append(model)
|
| 97 |
+
return loaded_model_names, loaded_models
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def load_single_patch_model_from_single_file(state_dict, model_name, model_class, base_model, extra_kwargs, torch_dtype, device):
|
| 101 |
+
print(f" model_name: {model_name} model_class: {model_class.__name__} extra_kwargs: {extra_kwargs}")
|
| 102 |
+
base_state_dict = base_model.state_dict()
|
| 103 |
+
base_model.to("cpu")
|
| 104 |
+
del base_model
|
| 105 |
+
model = model_class(**extra_kwargs)
|
| 106 |
+
model.load_state_dict(base_state_dict, strict=False)
|
| 107 |
+
model.load_state_dict(state_dict, strict=False)
|
| 108 |
+
model.to(dtype=torch_dtype, device=device)
|
| 109 |
+
return model
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def load_patch_model_from_single_file(state_dict, model_names, model_classes, extra_kwargs, model_manager, torch_dtype, device):
|
| 113 |
+
loaded_model_names, loaded_models = [], []
|
| 114 |
+
for model_name, model_class in zip(model_names, model_classes):
|
| 115 |
+
while True:
|
| 116 |
+
for model_id in range(len(model_manager.model)):
|
| 117 |
+
base_model_name = model_manager.model_name[model_id]
|
| 118 |
+
if base_model_name == model_name:
|
| 119 |
+
base_model_path = model_manager.model_path[model_id]
|
| 120 |
+
base_model = model_manager.model[model_id]
|
| 121 |
+
print(f" Adding patch model to {base_model_name} ({base_model_path})")
|
| 122 |
+
patched_model = load_single_patch_model_from_single_file(
|
| 123 |
+
state_dict, model_name, model_class, base_model, extra_kwargs, torch_dtype, device)
|
| 124 |
+
loaded_model_names.append(base_model_name)
|
| 125 |
+
loaded_models.append(patched_model)
|
| 126 |
+
model_manager.model.pop(model_id)
|
| 127 |
+
model_manager.model_path.pop(model_id)
|
| 128 |
+
model_manager.model_name.pop(model_id)
|
| 129 |
+
break
|
| 130 |
+
else:
|
| 131 |
+
break
|
| 132 |
+
return loaded_model_names, loaded_models
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
class ModelDetectorTemplate:
|
| 137 |
+
def __init__(self):
|
| 138 |
+
pass
|
| 139 |
+
|
| 140 |
+
def match(self, file_path="", state_dict={}):
|
| 141 |
+
return False
|
| 142 |
+
|
| 143 |
+
def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):
|
| 144 |
+
return [], []
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
class ModelDetectorFromSingleFile:
|
| 149 |
+
def __init__(self, model_loader_configs=[]):
|
| 150 |
+
self.keys_hash_with_shape_dict = {}
|
| 151 |
+
self.keys_hash_dict = {}
|
| 152 |
+
for metadata in model_loader_configs:
|
| 153 |
+
self.add_model_metadata(*metadata)
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
def add_model_metadata(self, keys_hash, keys_hash_with_shape, model_names, model_classes, model_resource):
|
| 157 |
+
self.keys_hash_with_shape_dict[keys_hash_with_shape] = (model_names, model_classes, model_resource)
|
| 158 |
+
if keys_hash is not None:
|
| 159 |
+
self.keys_hash_dict[keys_hash] = (model_names, model_classes, model_resource)
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def match(self, file_path="", state_dict={}):
|
| 163 |
+
if isinstance(file_path, str) and os.path.isdir(file_path):
|
| 164 |
+
return False
|
| 165 |
+
if len(state_dict) == 0:
|
| 166 |
+
state_dict = load_state_dict(file_path)
|
| 167 |
+
keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)
|
| 168 |
+
if keys_hash_with_shape in self.keys_hash_with_shape_dict:
|
| 169 |
+
return True
|
| 170 |
+
keys_hash = hash_state_dict_keys(state_dict, with_shape=False)
|
| 171 |
+
if keys_hash in self.keys_hash_dict:
|
| 172 |
+
return True
|
| 173 |
+
return False
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):
|
| 177 |
+
if len(state_dict) == 0:
|
| 178 |
+
state_dict = load_state_dict(file_path)
|
| 179 |
+
|
| 180 |
+
# Load models with strict matching
|
| 181 |
+
keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)
|
| 182 |
+
if keys_hash_with_shape in self.keys_hash_with_shape_dict:
|
| 183 |
+
model_names, model_classes, model_resource = self.keys_hash_with_shape_dict[keys_hash_with_shape]
|
| 184 |
+
loaded_model_names, loaded_models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device)
|
| 185 |
+
return loaded_model_names, loaded_models
|
| 186 |
+
|
| 187 |
+
# Load models without strict matching
|
| 188 |
+
# (the shape of parameters may be inconsistent, and the state_dict_converter will modify the model architecture)
|
| 189 |
+
keys_hash = hash_state_dict_keys(state_dict, with_shape=False)
|
| 190 |
+
if keys_hash in self.keys_hash_dict:
|
| 191 |
+
model_names, model_classes, model_resource = self.keys_hash_dict[keys_hash]
|
| 192 |
+
loaded_model_names, loaded_models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, torch_dtype, device)
|
| 193 |
+
return loaded_model_names, loaded_models
|
| 194 |
+
|
| 195 |
+
return loaded_model_names, loaded_models
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
class ModelDetectorFromSplitedSingleFile(ModelDetectorFromSingleFile):
|
| 200 |
+
def __init__(self, model_loader_configs=[]):
|
| 201 |
+
super().__init__(model_loader_configs)
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def match(self, file_path="", state_dict={}):
|
| 205 |
+
if isinstance(file_path, str) and os.path.isdir(file_path):
|
| 206 |
+
return False
|
| 207 |
+
if len(state_dict) == 0:
|
| 208 |
+
state_dict = load_state_dict(file_path)
|
| 209 |
+
splited_state_dict = split_state_dict_with_prefix(state_dict)
|
| 210 |
+
for sub_state_dict in splited_state_dict:
|
| 211 |
+
if super().match(file_path, sub_state_dict):
|
| 212 |
+
return True
|
| 213 |
+
return False
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):
|
| 217 |
+
# Split the state_dict and load from each component
|
| 218 |
+
splited_state_dict = split_state_dict_with_prefix(state_dict)
|
| 219 |
+
valid_state_dict = {}
|
| 220 |
+
for sub_state_dict in splited_state_dict:
|
| 221 |
+
if super().match(file_path, sub_state_dict):
|
| 222 |
+
valid_state_dict.update(sub_state_dict)
|
| 223 |
+
if super().match(file_path, valid_state_dict):
|
| 224 |
+
loaded_model_names, loaded_models = super().load(file_path, valid_state_dict, device, torch_dtype)
|
| 225 |
+
else:
|
| 226 |
+
loaded_model_names, loaded_models = [], []
|
| 227 |
+
for sub_state_dict in splited_state_dict:
|
| 228 |
+
if super().match(file_path, sub_state_dict):
|
| 229 |
+
loaded_model_names_, loaded_models_ = super().load(file_path, valid_state_dict, device, torch_dtype)
|
| 230 |
+
loaded_model_names += loaded_model_names_
|
| 231 |
+
loaded_models += loaded_models_
|
| 232 |
+
return loaded_model_names, loaded_models
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
class ModelDetectorFromHuggingfaceFolder:
|
| 237 |
+
def __init__(self, model_loader_configs=[]):
|
| 238 |
+
self.architecture_dict = {}
|
| 239 |
+
for metadata in model_loader_configs:
|
| 240 |
+
self.add_model_metadata(*metadata)
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
def add_model_metadata(self, architecture, huggingface_lib, model_name, redirected_architecture):
|
| 244 |
+
self.architecture_dict[architecture] = (huggingface_lib, model_name, redirected_architecture)
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
def match(self, file_path="", state_dict={}):
|
| 248 |
+
if not isinstance(file_path, str) or os.path.isfile(file_path):
|
| 249 |
+
return False
|
| 250 |
+
file_list = os.listdir(file_path)
|
| 251 |
+
if "config.json" not in file_list:
|
| 252 |
+
return False
|
| 253 |
+
with open(os.path.join(file_path, "config.json"), "r") as f:
|
| 254 |
+
config = json.load(f)
|
| 255 |
+
if "architectures" not in config and "_class_name" not in config:
|
| 256 |
+
return False
|
| 257 |
+
return True
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, **kwargs):
|
| 261 |
+
with open(os.path.join(file_path, "config.json"), "r") as f:
|
| 262 |
+
config = json.load(f)
|
| 263 |
+
loaded_model_names, loaded_models = [], []
|
| 264 |
+
architectures = config["architectures"] if "architectures" in config else [config["_class_name"]]
|
| 265 |
+
for architecture in architectures:
|
| 266 |
+
huggingface_lib, model_name, redirected_architecture = self.architecture_dict[architecture]
|
| 267 |
+
if redirected_architecture is not None:
|
| 268 |
+
architecture = redirected_architecture
|
| 269 |
+
model_class = importlib.import_module(huggingface_lib).__getattribute__(architecture)
|
| 270 |
+
loaded_model_names_, loaded_models_ = load_model_from_huggingface_folder(file_path, [model_name], [model_class], torch_dtype, device)
|
| 271 |
+
loaded_model_names += loaded_model_names_
|
| 272 |
+
loaded_models += loaded_models_
|
| 273 |
+
return loaded_model_names, loaded_models
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
class ModelDetectorFromPatchedSingleFile:
|
| 278 |
+
def __init__(self, model_loader_configs=[]):
|
| 279 |
+
self.keys_hash_with_shape_dict = {}
|
| 280 |
+
for metadata in model_loader_configs:
|
| 281 |
+
self.add_model_metadata(*metadata)
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def add_model_metadata(self, keys_hash_with_shape, model_name, model_class, extra_kwargs):
|
| 285 |
+
self.keys_hash_with_shape_dict[keys_hash_with_shape] = (model_name, model_class, extra_kwargs)
|
| 286 |
+
|
| 287 |
+
|
| 288 |
+
def match(self, file_path="", state_dict={}):
|
| 289 |
+
if not isinstance(file_path, str) or os.path.isdir(file_path):
|
| 290 |
+
return False
|
| 291 |
+
if len(state_dict) == 0:
|
| 292 |
+
state_dict = load_state_dict(file_path)
|
| 293 |
+
keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)
|
| 294 |
+
if keys_hash_with_shape in self.keys_hash_with_shape_dict:
|
| 295 |
+
return True
|
| 296 |
+
return False
|
| 297 |
+
|
| 298 |
+
|
| 299 |
+
def load(self, file_path="", state_dict={}, device="cuda", torch_dtype=torch.float16, model_manager=None, **kwargs):
|
| 300 |
+
if len(state_dict) == 0:
|
| 301 |
+
state_dict = load_state_dict(file_path)
|
| 302 |
+
|
| 303 |
+
# Load models with strict matching
|
| 304 |
+
loaded_model_names, loaded_models = [], []
|
| 305 |
+
keys_hash_with_shape = hash_state_dict_keys(state_dict, with_shape=True)
|
| 306 |
+
if keys_hash_with_shape in self.keys_hash_with_shape_dict:
|
| 307 |
+
model_names, model_classes, extra_kwargs = self.keys_hash_with_shape_dict[keys_hash_with_shape]
|
| 308 |
+
loaded_model_names_, loaded_models_ = load_patch_model_from_single_file(
|
| 309 |
+
state_dict, model_names, model_classes, extra_kwargs, model_manager, torch_dtype, device)
|
| 310 |
+
loaded_model_names += loaded_model_names_
|
| 311 |
+
loaded_models += loaded_models_
|
| 312 |
+
return loaded_model_names, loaded_models
|
| 313 |
+
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
class ModelManager:
|
| 317 |
+
def __init__(
|
| 318 |
+
self,
|
| 319 |
+
torch_dtype=torch.float16,
|
| 320 |
+
device="cuda",
|
| 321 |
+
model_id_list: List[Preset_model_id] = [],
|
| 322 |
+
downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],
|
| 323 |
+
file_path_list: List[str] = [],
|
| 324 |
+
):
|
| 325 |
+
self.torch_dtype = torch_dtype
|
| 326 |
+
self.device = device
|
| 327 |
+
self.model = []
|
| 328 |
+
self.model_path = []
|
| 329 |
+
self.model_name = []
|
| 330 |
+
downloaded_files = download_models(model_id_list, downloading_priority) if len(model_id_list) > 0 else []
|
| 331 |
+
self.model_detector = [
|
| 332 |
+
ModelDetectorFromSingleFile(model_loader_configs),
|
| 333 |
+
ModelDetectorFromSplitedSingleFile(model_loader_configs),
|
| 334 |
+
ModelDetectorFromHuggingfaceFolder(huggingface_model_loader_configs),
|
| 335 |
+
ModelDetectorFromPatchedSingleFile(patch_model_loader_configs),
|
| 336 |
+
]
|
| 337 |
+
self.load_models(downloaded_files + file_path_list)
|
| 338 |
+
|
| 339 |
+
|
| 340 |
+
def load_model_from_single_file(self, file_path="", state_dict={}, model_names=[], model_classes=[], model_resource=None):
|
| 341 |
+
print(f"Loading models from file: {file_path}")
|
| 342 |
+
if len(state_dict) == 0:
|
| 343 |
+
state_dict = load_state_dict(file_path)
|
| 344 |
+
model_names, models = load_model_from_single_file(state_dict, model_names, model_classes, model_resource, self.torch_dtype, self.device)
|
| 345 |
+
for model_name, model in zip(model_names, models):
|
| 346 |
+
self.model.append(model)
|
| 347 |
+
self.model_path.append(file_path)
|
| 348 |
+
self.model_name.append(model_name)
|
| 349 |
+
print(f" The following models are loaded: {model_names}.")
|
| 350 |
+
|
| 351 |
+
|
| 352 |
+
def load_model_from_huggingface_folder(self, file_path="", model_names=[], model_classes=[]):
|
| 353 |
+
print(f"Loading models from folder: {file_path}")
|
| 354 |
+
model_names, models = load_model_from_huggingface_folder(file_path, model_names, model_classes, self.torch_dtype, self.device)
|
| 355 |
+
for model_name, model in zip(model_names, models):
|
| 356 |
+
self.model.append(model)
|
| 357 |
+
self.model_path.append(file_path)
|
| 358 |
+
self.model_name.append(model_name)
|
| 359 |
+
print(f" The following models are loaded: {model_names}.")
|
| 360 |
+
|
| 361 |
+
|
| 362 |
+
def load_patch_model_from_single_file(self, file_path="", state_dict={}, model_names=[], model_classes=[], extra_kwargs={}):
|
| 363 |
+
print(f"Loading patch models from file: {file_path}")
|
| 364 |
+
model_names, models = load_patch_model_from_single_file(
|
| 365 |
+
state_dict, model_names, model_classes, extra_kwargs, self, self.torch_dtype, self.device)
|
| 366 |
+
for model_name, model in zip(model_names, models):
|
| 367 |
+
self.model.append(model)
|
| 368 |
+
self.model_path.append(file_path)
|
| 369 |
+
self.model_name.append(model_name)
|
| 370 |
+
print(f" The following patched models are loaded: {model_names}.")
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
def load_lora(self, file_path="", state_dict={}, lora_alpha=1.0):
|
| 374 |
+
if isinstance(file_path, list):
|
| 375 |
+
for file_path_ in file_path:
|
| 376 |
+
self.load_lora(file_path_, state_dict=state_dict, lora_alpha=lora_alpha)
|
| 377 |
+
else:
|
| 378 |
+
print(f"Loading LoRA models from file: {file_path}")
|
| 379 |
+
is_loaded = False
|
| 380 |
+
if len(state_dict) == 0:
|
| 381 |
+
state_dict = load_state_dict(file_path)
|
| 382 |
+
for model_name, model, model_path in zip(self.model_name, self.model, self.model_path):
|
| 383 |
+
for lora in get_lora_loaders():
|
| 384 |
+
match_results = lora.match(model, state_dict)
|
| 385 |
+
if match_results is not None:
|
| 386 |
+
print(f" Adding LoRA to {model_name} ({model_path}).")
|
| 387 |
+
lora_prefix, model_resource = match_results
|
| 388 |
+
lora.load(model, state_dict, lora_prefix, alpha=lora_alpha, model_resource=model_resource)
|
| 389 |
+
is_loaded = True
|
| 390 |
+
break
|
| 391 |
+
if not is_loaded:
|
| 392 |
+
print(f" Cannot load LoRA: {file_path}")
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
def load_model(self, file_path, model_names=None, device=None, torch_dtype=None):
|
| 396 |
+
print(f"Loading models from: {file_path}")
|
| 397 |
+
if device is None: device = self.device
|
| 398 |
+
if torch_dtype is None: torch_dtype = self.torch_dtype
|
| 399 |
+
if isinstance(file_path, list):
|
| 400 |
+
state_dict = {}
|
| 401 |
+
for path in file_path:
|
| 402 |
+
state_dict.update(load_state_dict(path))
|
| 403 |
+
elif os.path.isfile(file_path):
|
| 404 |
+
state_dict = load_state_dict(file_path)
|
| 405 |
+
else:
|
| 406 |
+
state_dict = None
|
| 407 |
+
for model_detector in self.model_detector:
|
| 408 |
+
if model_detector.match(file_path, state_dict):
|
| 409 |
+
model_names, models = model_detector.load(
|
| 410 |
+
file_path, state_dict,
|
| 411 |
+
device=device, torch_dtype=torch_dtype,
|
| 412 |
+
allowed_model_names=model_names, model_manager=self
|
| 413 |
+
)
|
| 414 |
+
for model_name, model in zip(model_names, models):
|
| 415 |
+
self.model.append(model)
|
| 416 |
+
self.model_path.append(file_path)
|
| 417 |
+
self.model_name.append(model_name)
|
| 418 |
+
print(f" The following models are loaded: {model_names}.")
|
| 419 |
+
break
|
| 420 |
+
else:
|
| 421 |
+
print(f" We cannot detect the model type. No models are loaded.")
|
| 422 |
+
|
| 423 |
+
|
| 424 |
+
def load_models(self, file_path_list, model_names=None, device=None, torch_dtype=None):
|
| 425 |
+
for file_path in file_path_list:
|
| 426 |
+
self.load_model(file_path, model_names, device=device, torch_dtype=torch_dtype)
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
def fetch_model(self, model_name, file_path=None, require_model_path=False, index=None):
|
| 430 |
+
fetched_models = []
|
| 431 |
+
fetched_model_paths = []
|
| 432 |
+
for model, model_path, model_name_ in zip(self.model, self.model_path, self.model_name):
|
| 433 |
+
if file_path is not None and file_path != model_path:
|
| 434 |
+
continue
|
| 435 |
+
if model_name == model_name_:
|
| 436 |
+
fetched_models.append(model)
|
| 437 |
+
fetched_model_paths.append(model_path)
|
| 438 |
+
if len(fetched_models) == 0:
|
| 439 |
+
print(f"No {model_name} models available.")
|
| 440 |
+
return None
|
| 441 |
+
if len(fetched_models) == 1:
|
| 442 |
+
print(f"Using {model_name} from {fetched_model_paths[0]}.")
|
| 443 |
+
model = fetched_models[0]
|
| 444 |
+
path = fetched_model_paths[0]
|
| 445 |
+
else:
|
| 446 |
+
if index is None:
|
| 447 |
+
model = fetched_models[0]
|
| 448 |
+
path = fetched_model_paths[0]
|
| 449 |
+
print(f"More than one {model_name} models are loaded in model manager: {fetched_model_paths}. Using {model_name} from {fetched_model_paths[0]}.")
|
| 450 |
+
elif isinstance(index, int):
|
| 451 |
+
model = fetched_models[:index]
|
| 452 |
+
path = fetched_model_paths[:index]
|
| 453 |
+
print(f"More than one {model_name} models are loaded in model manager: {fetched_model_paths}. Using {model_name} from {fetched_model_paths[:index]}.")
|
| 454 |
+
else:
|
| 455 |
+
model = fetched_models
|
| 456 |
+
path = fetched_model_paths
|
| 457 |
+
print(f"More than one {model_name} models are loaded in model manager: {fetched_model_paths}. Using {model_name} from {fetched_model_paths}.")
|
| 458 |
+
if require_model_path:
|
| 459 |
+
return model, path
|
| 460 |
+
else:
|
| 461 |
+
return model
|
| 462 |
+
|
| 463 |
+
|
| 464 |
+
def to(self, device):
|
| 465 |
+
for model in self.model:
|
| 466 |
+
model.to(device)
|
| 467 |
+
|
diffsynth/models/nexus_gen.py
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from PIL import Image
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class NexusGenAutoregressiveModel(torch.nn.Module):
|
| 6 |
+
def __init__(self, max_length=1024, max_pixels=262640):
|
| 7 |
+
super(NexusGenAutoregressiveModel, self).__init__()
|
| 8 |
+
from .nexus_gen_ar_model import Qwen2_5_VLForConditionalGeneration
|
| 9 |
+
from transformers import Qwen2_5_VLConfig
|
| 10 |
+
self.max_length = max_length
|
| 11 |
+
self.max_pixels = max_pixels
|
| 12 |
+
model_config = Qwen2_5_VLConfig(**{
|
| 13 |
+
"_name_or_path": "DiffSynth-Studio/Nexus-GenV2",
|
| 14 |
+
"architectures": [
|
| 15 |
+
"Qwen2_5_VLForConditionalGeneration"
|
| 16 |
+
],
|
| 17 |
+
"attention_dropout": 0.0,
|
| 18 |
+
"auto_map": {
|
| 19 |
+
"AutoConfig": "configuration_qwen2_5_vl.Qwen2_5_VLConfig",
|
| 20 |
+
"AutoModel": "modeling_qwen2_5_vl.Qwen2_5_VLModel",
|
| 21 |
+
"AutoModelForCausalLM": "modeling_qwen2_5_vl.Qwen2_5_VLForConditionalGeneration"
|
| 22 |
+
},
|
| 23 |
+
"bos_token_id": 151643,
|
| 24 |
+
"eos_token_id": 151645,
|
| 25 |
+
"hidden_act": "silu",
|
| 26 |
+
"hidden_size": 3584,
|
| 27 |
+
"image_token_id": 151655,
|
| 28 |
+
"initializer_range": 0.02,
|
| 29 |
+
"intermediate_size": 18944,
|
| 30 |
+
"max_position_embeddings": 128000,
|
| 31 |
+
"max_window_layers": 28,
|
| 32 |
+
"model_type": "qwen2_5_vl",
|
| 33 |
+
"num_attention_heads": 28,
|
| 34 |
+
"num_hidden_layers": 28,
|
| 35 |
+
"num_key_value_heads": 4,
|
| 36 |
+
"pad_token_id": 151643,
|
| 37 |
+
"rms_norm_eps": 1e-06,
|
| 38 |
+
"rope_scaling": {
|
| 39 |
+
"mrope_section": [
|
| 40 |
+
16,
|
| 41 |
+
24,
|
| 42 |
+
24
|
| 43 |
+
],
|
| 44 |
+
"rope_type": "default",
|
| 45 |
+
"type": "default"
|
| 46 |
+
},
|
| 47 |
+
"rope_theta": 1000000.0,
|
| 48 |
+
"sliding_window": 32768,
|
| 49 |
+
"tie_word_embeddings": False,
|
| 50 |
+
"torch_dtype": "bfloat16",
|
| 51 |
+
"transformers_version": "4.49.0",
|
| 52 |
+
"use_cache": False,
|
| 53 |
+
"use_sliding_window": False,
|
| 54 |
+
"video_token_id": 151656,
|
| 55 |
+
"vision_config": {
|
| 56 |
+
"hidden_size": 1280,
|
| 57 |
+
"in_chans": 3,
|
| 58 |
+
"model_type": "qwen2_5_vl",
|
| 59 |
+
"spatial_patch_size": 14,
|
| 60 |
+
"tokens_per_second": 2,
|
| 61 |
+
"torch_dtype": "bfloat16"
|
| 62 |
+
},
|
| 63 |
+
"vision_end_token_id": 151653,
|
| 64 |
+
"vision_start_token_id": 151652,
|
| 65 |
+
"vision_token_id": 151654,
|
| 66 |
+
"vocab_size": 152064
|
| 67 |
+
})
|
| 68 |
+
self.model = Qwen2_5_VLForConditionalGeneration(model_config)
|
| 69 |
+
self.processor = None
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
def load_processor(self, path):
|
| 73 |
+
from .nexus_gen_ar_model import Qwen2_5_VLProcessor
|
| 74 |
+
self.processor = Qwen2_5_VLProcessor.from_pretrained(path)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
@staticmethod
|
| 78 |
+
def state_dict_converter():
|
| 79 |
+
return NexusGenAutoregressiveModelStateDictConverter()
|
| 80 |
+
|
| 81 |
+
def bound_image(self, image, max_pixels=262640):
|
| 82 |
+
from qwen_vl_utils import smart_resize
|
| 83 |
+
resized_height, resized_width = smart_resize(
|
| 84 |
+
image.height,
|
| 85 |
+
image.width,
|
| 86 |
+
max_pixels=max_pixels,
|
| 87 |
+
)
|
| 88 |
+
return image.resize((resized_width, resized_height))
|
| 89 |
+
|
| 90 |
+
def get_editing_msg(self, instruction):
|
| 91 |
+
if '<image>' not in instruction:
|
| 92 |
+
instruction = '<image> ' + instruction
|
| 93 |
+
messages = [{"role":"user", "content":instruction}, {"role":"assistant", "content":"Here is the image: <image>"}]
|
| 94 |
+
return messages
|
| 95 |
+
|
| 96 |
+
def get_generation_msg(self, instruction):
|
| 97 |
+
instruction = "Generate an image according to the following description: {}".format(instruction)
|
| 98 |
+
messages = [{"role":"user", "content":instruction}, {"role":"assistant", "content":"Here is an image based on the description: <image>"}]
|
| 99 |
+
return messages
|
| 100 |
+
|
| 101 |
+
def forward(self, instruction, ref_image=None, num_img_tokens=81):
|
| 102 |
+
"""
|
| 103 |
+
Generate target embeddings for the given instruction and reference image.
|
| 104 |
+
"""
|
| 105 |
+
if ref_image is not None:
|
| 106 |
+
messages = self.get_editing_msg(instruction)
|
| 107 |
+
images = [self.bound_image(ref_image)] + [Image.new(mode='RGB', size=(252, 252), color=(255, 255, 255))]
|
| 108 |
+
output_image_embeddings = self.get_target_embeddings(images, messages, self.processor, self.model, num_img_tokens)
|
| 109 |
+
else:
|
| 110 |
+
messages = self.get_generation_msg(instruction)
|
| 111 |
+
images = [Image.new(mode='RGB', size=(252, 252), color=(255, 255, 255))]
|
| 112 |
+
output_image_embeddings = self.get_target_embeddings(images, messages, self.processor, self.model, num_img_tokens)
|
| 113 |
+
|
| 114 |
+
return output_image_embeddings
|
| 115 |
+
|
| 116 |
+
def get_target_embeddings(self, images, messages, processor, model, num_img_tokens=81):
|
| 117 |
+
text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=False)
|
| 118 |
+
text = text.replace('<image>', '<|vision_start|><|image_pad|><|vision_end|>')
|
| 119 |
+
inputs = processor(
|
| 120 |
+
text=[text],
|
| 121 |
+
images=images,
|
| 122 |
+
padding=True,
|
| 123 |
+
return_tensors="pt",
|
| 124 |
+
)
|
| 125 |
+
inputs = inputs.to(model.device)
|
| 126 |
+
|
| 127 |
+
input_embeds = model.model.embed_tokens(inputs['input_ids'])
|
| 128 |
+
image_embeds = model.visual(inputs['pixel_values'], grid_thw=inputs['image_grid_thw'])
|
| 129 |
+
ground_truth_image_embeds = image_embeds[-num_img_tokens:]
|
| 130 |
+
input_image_embeds = image_embeds[:-num_img_tokens]
|
| 131 |
+
|
| 132 |
+
image_mask = inputs['input_ids'] == model.config.image_token_id
|
| 133 |
+
indices = image_mask.cumsum(dim=1)
|
| 134 |
+
input_image_mask = torch.logical_and(indices <= (image_embeds.shape[0] - ground_truth_image_embeds.shape[0]), image_mask)
|
| 135 |
+
gt_image_mask = torch.logical_and(image_mask, ~input_image_mask)
|
| 136 |
+
input_image_mask = input_image_mask.unsqueeze(-1).expand_as(input_embeds)
|
| 137 |
+
input_embeds = input_embeds.masked_scatter(input_image_mask, input_image_embeds)
|
| 138 |
+
|
| 139 |
+
image_prefill_embeds = model.image_prefill_embeds(
|
| 140 |
+
torch.arange(81, device=model.device).long()
|
| 141 |
+
)
|
| 142 |
+
input_embeds = input_embeds.masked_scatter(gt_image_mask.unsqueeze(-1).expand_as(input_embeds), image_prefill_embeds)
|
| 143 |
+
|
| 144 |
+
position_ids, _ = model.get_rope_index(
|
| 145 |
+
inputs['input_ids'],
|
| 146 |
+
inputs['image_grid_thw'],
|
| 147 |
+
attention_mask=inputs['attention_mask'])
|
| 148 |
+
position_ids = position_ids.contiguous()
|
| 149 |
+
outputs = model(inputs_embeds=input_embeds, position_ids=position_ids, attention_mask=inputs['attention_mask'], return_dict=True)
|
| 150 |
+
output_image_embeddings = outputs.image_embeddings[:, :-1, :]
|
| 151 |
+
output_image_embeddings = output_image_embeddings[gt_image_mask[:, 1:]]
|
| 152 |
+
return output_image_embeddings, input_image_embeds, inputs['image_grid_thw']
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
class NexusGenAutoregressiveModelStateDictConverter:
|
| 156 |
+
def __init__(self):
|
| 157 |
+
pass
|
| 158 |
+
|
| 159 |
+
def from_civitai(self, state_dict):
|
| 160 |
+
state_dict = {"model." + key: value for key, value in state_dict.items()}
|
| 161 |
+
return state_dict
|
diffsynth/models/nexus_gen_ar_model.py
ADDED
|
@@ -0,0 +1,1143 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import re
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
from typing import Any, Dict, List, Optional, Tuple, Union
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
from torch.nn import CrossEntropyLoss
|
| 9 |
+
|
| 10 |
+
from transformers.cache_utils import Cache
|
| 11 |
+
from transformers.generation import GenerationMixin, LogitsProcessorList, StoppingCriteriaList, GenerationConfig, GenerateDecoderOnlyOutput, GenerateEncoderDecoderOutput
|
| 12 |
+
from transformers.utils import add_start_docstrings_to_model_forward, logging, replace_return_docstrings
|
| 13 |
+
from transformers.modeling_outputs import ModelOutput
|
| 14 |
+
from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLConfig
|
| 15 |
+
from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
|
| 16 |
+
Qwen2_5_VisionTransformerPretrainedModel,
|
| 17 |
+
Qwen2_5_VLModel,
|
| 18 |
+
Qwen2_5_VLPreTrainedModel,
|
| 19 |
+
QWEN2_5_VL_INPUTS_DOCSTRING,
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
from transformers.feature_extraction_utils import BatchFeature
|
| 23 |
+
from transformers.image_utils import ImageInput, VideoInput
|
| 24 |
+
from transformers.processing_utils import ProcessingKwargs, ProcessorMixin, Unpack, VideosKwargs
|
| 25 |
+
from transformers.tokenization_utils_base import PreTokenizedInput, TextInput
|
| 26 |
+
|
| 27 |
+
GenerateNonBeamOutput = Union[GenerateDecoderOnlyOutput, GenerateEncoderDecoderOutput]
|
| 28 |
+
|
| 29 |
+
logger = logging.get_logger(__name__)
|
| 30 |
+
|
| 31 |
+
_CONFIG_FOR_DOC = "Qwen2_5_VLConfig"
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
@dataclass
|
| 35 |
+
class Qwen2_5_VLCausalLMOutputWithPast(ModelOutput):
|
| 36 |
+
"""
|
| 37 |
+
Base class for Qwen2_5_VL causal language model (or autoregressive) outputs.
|
| 38 |
+
|
| 39 |
+
Args:
|
| 40 |
+
loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
|
| 41 |
+
Language modeling loss (for next-token prediction).
|
| 42 |
+
logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
|
| 43 |
+
Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
|
| 44 |
+
past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
|
| 45 |
+
Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
|
| 46 |
+
`(batch_size, num_heads, sequence_length, embed_size_per_head)`)
|
| 47 |
+
|
| 48 |
+
Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
|
| 49 |
+
`past_key_values` input) to speed up sequential decoding.
|
| 50 |
+
hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
|
| 51 |
+
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
|
| 52 |
+
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
|
| 53 |
+
|
| 54 |
+
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
|
| 55 |
+
attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
|
| 56 |
+
Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
|
| 57 |
+
sequence_length)`.
|
| 58 |
+
|
| 59 |
+
Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
|
| 60 |
+
heads.
|
| 61 |
+
rope_deltas (`torch.LongTensor` of shape `(batch_size, )`, *optional*):
|
| 62 |
+
The rope index difference between sequence length and multimodal rope.
|
| 63 |
+
"""
|
| 64 |
+
|
| 65 |
+
loss: Optional[torch.FloatTensor] = None
|
| 66 |
+
logits: torch.FloatTensor = None
|
| 67 |
+
image_embeddings: torch.FloatTensor = None
|
| 68 |
+
past_key_values: Optional[List[torch.FloatTensor]] = None
|
| 69 |
+
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
|
| 70 |
+
attentions: Optional[Tuple[torch.FloatTensor]] = None
|
| 71 |
+
rope_deltas: Optional[torch.LongTensor] = None
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
class Qwen2_5_VLForConditionalGeneration(Qwen2_5_VLPreTrainedModel, GenerationMixin):
|
| 75 |
+
_tied_weights_keys = ["lm_head.weight"]
|
| 76 |
+
config_class = Qwen2_5_VLConfig
|
| 77 |
+
_no_split_modules = ["Qwen2_5_VLDecoderLayer", "Qwen2_5_VLVisionBlock"]
|
| 78 |
+
|
| 79 |
+
def __init__(self, config):
|
| 80 |
+
super().__init__(config)
|
| 81 |
+
self.visual = Qwen2_5_VisionTransformerPretrainedModel._from_config(config.vision_config)
|
| 82 |
+
self.model = Qwen2_5_VLModel(config)
|
| 83 |
+
self.vocab_size = config.vocab_size
|
| 84 |
+
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
| 85 |
+
self.vision_head = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
|
| 86 |
+
self.rope_deltas = None # cache rope_deltas here
|
| 87 |
+
self.image_prefill_embeds = nn.Embedding(81, config.hidden_size)
|
| 88 |
+
|
| 89 |
+
# Initialize weights and apply final processing
|
| 90 |
+
self.post_init()
|
| 91 |
+
|
| 92 |
+
def get_input_embeddings(self):
|
| 93 |
+
return self.model.embed_tokens
|
| 94 |
+
|
| 95 |
+
def set_input_embeddings(self, value):
|
| 96 |
+
self.model.embed_tokens = value
|
| 97 |
+
|
| 98 |
+
def get_output_embeddings(self):
|
| 99 |
+
return self.lm_head
|
| 100 |
+
|
| 101 |
+
def set_output_embeddings(self, new_embeddings):
|
| 102 |
+
self.lm_head = new_embeddings
|
| 103 |
+
|
| 104 |
+
def set_decoder(self, decoder):
|
| 105 |
+
self.model = decoder
|
| 106 |
+
|
| 107 |
+
def get_decoder(self):
|
| 108 |
+
return self.model
|
| 109 |
+
|
| 110 |
+
def get_rope_index(
|
| 111 |
+
self,
|
| 112 |
+
input_ids: Optional[torch.LongTensor] = None,
|
| 113 |
+
image_grid_thw: Optional[torch.LongTensor] = None,
|
| 114 |
+
video_grid_thw: Optional[torch.LongTensor] = None,
|
| 115 |
+
second_per_grid_ts: Optional[torch.Tensor] = None,
|
| 116 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 117 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 118 |
+
"""
|
| 119 |
+
Calculate the 3D rope index based on image and video's temporal, height and width in LLM.
|
| 120 |
+
|
| 121 |
+
Explanation:
|
| 122 |
+
Each embedding sequence contains vision embedding and text embedding or just contains text embedding.
|
| 123 |
+
|
| 124 |
+
For pure text embedding sequence, the rotary position embedding has no difference with modern LLMs.
|
| 125 |
+
Examples:
|
| 126 |
+
input_ids: [T T T T T], here T is for text.
|
| 127 |
+
temporal position_ids: [0, 1, 2, 3, 4]
|
| 128 |
+
height position_ids: [0, 1, 2, 3, 4]
|
| 129 |
+
width position_ids: [0, 1, 2, 3, 4]
|
| 130 |
+
|
| 131 |
+
For vision and text embedding sequence, we calculate 3D rotary position embedding for vision part
|
| 132 |
+
and 1D rotary position embedding for text part.
|
| 133 |
+
Examples:
|
| 134 |
+
Temporal (Time): 3 patches, representing different segments of the video in time.
|
| 135 |
+
Height: 2 patches, dividing each frame vertically.
|
| 136 |
+
Width: 2 patches, dividing each frame horizontally.
|
| 137 |
+
We also have some important parameters:
|
| 138 |
+
fps (Frames Per Second): The video's frame rate, set to 1. This means one frame is processed each second.
|
| 139 |
+
tokens_per_second: This is a crucial parameter. It dictates how many "time-steps" or "temporal tokens" are conceptually packed into a one-second interval of the video. In this case, we have 25 tokens per second. So each second of the video will be represented with 25 separate time points. It essentially defines the temporal granularity.
|
| 140 |
+
temporal_patch_size: The number of frames that compose one temporal patch. Here, it's 2 frames.
|
| 141 |
+
interval: The step size for the temporal position IDs, calculated as tokens_per_second * temporal_patch_size / fps. In this case, 25 * 2 / 1 = 50. This means that each temporal patch will be have a difference of 50 in the temporal position IDs.
|
| 142 |
+
input_ids: [V V V V V V V V V V V V T T T T T], here V is for vision.
|
| 143 |
+
vision temporal position_ids: [0, 0, 0, 0, 50, 50, 50, 50, 100, 100, 100, 100]
|
| 144 |
+
vision height position_ids: [0, 0, 1, 1, 0, 0, 1, 1, 0, 0, 1, 1]
|
| 145 |
+
vision width position_ids: [0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1]
|
| 146 |
+
text temporal position_ids: [101, 102, 103, 104, 105]
|
| 147 |
+
text height position_ids: [101, 102, 103, 104, 105]
|
| 148 |
+
text width position_ids: [101, 102, 103, 104, 105]
|
| 149 |
+
Here we calculate the text start position_ids as the max vision position_ids plus 1.
|
| 150 |
+
|
| 151 |
+
Args:
|
| 152 |
+
input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
|
| 153 |
+
Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
|
| 154 |
+
it.
|
| 155 |
+
image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
|
| 156 |
+
The temporal, height and width of feature shape of each image in LLM.
|
| 157 |
+
video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
|
| 158 |
+
The temporal, height and width of feature shape of each video in LLM.
|
| 159 |
+
second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):
|
| 160 |
+
The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.
|
| 161 |
+
attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
|
| 162 |
+
Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
|
| 163 |
+
|
| 164 |
+
- 1 for tokens that are **not masked**,
|
| 165 |
+
- 0 for tokens that are **masked**.
|
| 166 |
+
|
| 167 |
+
Returns:
|
| 168 |
+
position_ids (`torch.LongTensor` of shape `(3, batch_size, sequence_length)`)
|
| 169 |
+
mrope_position_deltas (`torch.Tensor` of shape `(batch_size)`)
|
| 170 |
+
"""
|
| 171 |
+
spatial_merge_size = self.config.vision_config.spatial_merge_size
|
| 172 |
+
image_token_id = self.config.image_token_id
|
| 173 |
+
video_token_id = self.config.video_token_id
|
| 174 |
+
vision_start_token_id = self.config.vision_start_token_id
|
| 175 |
+
mrope_position_deltas = []
|
| 176 |
+
if input_ids is not None and (image_grid_thw is not None or video_grid_thw is not None):
|
| 177 |
+
total_input_ids = input_ids
|
| 178 |
+
if attention_mask is None:
|
| 179 |
+
attention_mask = torch.ones_like(total_input_ids)
|
| 180 |
+
position_ids = torch.ones(
|
| 181 |
+
3,
|
| 182 |
+
input_ids.shape[0],
|
| 183 |
+
input_ids.shape[1],
|
| 184 |
+
dtype=input_ids.dtype,
|
| 185 |
+
device=input_ids.device,
|
| 186 |
+
)
|
| 187 |
+
image_index, video_index = 0, 0
|
| 188 |
+
attention_mask = attention_mask.to(total_input_ids.device)
|
| 189 |
+
for i, input_ids in enumerate(total_input_ids):
|
| 190 |
+
input_ids = input_ids[attention_mask[i] == 1]
|
| 191 |
+
image_nums, video_nums = 0, 0
|
| 192 |
+
vision_start_indices = torch.argwhere(input_ids == vision_start_token_id).squeeze(1)
|
| 193 |
+
vision_tokens = input_ids[vision_start_indices + 1]
|
| 194 |
+
image_nums = (vision_tokens == image_token_id).sum()
|
| 195 |
+
video_nums = (vision_tokens == video_token_id).sum()
|
| 196 |
+
input_tokens = input_ids.tolist()
|
| 197 |
+
llm_pos_ids_list: list = []
|
| 198 |
+
st = 0
|
| 199 |
+
remain_images, remain_videos = image_nums, video_nums
|
| 200 |
+
for _ in range(image_nums + video_nums):
|
| 201 |
+
if image_token_id in input_tokens and remain_images > 0:
|
| 202 |
+
ed_image = input_tokens.index(image_token_id, st)
|
| 203 |
+
else:
|
| 204 |
+
ed_image = len(input_tokens) + 1
|
| 205 |
+
if video_token_id in input_tokens and remain_videos > 0:
|
| 206 |
+
ed_video = input_tokens.index(video_token_id, st)
|
| 207 |
+
else:
|
| 208 |
+
ed_video = len(input_tokens) + 1
|
| 209 |
+
if ed_image < ed_video:
|
| 210 |
+
t, h, w = (
|
| 211 |
+
image_grid_thw[image_index][0],
|
| 212 |
+
image_grid_thw[image_index][1],
|
| 213 |
+
image_grid_thw[image_index][2],
|
| 214 |
+
)
|
| 215 |
+
second_per_grid_t = 0
|
| 216 |
+
image_index += 1
|
| 217 |
+
remain_images -= 1
|
| 218 |
+
ed = ed_image
|
| 219 |
+
|
| 220 |
+
else:
|
| 221 |
+
t, h, w = (
|
| 222 |
+
video_grid_thw[video_index][0],
|
| 223 |
+
video_grid_thw[video_index][1],
|
| 224 |
+
video_grid_thw[video_index][2],
|
| 225 |
+
)
|
| 226 |
+
if second_per_grid_ts is not None:
|
| 227 |
+
second_per_grid_t = second_per_grid_ts[video_index]
|
| 228 |
+
else:
|
| 229 |
+
second_per_grid_t = 1.0
|
| 230 |
+
video_index += 1
|
| 231 |
+
remain_videos -= 1
|
| 232 |
+
ed = ed_video
|
| 233 |
+
llm_grid_t, llm_grid_h, llm_grid_w = (
|
| 234 |
+
t.item(),
|
| 235 |
+
h.item() // spatial_merge_size,
|
| 236 |
+
w.item() // spatial_merge_size,
|
| 237 |
+
)
|
| 238 |
+
text_len = ed - st
|
| 239 |
+
|
| 240 |
+
st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0
|
| 241 |
+
llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
|
| 242 |
+
|
| 243 |
+
range_tensor = torch.arange(llm_grid_t).view(-1, 1)
|
| 244 |
+
expanded_range = range_tensor.expand(-1, llm_grid_h * llm_grid_w)
|
| 245 |
+
|
| 246 |
+
time_tensor = expanded_range * second_per_grid_t * self.config.vision_config.tokens_per_second
|
| 247 |
+
|
| 248 |
+
time_tensor_long = time_tensor.long()
|
| 249 |
+
t_index = time_tensor_long.flatten()
|
| 250 |
+
|
| 251 |
+
h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten()
|
| 252 |
+
w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten()
|
| 253 |
+
llm_pos_ids_list.append(torch.stack([t_index, h_index, w_index]) + text_len + st_idx)
|
| 254 |
+
st = ed + llm_grid_t * llm_grid_h * llm_grid_w
|
| 255 |
+
|
| 256 |
+
if st < len(input_tokens):
|
| 257 |
+
st_idx = llm_pos_ids_list[-1].max() + 1 if len(llm_pos_ids_list) > 0 else 0
|
| 258 |
+
text_len = len(input_tokens) - st
|
| 259 |
+
llm_pos_ids_list.append(torch.arange(text_len).view(1, -1).expand(3, -1) + st_idx)
|
| 260 |
+
|
| 261 |
+
llm_positions = torch.cat(llm_pos_ids_list, dim=1).reshape(3, -1)
|
| 262 |
+
position_ids[..., i, attention_mask[i] == 1] = llm_positions.to(position_ids.device)
|
| 263 |
+
mrope_position_deltas.append(llm_positions.max() + 1 - len(total_input_ids[i]))
|
| 264 |
+
mrope_position_deltas = torch.tensor(mrope_position_deltas, device=input_ids.device).unsqueeze(1)
|
| 265 |
+
return position_ids, mrope_position_deltas
|
| 266 |
+
else:
|
| 267 |
+
if attention_mask is not None:
|
| 268 |
+
position_ids = attention_mask.long().cumsum(-1) - 1
|
| 269 |
+
position_ids.masked_fill_(attention_mask == 0, 1)
|
| 270 |
+
position_ids = position_ids.unsqueeze(0).expand(3, -1, -1).to(attention_mask.device)
|
| 271 |
+
max_position_ids = position_ids.max(0, keepdim=False)[0].max(-1, keepdim=True)[0]
|
| 272 |
+
mrope_position_deltas = max_position_ids + 1 - attention_mask.shape[-1]
|
| 273 |
+
else:
|
| 274 |
+
position_ids = (
|
| 275 |
+
torch.arange(input_ids.shape[1], device=input_ids.device)
|
| 276 |
+
.view(1, 1, -1)
|
| 277 |
+
.expand(3, input_ids.shape[0], -1)
|
| 278 |
+
)
|
| 279 |
+
mrope_position_deltas = torch.zeros(
|
| 280 |
+
[input_ids.shape[0], 1],
|
| 281 |
+
device=input_ids.device,
|
| 282 |
+
dtype=input_ids.dtype,
|
| 283 |
+
)
|
| 284 |
+
|
| 285 |
+
return position_ids, mrope_position_deltas
|
| 286 |
+
|
| 287 |
+
@add_start_docstrings_to_model_forward(QWEN2_5_VL_INPUTS_DOCSTRING)
|
| 288 |
+
@replace_return_docstrings(output_type=Qwen2_5_VLCausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
|
| 289 |
+
def forward(
|
| 290 |
+
self,
|
| 291 |
+
input_ids: torch.LongTensor = None,
|
| 292 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 293 |
+
position_ids: Optional[torch.LongTensor] = None,
|
| 294 |
+
past_key_values: Optional[List[torch.FloatTensor]] = None,
|
| 295 |
+
inputs_embeds: Optional[torch.FloatTensor] = None,
|
| 296 |
+
labels: Optional[torch.LongTensor] = None,
|
| 297 |
+
use_cache: Optional[bool] = None,
|
| 298 |
+
output_attentions: Optional[bool] = None,
|
| 299 |
+
output_hidden_states: Optional[bool] = None,
|
| 300 |
+
return_dict: Optional[bool] = None,
|
| 301 |
+
pixel_values: Optional[torch.Tensor] = None,
|
| 302 |
+
pixel_values_videos: Optional[torch.FloatTensor] = None,
|
| 303 |
+
image_grid_thw: Optional[torch.LongTensor] = None,
|
| 304 |
+
video_grid_thw: Optional[torch.LongTensor] = None,
|
| 305 |
+
rope_deltas: Optional[torch.LongTensor] = None,
|
| 306 |
+
cache_position: Optional[torch.LongTensor] = None,
|
| 307 |
+
second_per_grid_ts: Optional[torch.Tensor] = None,
|
| 308 |
+
image_embeddings: Optional[torch.Tensor] = None,
|
| 309 |
+
token_loss_weight: Optional[float] = 0.1,
|
| 310 |
+
img_loss_weight: Optional[float] = 1.0,
|
| 311 |
+
) -> Union[Tuple, Qwen2_5_VLCausalLMOutputWithPast]:
|
| 312 |
+
r"""
|
| 313 |
+
labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
| 314 |
+
Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
|
| 315 |
+
config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
|
| 316 |
+
(masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
|
| 317 |
+
|
| 318 |
+
Returns:
|
| 319 |
+
|
| 320 |
+
Example:
|
| 321 |
+
|
| 322 |
+
```python
|
| 323 |
+
>>> from PIL import Image
|
| 324 |
+
>>> import requests
|
| 325 |
+
>>> from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration
|
| 326 |
+
|
| 327 |
+
>>> model = Qwen2_5_VLForConditionalGeneration.from_pretrained("Qwen/Qwen2.5-VL-7B-Instruct")
|
| 328 |
+
>>> processor = AutoProcessor.from_pretrained("Qwen/Qwen2.5-VL-7B-Instruct")
|
| 329 |
+
|
| 330 |
+
>>> messages = [
|
| 331 |
+
{
|
| 332 |
+
"role": "user",
|
| 333 |
+
"content": [
|
| 334 |
+
{"type": "image"},
|
| 335 |
+
{"type": "text", "text": "What is shown in this image?"},
|
| 336 |
+
],
|
| 337 |
+
},
|
| 338 |
+
]
|
| 339 |
+
>>> url = "https://www.ilankelman.org/stopsigns/australia.jpg"
|
| 340 |
+
>>> image = Image.open(requests.get(url, stream=True).raw)
|
| 341 |
+
|
| 342 |
+
>>> text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 343 |
+
>>> inputs = processor(text=[text], images=[image], vision_infos=[vision_infos])
|
| 344 |
+
|
| 345 |
+
>>> # Generate
|
| 346 |
+
>>> generate_ids = model.generate(inputs.input_ids, max_length=30)
|
| 347 |
+
>>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
|
| 348 |
+
"The image shows a street scene with a red stop sign in the foreground. In the background, there is a large red gate with Chinese characters ..."
|
| 349 |
+
```"""
|
| 350 |
+
|
| 351 |
+
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
| 352 |
+
output_hidden_states = (
|
| 353 |
+
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
| 354 |
+
)
|
| 355 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 356 |
+
|
| 357 |
+
if inputs_embeds is None:
|
| 358 |
+
# test feature
|
| 359 |
+
inputs_embeds = self.model.embed_tokens(input_ids)
|
| 360 |
+
# for image encoding and training
|
| 361 |
+
if pixel_values is not None:
|
| 362 |
+
pixel_values = pixel_values.type(self.visual.dtype)
|
| 363 |
+
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
|
| 364 |
+
n_image_tokens = (input_ids == self.config.image_token_id).sum().item()
|
| 365 |
+
n_image_features = image_embeds.shape[0]
|
| 366 |
+
if n_image_tokens != n_image_features:
|
| 367 |
+
raise ValueError(
|
| 368 |
+
f"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}"
|
| 369 |
+
)
|
| 370 |
+
|
| 371 |
+
mask = input_ids == self.config.image_token_id
|
| 372 |
+
mask_unsqueezed = mask.unsqueeze(-1)
|
| 373 |
+
mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)
|
| 374 |
+
image_mask = mask_expanded.to(inputs_embeds.device)
|
| 375 |
+
|
| 376 |
+
image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
|
| 377 |
+
inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
|
| 378 |
+
|
| 379 |
+
if pixel_values_videos is not None:
|
| 380 |
+
pixel_values_videos = pixel_values_videos.type(self.visual.dtype)
|
| 381 |
+
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw)
|
| 382 |
+
n_video_tokens = (input_ids == self.config.video_token_id).sum().item()
|
| 383 |
+
n_video_features = video_embeds.shape[0]
|
| 384 |
+
if n_video_tokens != n_video_features:
|
| 385 |
+
raise ValueError(
|
| 386 |
+
f"Video features and video tokens do not match: tokens: {n_video_tokens}, features {n_video_features}"
|
| 387 |
+
)
|
| 388 |
+
|
| 389 |
+
mask = input_ids == self.config.video_token_id
|
| 390 |
+
mask_unsqueezed = mask.unsqueeze(-1)
|
| 391 |
+
mask_expanded = mask_unsqueezed.expand_as(inputs_embeds)
|
| 392 |
+
video_mask = mask_expanded.to(inputs_embeds.device)
|
| 393 |
+
|
| 394 |
+
video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
|
| 395 |
+
inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
|
| 396 |
+
|
| 397 |
+
if attention_mask is not None:
|
| 398 |
+
attention_mask = attention_mask.to(inputs_embeds.device)
|
| 399 |
+
|
| 400 |
+
# if we get 4D attention mask we cannot calculate rope deltas anymore. TODO @raushan fixme
|
| 401 |
+
if position_ids is None and (attention_mask is None or attention_mask.ndim == 2):
|
| 402 |
+
# calculate RoPE index once per generation in the pre-fill stage only
|
| 403 |
+
if (
|
| 404 |
+
(cache_position is not None and cache_position[0] == 0)
|
| 405 |
+
or self.rope_deltas is None
|
| 406 |
+
or (past_key_values is None or past_key_values.get_seq_length() == 0)
|
| 407 |
+
):
|
| 408 |
+
position_ids, rope_deltas = self.get_rope_index(
|
| 409 |
+
input_ids,
|
| 410 |
+
image_grid_thw,
|
| 411 |
+
video_grid_thw,
|
| 412 |
+
second_per_grid_ts,
|
| 413 |
+
attention_mask,
|
| 414 |
+
)
|
| 415 |
+
self.rope_deltas = rope_deltas
|
| 416 |
+
# then use the prev pre-calculated rope-deltas to get the correct position ids
|
| 417 |
+
else:
|
| 418 |
+
batch_size, seq_length, _ = inputs_embeds.shape
|
| 419 |
+
delta = (
|
| 420 |
+
(cache_position[0] + self.rope_deltas).to(inputs_embeds.device)
|
| 421 |
+
if cache_position is not None
|
| 422 |
+
else 0
|
| 423 |
+
)
|
| 424 |
+
position_ids = torch.arange(seq_length, device=inputs_embeds.device)
|
| 425 |
+
position_ids = position_ids.view(1, -1).expand(batch_size, -1)
|
| 426 |
+
if cache_position is not None: # otherwise `deltas` is an int `0`
|
| 427 |
+
delta = delta.repeat_interleave(batch_size // delta.shape[0], dim=0)
|
| 428 |
+
position_ids = position_ids.add(delta)
|
| 429 |
+
position_ids = position_ids.unsqueeze(0).expand(3, -1, -1)
|
| 430 |
+
# position_ids [3, B, L]
|
| 431 |
+
|
| 432 |
+
outputs = self.model(
|
| 433 |
+
input_ids=None,
|
| 434 |
+
position_ids=position_ids,
|
| 435 |
+
attention_mask=attention_mask,
|
| 436 |
+
past_key_values=past_key_values,
|
| 437 |
+
inputs_embeds=inputs_embeds,
|
| 438 |
+
use_cache=use_cache,
|
| 439 |
+
output_attentions=output_attentions,
|
| 440 |
+
output_hidden_states=output_hidden_states,
|
| 441 |
+
return_dict=return_dict,
|
| 442 |
+
cache_position=cache_position,
|
| 443 |
+
)
|
| 444 |
+
|
| 445 |
+
hidden_states = outputs[0]
|
| 446 |
+
logits = self.lm_head(hidden_states)
|
| 447 |
+
image_embeds = self.vision_head(hidden_states)
|
| 448 |
+
|
| 449 |
+
loss = None
|
| 450 |
+
if labels is not None:
|
| 451 |
+
# Upcast to float if we need to compute the loss to avoid potential precision issues
|
| 452 |
+
# prepare labels for logits
|
| 453 |
+
logits_labels = labels.clone().detach()
|
| 454 |
+
image_tokens = (labels == self.config.image_token_id)
|
| 455 |
+
logits_labels[image_tokens] = -100
|
| 456 |
+
|
| 457 |
+
logits = logits.float()
|
| 458 |
+
# Shift so that tokens < n predict n
|
| 459 |
+
shift_logits = logits[..., :-1, :].contiguous()
|
| 460 |
+
shift_labels = logits_labels[..., 1:].contiguous()
|
| 461 |
+
# Flatten the tokens
|
| 462 |
+
loss_fct = CrossEntropyLoss()
|
| 463 |
+
shift_logits = shift_logits.view(-1, self.config.vocab_size)
|
| 464 |
+
shift_labels = shift_labels.view(-1)
|
| 465 |
+
# Enable model parallelism
|
| 466 |
+
shift_labels = shift_labels.to(shift_logits.device)
|
| 467 |
+
loss = loss_fct(shift_logits, shift_labels) * token_loss_weight
|
| 468 |
+
|
| 469 |
+
shift_image_tokens_2d = (labels[..., 1:].contiguous() == self.config.image_token_id) # (B, L-1)
|
| 470 |
+
shifted_image_embeds = image_embeds[:, :-1, :].contiguous() # (B, L-1, D)
|
| 471 |
+
masked_image_embeds = shifted_image_embeds[shift_image_tokens_2d] # (num_image_tokens, D)
|
| 472 |
+
|
| 473 |
+
mse_loss_fct = nn.MSELoss()
|
| 474 |
+
mse_loss_fct = mse_loss_fct.to(shift_logits.device)
|
| 475 |
+
if image_embeddings is None:
|
| 476 |
+
image_embeddings = torch.zeros_like(masked_image_embeds)
|
| 477 |
+
img_loss = mse_loss_fct(masked_image_embeds, image_embeddings)
|
| 478 |
+
|
| 479 |
+
cos_sim = torch.cosine_similarity(
|
| 480 |
+
masked_image_embeds,
|
| 481 |
+
image_embeddings,
|
| 482 |
+
dim=-1
|
| 483 |
+
)
|
| 484 |
+
cos_loss = (1 - cos_sim).mean()
|
| 485 |
+
img_loss = 0.5 * img_loss + 0.5 * cos_loss
|
| 486 |
+
# fix nan for empty image tokens
|
| 487 |
+
if image_embeddings.size(0) == 0:
|
| 488 |
+
img_loss = img_loss.nan_to_num(0.0)
|
| 489 |
+
# combine the loss
|
| 490 |
+
loss = loss + img_loss_weight * img_loss
|
| 491 |
+
|
| 492 |
+
if not return_dict:
|
| 493 |
+
output = (logits,) + outputs[1:]
|
| 494 |
+
return (loss,) + output if loss is not None else output
|
| 495 |
+
|
| 496 |
+
return Qwen2_5_VLCausalLMOutputWithPast(
|
| 497 |
+
loss=loss,
|
| 498 |
+
logits=logits,
|
| 499 |
+
image_embeddings=image_embeds,
|
| 500 |
+
past_key_values=outputs.past_key_values,
|
| 501 |
+
hidden_states=outputs.hidden_states,
|
| 502 |
+
attentions=outputs.attentions,
|
| 503 |
+
rope_deltas=self.rope_deltas,
|
| 504 |
+
)
|
| 505 |
+
|
| 506 |
+
|
| 507 |
+
|
| 508 |
+
def _sample(
|
| 509 |
+
self,
|
| 510 |
+
input_ids: torch.LongTensor,
|
| 511 |
+
logits_processor: LogitsProcessorList,
|
| 512 |
+
stopping_criteria: StoppingCriteriaList,
|
| 513 |
+
generation_config: GenerationConfig,
|
| 514 |
+
synced_gpus: bool,
|
| 515 |
+
streamer: Optional["BaseStreamer"],
|
| 516 |
+
**model_kwargs,
|
| 517 |
+
) -> Union[GenerateNonBeamOutput, torch.LongTensor]:
|
| 518 |
+
r"""
|
| 519 |
+
Generates sequences of token ids for models with a language modeling head using **multinomial sampling** and
|
| 520 |
+
can be used for text-decoder, text-to-text, speech-to-text, and vision-to-text models.
|
| 521 |
+
|
| 522 |
+
Parameters:
|
| 523 |
+
input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
|
| 524 |
+
The sequence used as a prompt for the generation.
|
| 525 |
+
logits_processor (`LogitsProcessorList`):
|
| 526 |
+
An instance of [`LogitsProcessorList`]. List of instances of class derived from [`LogitsProcessor`]
|
| 527 |
+
used to modify the prediction scores of the language modeling head applied at each generation step.
|
| 528 |
+
stopping_criteria (`StoppingCriteriaList`):
|
| 529 |
+
An instance of [`StoppingCriteriaList`]. List of instances of class derived from [`StoppingCriteria`]
|
| 530 |
+
used to tell if the generation loop should stop.
|
| 531 |
+
generation_config ([`~generation.GenerationConfig`]):
|
| 532 |
+
The generation configuration to be used as parametrization of the decoding method.
|
| 533 |
+
synced_gpus (`bool`):
|
| 534 |
+
Whether to continue running the while loop until max_length (needed to avoid deadlocking with
|
| 535 |
+
`FullyShardedDataParallel` and DeepSpeed ZeRO Stage 3).
|
| 536 |
+
streamer (`BaseStreamer`, *optional*):
|
| 537 |
+
Streamer object that will be used to stream the generated sequences. Generated tokens are passed
|
| 538 |
+
through `streamer.put(token_ids)` and the streamer is responsible for any further processing.
|
| 539 |
+
model_kwargs:
|
| 540 |
+
Additional model specific kwargs will be forwarded to the `forward` function of the model. If model is
|
| 541 |
+
an encoder-decoder model the kwargs should include `encoder_outputs`.
|
| 542 |
+
|
| 543 |
+
Return:
|
| 544 |
+
[`~generation.GenerateDecoderOnlyOutput`], [`~generation.GenerateEncoderDecoderOutput`] or `torch.LongTensor`:
|
| 545 |
+
A `torch.LongTensor` containing the generated tokens (default behaviour) or a
|
| 546 |
+
[`~generation.GenerateDecoderOnlyOutput`] if `model.config.is_encoder_decoder=False` and
|
| 547 |
+
`return_dict_in_generate=True` or a [`~generation.GenerateEncoderDecoderOutput`] if
|
| 548 |
+
`model.config.is_encoder_decoder=True`.
|
| 549 |
+
"""
|
| 550 |
+
# init values
|
| 551 |
+
pad_token_id = generation_config._pad_token_tensor
|
| 552 |
+
output_attentions = generation_config.output_attentions
|
| 553 |
+
output_hidden_states = generation_config.output_hidden_states
|
| 554 |
+
output_scores = generation_config.output_scores
|
| 555 |
+
output_logits = generation_config.output_logits
|
| 556 |
+
return_dict_in_generate = generation_config.return_dict_in_generate
|
| 557 |
+
max_length = generation_config.max_length
|
| 558 |
+
has_eos_stopping_criteria = any(hasattr(criteria, "eos_token_id") for criteria in stopping_criteria)
|
| 559 |
+
do_sample = generation_config.do_sample
|
| 560 |
+
|
| 561 |
+
# init attention / hidden states / scores tuples
|
| 562 |
+
scores = () if (return_dict_in_generate and output_scores) else None
|
| 563 |
+
raw_logits = () if (return_dict_in_generate and output_logits) else None
|
| 564 |
+
decoder_attentions = () if (return_dict_in_generate and output_attentions) else None
|
| 565 |
+
cross_attentions = () if (return_dict_in_generate and output_attentions) else None
|
| 566 |
+
decoder_hidden_states = () if (return_dict_in_generate and output_hidden_states) else None
|
| 567 |
+
|
| 568 |
+
# if model is an encoder-decoder, retrieve encoder attention weights and hidden states
|
| 569 |
+
if return_dict_in_generate and self.config.is_encoder_decoder:
|
| 570 |
+
encoder_attentions = model_kwargs["encoder_outputs"].get("attentions") if output_attentions else None
|
| 571 |
+
encoder_hidden_states = (
|
| 572 |
+
model_kwargs["encoder_outputs"].get("hidden_states") if output_hidden_states else None
|
| 573 |
+
)
|
| 574 |
+
|
| 575 |
+
# keep track of which sequences are already finished
|
| 576 |
+
batch_size, cur_len = input_ids.shape
|
| 577 |
+
this_peer_finished = False
|
| 578 |
+
unfinished_sequences = torch.ones(batch_size, dtype=torch.long, device=input_ids.device)
|
| 579 |
+
model_kwargs = self._get_initial_cache_position(input_ids, model_kwargs)
|
| 580 |
+
|
| 581 |
+
model_forward = self.__call__
|
| 582 |
+
if isinstance(model_kwargs.get("past_key_values"), Cache):
|
| 583 |
+
is_compileable = model_kwargs["past_key_values"].is_compileable and self._supports_static_cache
|
| 584 |
+
is_compileable = is_compileable and not self.generation_config.disable_compile
|
| 585 |
+
if is_compileable and (
|
| 586 |
+
self.device.type == "cuda" or generation_config.compile_config._compile_all_devices
|
| 587 |
+
):
|
| 588 |
+
os.environ["TOKENIZERS_PARALLELISM"] = "0"
|
| 589 |
+
model_forward = self.get_compiled_call(generation_config.compile_config)
|
| 590 |
+
|
| 591 |
+
is_prefill = True
|
| 592 |
+
is_sampling_img = input_ids[:, -1] == self.config.vision_start_token_id
|
| 593 |
+
generation_image_grid_thw = model_kwargs.pop("generation_image_grid_thw", self.get_default_image_grid_thw())
|
| 594 |
+
num_img_tokens = self.get_num_image_tokens(generation_image_grid_thw)
|
| 595 |
+
output_image_embeddings = []
|
| 596 |
+
while self._has_unfinished_sequences(
|
| 597 |
+
this_peer_finished, synced_gpus, device=input_ids.device, cur_len=cur_len, max_length=max_length
|
| 598 |
+
):
|
| 599 |
+
# prepare model inputs
|
| 600 |
+
model_inputs = self.prepare_inputs_for_generation(input_ids, **model_kwargs)
|
| 601 |
+
|
| 602 |
+
# prepare prefilled embeds
|
| 603 |
+
model_inputs.update(self.prepare_prefilled_image_embeds(len(output_image_embeddings), num_img_tokens, is_sampling_img, **model_kwargs))
|
| 604 |
+
|
| 605 |
+
# parse position_ids from model_kwargs
|
| 606 |
+
model_inputs.update(self.prepare_image_position_ids(input_ids, generation_image_grid_thw, is_sampling_img, **model_kwargs))
|
| 607 |
+
|
| 608 |
+
# prepare variable output controls (note: some models won't accept all output controls)
|
| 609 |
+
model_inputs.update({"output_attentions": output_attentions} if output_attentions else {})
|
| 610 |
+
model_inputs.update({"output_hidden_states": output_hidden_states} if output_hidden_states else {})
|
| 611 |
+
|
| 612 |
+
if is_prefill:
|
| 613 |
+
outputs = self(**model_inputs, return_dict=True)
|
| 614 |
+
is_prefill = False
|
| 615 |
+
else:
|
| 616 |
+
outputs = model_forward(**model_inputs, return_dict=True)
|
| 617 |
+
|
| 618 |
+
# synced_gpus: don't waste resources running the code we don't need; kwargs must be updated before skipping
|
| 619 |
+
model_kwargs = self._update_model_kwargs_for_generation(
|
| 620 |
+
outputs,
|
| 621 |
+
model_kwargs,
|
| 622 |
+
is_encoder_decoder=self.config.is_encoder_decoder,
|
| 623 |
+
)
|
| 624 |
+
# TODO: support batch image sampling
|
| 625 |
+
if bool(is_sampling_img) and len(output_image_embeddings) < num_img_tokens:
|
| 626 |
+
output_image_embeddings.append(outputs.image_embeddings[:, -1, :].unsqueeze(1))
|
| 627 |
+
|
| 628 |
+
if synced_gpus and this_peer_finished:
|
| 629 |
+
continue
|
| 630 |
+
# Clone is needed to avoid keeping a hanging ref to outputs.logits which may be very large for first iteration
|
| 631 |
+
# (the clone itself is always small)
|
| 632 |
+
next_token_logits = outputs.logits[:, -1, :].clone().float()
|
| 633 |
+
next_token_logits = next_token_logits.to(input_ids.device)
|
| 634 |
+
|
| 635 |
+
# do not sample <vision_end> token
|
| 636 |
+
next_token_logits[:, self.config.vision_end_token_id] = -float('inf')
|
| 637 |
+
# pre-process distribution
|
| 638 |
+
next_token_scores = logits_processor(input_ids, next_token_logits)
|
| 639 |
+
# Store scores, attentions and hidden_states when required
|
| 640 |
+
if return_dict_in_generate:
|
| 641 |
+
if output_scores:
|
| 642 |
+
scores += (next_token_scores,)
|
| 643 |
+
if output_logits:
|
| 644 |
+
raw_logits += (next_token_logits,)
|
| 645 |
+
if output_attentions:
|
| 646 |
+
decoder_attentions += (
|
| 647 |
+
(outputs.decoder_attentions,) if self.config.is_encoder_decoder else (outputs.attentions,)
|
| 648 |
+
)
|
| 649 |
+
if self.config.is_encoder_decoder:
|
| 650 |
+
cross_attentions += (outputs.cross_attentions,)
|
| 651 |
+
|
| 652 |
+
if output_hidden_states:
|
| 653 |
+
decoder_hidden_states += (
|
| 654 |
+
(outputs.decoder_hidden_states,)
|
| 655 |
+
if self.config.is_encoder_decoder
|
| 656 |
+
else (outputs.hidden_states,)
|
| 657 |
+
)
|
| 658 |
+
|
| 659 |
+
# token selection
|
| 660 |
+
if do_sample:
|
| 661 |
+
probs = nn.functional.softmax(next_token_scores, dim=-1)
|
| 662 |
+
# TODO (joao): this OP throws "skipping cudagraphs due to ['incompatible ops']", find solution
|
| 663 |
+
next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)
|
| 664 |
+
# while not bool(is_sampling_img) and torch.any(next_tokens == self.config.vision_end_token_id):
|
| 665 |
+
# probs[:, self.config.vision_end_token_id] = 0
|
| 666 |
+
# next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)
|
| 667 |
+
else:
|
| 668 |
+
next_tokens = torch.argmax(next_token_scores, dim=-1)
|
| 669 |
+
|
| 670 |
+
# finished sentences should have their next token be a padding token
|
| 671 |
+
if has_eos_stopping_criteria:
|
| 672 |
+
next_tokens = next_tokens * unfinished_sequences + pad_token_id * (1 - unfinished_sequences)
|
| 673 |
+
|
| 674 |
+
#TODO: support batch image sample
|
| 675 |
+
if num_img_tokens is not None:
|
| 676 |
+
cur_img_tokens = (input_ids == self.config.vision_start_token_id).flip(dims=[1]).float().argmax(dim=1)
|
| 677 |
+
# check whether is sampling images
|
| 678 |
+
is_end_img = torch.logical_and(cur_img_tokens == num_img_tokens, is_sampling_img)
|
| 679 |
+
is_sampling_img = torch.logical_and(is_sampling_img, cur_img_tokens < num_img_tokens)
|
| 680 |
+
next_tokens[is_sampling_img] = self.config.image_token_id
|
| 681 |
+
# check whether to end sampling images
|
| 682 |
+
next_tokens[is_end_img] = self.config.vision_end_token_id
|
| 683 |
+
else:
|
| 684 |
+
# check whether to end sampling images
|
| 685 |
+
is_sampling_img = torch.logical_and(is_sampling_img, (next_tokens != self.config.vision_end_token_id))
|
| 686 |
+
# replace the next token with the image token if is sampling image
|
| 687 |
+
next_tokens[is_sampling_img] = self.config.image_token_id
|
| 688 |
+
# check whether to start sampling images
|
| 689 |
+
is_sampling_img = torch.logical_or(is_sampling_img, (next_tokens == self.config.vision_start_token_id))
|
| 690 |
+
|
| 691 |
+
# update generated ids, model inputs, and length for next step
|
| 692 |
+
input_ids = torch.cat([input_ids, next_tokens[:, None]], dim=-1)
|
| 693 |
+
|
| 694 |
+
if streamer is not None:
|
| 695 |
+
streamer.put(next_tokens.cpu())
|
| 696 |
+
|
| 697 |
+
unfinished_sequences = unfinished_sequences & ~stopping_criteria(input_ids, scores)
|
| 698 |
+
this_peer_finished = unfinished_sequences.max() == 0
|
| 699 |
+
cur_len += 1
|
| 700 |
+
|
| 701 |
+
# This is needed to properly delete outputs.logits which may be very large for first iteration
|
| 702 |
+
# Otherwise a reference to outputs is kept which keeps the logits alive in the next iteration
|
| 703 |
+
del outputs
|
| 704 |
+
|
| 705 |
+
if streamer is not None:
|
| 706 |
+
streamer.end()
|
| 707 |
+
|
| 708 |
+
# output the image embeddings
|
| 709 |
+
output_image_embeddings = torch.cat(output_image_embeddings, dim=1) if len(output_image_embeddings) > 0 else None
|
| 710 |
+
|
| 711 |
+
if return_dict_in_generate:
|
| 712 |
+
return GenerateDecoderOnlyAll2AllOutput(
|
| 713 |
+
sequences=input_ids,
|
| 714 |
+
scores=scores,
|
| 715 |
+
logits=raw_logits,
|
| 716 |
+
attentions=decoder_attentions,
|
| 717 |
+
hidden_states=decoder_hidden_states,
|
| 718 |
+
past_key_values=model_kwargs.get("past_key_values"),
|
| 719 |
+
output_image_embeddings=output_image_embeddings,
|
| 720 |
+
)
|
| 721 |
+
else:
|
| 722 |
+
return input_ids
|
| 723 |
+
|
| 724 |
+
|
| 725 |
+
def prepare_prefilled_image_embeds(self, cur_image_tokens, num_img_tokens, is_sampling_img, **model_kwargs):
|
| 726 |
+
if cur_image_tokens == 0 or cur_image_tokens > num_img_tokens or not bool(is_sampling_img):
|
| 727 |
+
return {}
|
| 728 |
+
# TODO: support batch image sample
|
| 729 |
+
image_idx = torch.tensor([cur_image_tokens-1]).to(self.device).long().unsqueeze(0)
|
| 730 |
+
inputs_embeds = self.image_prefill_embeds(image_idx)
|
| 731 |
+
return {"inputs_embeds": inputs_embeds}
|
| 732 |
+
|
| 733 |
+
|
| 734 |
+
def get_default_image_grid_thw(self,):
|
| 735 |
+
return torch.tensor([[1, 18, 18]]).to(self.device)
|
| 736 |
+
|
| 737 |
+
|
| 738 |
+
def get_num_image_tokens(self, image_grid_thw):
|
| 739 |
+
return int(torch.prod(image_grid_thw, dim=1).sum() // 4)
|
| 740 |
+
|
| 741 |
+
|
| 742 |
+
def _validate_model_kwargs(self, model_kwargs: Dict[str, Any]):
|
| 743 |
+
num_img_tokens = model_kwargs.pop("generation_image_grid_thw", None)
|
| 744 |
+
super()._validate_model_kwargs(model_kwargs)
|
| 745 |
+
model_kwargs["generation_image_grid_thw"] = num_img_tokens
|
| 746 |
+
|
| 747 |
+
def prepare_image_position_ids(self, input_ids, generation_image_grid_thw, is_sampling_img, **model_kwargs):
|
| 748 |
+
# Overwritten -- prepare position_ids for image tokens
|
| 749 |
+
cur_img_tokens = int((input_ids == self.config.vision_start_token_id).flip(dims=[1]).float().argmax(dim=1))
|
| 750 |
+
# TODO: support batch image sample
|
| 751 |
+
if cur_img_tokens > 0 and bool(is_sampling_img):
|
| 752 |
+
image_grid_thw = generation_image_grid_thw
|
| 753 |
+
if model_kwargs.get('image_grid_thw') is not None:
|
| 754 |
+
image_grid_thw = torch.cat([model_kwargs.get('image_grid_thw'), image_grid_thw])
|
| 755 |
+
remaining_img_tokens = self.get_num_image_tokens(generation_image_grid_thw) - cur_img_tokens
|
| 756 |
+
padding_ids = input_ids.new_full((1, remaining_img_tokens), fill_value=self.config.image_token_id)
|
| 757 |
+
padded_ids = torch.cat([input_ids, padding_ids], dim=1)
|
| 758 |
+
position_ids, _ = self.get_rope_index(padded_ids, image_grid_thw, None, None)
|
| 759 |
+
if model_kwargs.get("use_cache", True):
|
| 760 |
+
position_ids = position_ids[:, :, input_ids.shape[1] - 1].unsqueeze(-1)
|
| 761 |
+
else:
|
| 762 |
+
position_ids = position_ids[:, :, :input_ids.shape[1]]
|
| 763 |
+
return {"position_ids": position_ids}
|
| 764 |
+
return {}
|
| 765 |
+
|
| 766 |
+
def prepare_inputs_for_generation(
|
| 767 |
+
self,
|
| 768 |
+
input_ids,
|
| 769 |
+
past_key_values=None,
|
| 770 |
+
attention_mask=None,
|
| 771 |
+
inputs_embeds=None,
|
| 772 |
+
cache_position=None,
|
| 773 |
+
position_ids=None,
|
| 774 |
+
use_cache=True,
|
| 775 |
+
pixel_values=None,
|
| 776 |
+
pixel_values_videos=None,
|
| 777 |
+
image_grid_thw=None,
|
| 778 |
+
video_grid_thw=None,
|
| 779 |
+
second_per_grid_ts=None,
|
| 780 |
+
image_embeddings=None,
|
| 781 |
+
**kwargs,
|
| 782 |
+
):
|
| 783 |
+
# Overwritten -- in specific circumstances we don't want to forward image inputs to the model
|
| 784 |
+
|
| 785 |
+
model_inputs = super().prepare_inputs_for_generation(
|
| 786 |
+
input_ids,
|
| 787 |
+
past_key_values=past_key_values,
|
| 788 |
+
attention_mask=attention_mask,
|
| 789 |
+
inputs_embeds=inputs_embeds,
|
| 790 |
+
cache_position=cache_position,
|
| 791 |
+
position_ids=position_ids,
|
| 792 |
+
pixel_values=pixel_values,
|
| 793 |
+
pixel_values_videos=pixel_values_videos,
|
| 794 |
+
image_grid_thw=image_grid_thw,
|
| 795 |
+
video_grid_thw=video_grid_thw,
|
| 796 |
+
second_per_grid_ts=second_per_grid_ts,
|
| 797 |
+
use_cache=use_cache,
|
| 798 |
+
**kwargs,
|
| 799 |
+
)
|
| 800 |
+
|
| 801 |
+
# Qwen2-5-VL position_ids are prepared with rope_deltas in forward
|
| 802 |
+
model_inputs["position_ids"] = None
|
| 803 |
+
|
| 804 |
+
if cache_position[0] != 0:
|
| 805 |
+
model_inputs["pixel_values"] = None
|
| 806 |
+
model_inputs["pixel_values_videos"] = None
|
| 807 |
+
return model_inputs
|
| 808 |
+
|
| 809 |
+
def _get_image_nums_and_video_nums(
|
| 810 |
+
self,
|
| 811 |
+
input_ids: Optional[torch.LongTensor],
|
| 812 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 813 |
+
"""
|
| 814 |
+
Get the number of images and videos for each sample to calculate the separation length of the sample tensor.
|
| 815 |
+
These parameters are not passed through the processor to avoid unpredictable impacts from interface modifications.
|
| 816 |
+
|
| 817 |
+
Args:
|
| 818 |
+
input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
|
| 819 |
+
Indices of input sequence tokens in the vocabulary.
|
| 820 |
+
|
| 821 |
+
Returns:
|
| 822 |
+
image_nums (`torch.LongTensor` of shape `(batch_size, num_images_sample)`)
|
| 823 |
+
video_nums (`torch.LongTensor` of shape `(batch_size, num_videos_sample)`)
|
| 824 |
+
"""
|
| 825 |
+
image_token_id = self.config.image_token_id
|
| 826 |
+
video_token_id = self.config.video_token_id
|
| 827 |
+
vision_start_token_id = self.config.vision_start_token_id
|
| 828 |
+
|
| 829 |
+
vision_start_mask = input_ids == vision_start_token_id
|
| 830 |
+
vision_first_mask = torch.roll(vision_start_mask, shifts=1, dims=1)
|
| 831 |
+
image_mask = input_ids == image_token_id
|
| 832 |
+
video_mask = input_ids == video_token_id
|
| 833 |
+
image_nums = torch.sum(vision_first_mask & image_mask, dim=1)
|
| 834 |
+
video_nums = torch.sum(vision_first_mask & video_mask, dim=1)
|
| 835 |
+
|
| 836 |
+
return image_nums, video_nums
|
| 837 |
+
|
| 838 |
+
def _expand_inputs_for_generation(
|
| 839 |
+
self,
|
| 840 |
+
expand_size: int = 1,
|
| 841 |
+
is_encoder_decoder: bool = False,
|
| 842 |
+
input_ids: Optional[torch.LongTensor] = None,
|
| 843 |
+
**model_kwargs,
|
| 844 |
+
) -> Tuple[torch.LongTensor, Dict[str, Any]]:
|
| 845 |
+
# Overwritten -- Support for expanding tensors without a batch size dimension
|
| 846 |
+
# e.g., pixel_values, image_grid_thw, pixel_values_videos, video_grid_thw, second_per_grid_t
|
| 847 |
+
# pixel_values.shape[0] is sum(seqlen_images for samples)
|
| 848 |
+
# image_grid_thw.shape[0] is sum(num_images for samples)
|
| 849 |
+
|
| 850 |
+
if expand_size == 1:
|
| 851 |
+
return input_ids, model_kwargs
|
| 852 |
+
|
| 853 |
+
visual_keys = ["pixel_values", "image_grid_thw", "pixel_values_videos", "video_grid_thw", "second_per_grid_ts"]
|
| 854 |
+
|
| 855 |
+
def _expand_dict_for_generation_visual(dict_to_expand):
|
| 856 |
+
image_grid_thw = model_kwargs.get("image_grid_thw", None)
|
| 857 |
+
video_grid_thw = model_kwargs.get("video_grid_thw", None)
|
| 858 |
+
image_nums, video_nums = self._get_image_nums_and_video_nums(input_ids)
|
| 859 |
+
|
| 860 |
+
def _repeat_interleave_samples(x, lengths, repeat_times):
|
| 861 |
+
samples = torch.split(x, lengths)
|
| 862 |
+
repeat_args = [repeat_times] + [1] * (x.dim() - 1)
|
| 863 |
+
result = torch.cat([sample.repeat(*repeat_args) for sample in samples], dim=0)
|
| 864 |
+
return result
|
| 865 |
+
|
| 866 |
+
for key in dict_to_expand:
|
| 867 |
+
if key == "pixel_values":
|
| 868 |
+
# split images into samples
|
| 869 |
+
samples = torch.split(image_grid_thw, list(image_nums))
|
| 870 |
+
# compute the sequence length of images for each sample
|
| 871 |
+
lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
|
| 872 |
+
dict_to_expand[key] = _repeat_interleave_samples(
|
| 873 |
+
dict_to_expand[key], lengths=lengths, repeat_times=expand_size
|
| 874 |
+
)
|
| 875 |
+
elif key == "image_grid_thw":
|
| 876 |
+
# get the num of images for each sample
|
| 877 |
+
lengths = list(image_nums)
|
| 878 |
+
dict_to_expand[key] = _repeat_interleave_samples(
|
| 879 |
+
dict_to_expand[key], lengths=lengths, repeat_times=expand_size
|
| 880 |
+
)
|
| 881 |
+
elif key == "pixel_values_videos":
|
| 882 |
+
samples = torch.split(video_grid_thw, list(video_nums))
|
| 883 |
+
lengths = [torch.prod(sample, dim=1).sum() for sample in samples]
|
| 884 |
+
dict_to_expand[key] = _repeat_interleave_samples(
|
| 885 |
+
dict_to_expand[key], lengths=lengths, repeat_times=expand_size
|
| 886 |
+
)
|
| 887 |
+
elif key == "video_grid_thw":
|
| 888 |
+
lengths = list(video_nums)
|
| 889 |
+
dict_to_expand[key] = _repeat_interleave_samples(
|
| 890 |
+
dict_to_expand[key], lengths=lengths, repeat_times=expand_size
|
| 891 |
+
)
|
| 892 |
+
elif key == "second_per_grid_ts":
|
| 893 |
+
if not isinstance(dict_to_expand[key], list):
|
| 894 |
+
raise TypeError(
|
| 895 |
+
f"Expected value for key '{key}' to be a list, but got {type(dict_to_expand[key])} instead."
|
| 896 |
+
)
|
| 897 |
+
tensor = torch.tensor(dict_to_expand[key])
|
| 898 |
+
lengths = list(video_nums)
|
| 899 |
+
tensor = _repeat_interleave_samples(tensor, lengths=lengths, repeat_times=expand_size)
|
| 900 |
+
dict_to_expand[key] = tensor.tolist()
|
| 901 |
+
return dict_to_expand
|
| 902 |
+
|
| 903 |
+
def _expand_dict_for_generation(dict_to_expand):
|
| 904 |
+
for key in dict_to_expand:
|
| 905 |
+
if (
|
| 906 |
+
key != "cache_position"
|
| 907 |
+
and dict_to_expand[key] is not None
|
| 908 |
+
and isinstance(dict_to_expand[key], torch.Tensor)
|
| 909 |
+
and key not in visual_keys
|
| 910 |
+
):
|
| 911 |
+
dict_to_expand[key] = dict_to_expand[key].repeat_interleave(expand_size, dim=0)
|
| 912 |
+
return dict_to_expand
|
| 913 |
+
|
| 914 |
+
# input_ids is required for expanding visual inputs
|
| 915 |
+
# If input_ids is unavailable, visual inputs will not be used; therefore, there is no need to expand visual inputs.
|
| 916 |
+
if input_ids is not None and input_ids.numel() != 0:
|
| 917 |
+
model_kwargs = _expand_dict_for_generation_visual(model_kwargs)
|
| 918 |
+
|
| 919 |
+
if input_ids is not None:
|
| 920 |
+
input_ids = input_ids.repeat_interleave(expand_size, dim=0)
|
| 921 |
+
|
| 922 |
+
model_kwargs = _expand_dict_for_generation(model_kwargs)
|
| 923 |
+
|
| 924 |
+
if is_encoder_decoder:
|
| 925 |
+
if model_kwargs.get("encoder_outputs") is None:
|
| 926 |
+
raise ValueError("If `is_encoder_decoder` is True, make sure that `encoder_outputs` is defined.")
|
| 927 |
+
model_kwargs["encoder_outputs"] = _expand_dict_for_generation(model_kwargs["encoder_outputs"])
|
| 928 |
+
|
| 929 |
+
return input_ids, model_kwargs
|
| 930 |
+
|
| 931 |
+
|
| 932 |
+
__all__ = ["Qwen2_5_VLForConditionalGeneration", "Qwen2_5_VLModel", "Qwen2_5_VLPreTrainedModel"]
|
| 933 |
+
|
| 934 |
+
|
| 935 |
+
|
| 936 |
+
class Qwen2_5_VLVideosProcessorKwargs(VideosKwargs, total=False):
|
| 937 |
+
fps: Union[List[float], float]
|
| 938 |
+
|
| 939 |
+
|
| 940 |
+
class Qwen2_5_VLProcessorKwargs(ProcessingKwargs, total=False):
|
| 941 |
+
videos_kwargs: Qwen2_5_VLVideosProcessorKwargs
|
| 942 |
+
_defaults = {
|
| 943 |
+
"text_kwargs": {
|
| 944 |
+
"padding": False,
|
| 945 |
+
},
|
| 946 |
+
"videos_kwargs": {"fps": 2.0},
|
| 947 |
+
}
|
| 948 |
+
|
| 949 |
+
|
| 950 |
+
class Qwen2_5_VLProcessor(ProcessorMixin):
|
| 951 |
+
r"""
|
| 952 |
+
Constructs a Qwen2.5-VL processor which wraps a Qwen2.5-VL image processor and a Qwen2 tokenizer into a single processor.
|
| 953 |
+
[`Qwen2_5_VLProcessor`] offers all the functionalities of [`Qwen2VLImageProcessor`] and [`Qwen2TokenizerFast`]. See the
|
| 954 |
+
[`~Qwen2_5_VLProcessor.__call__`] and [`~Qwen2_5_VLProcessor.decode`] for more information.
|
| 955 |
+
Args:
|
| 956 |
+
image_processor ([`Qwen2VLImageProcessor`], *optional*):
|
| 957 |
+
The image processor is a required input.
|
| 958 |
+
tokenizer ([`Qwen2TokenizerFast`], *optional*):
|
| 959 |
+
The tokenizer is a required input.
|
| 960 |
+
chat_template (`str`, *optional*): A Jinja template which will be used to convert lists of messages
|
| 961 |
+
in a chat into a tokenizable string.
|
| 962 |
+
"""
|
| 963 |
+
|
| 964 |
+
attributes = ["image_processor", "tokenizer"]
|
| 965 |
+
valid_kwargs = ["chat_template"]
|
| 966 |
+
|
| 967 |
+
image_processor_class = "AutoImageProcessor"
|
| 968 |
+
tokenizer_class = ("Qwen2Tokenizer", "Qwen2TokenizerFast")
|
| 969 |
+
|
| 970 |
+
def __init__(self, image_processor=None, tokenizer=None, chat_template=None, **kwargs):
|
| 971 |
+
self.image_token = "<|image_pad|>" if not hasattr(tokenizer, "image_token") else tokenizer.image_token
|
| 972 |
+
self.video_token = "<|video_pad|>" if not hasattr(tokenizer, "video_token") else tokenizer.video_token
|
| 973 |
+
super().__init__(image_processor, tokenizer, chat_template=chat_template)
|
| 974 |
+
|
| 975 |
+
def __call__(
|
| 976 |
+
self,
|
| 977 |
+
images: ImageInput = None,
|
| 978 |
+
text: Union[TextInput, PreTokenizedInput, List[TextInput], List[PreTokenizedInput]] = None,
|
| 979 |
+
videos: VideoInput = None,
|
| 980 |
+
**kwargs: Unpack[Qwen2_5_VLProcessorKwargs],
|
| 981 |
+
) -> BatchFeature:
|
| 982 |
+
"""
|
| 983 |
+
Main method to prepare for the model one or several sequences(s) and image(s). This method forwards the `text`
|
| 984 |
+
and `kwargs` arguments to Qwen2TokenizerFast's [`~Qwen2TokenizerFast.__call__`] if `text` is not `None` to encode
|
| 985 |
+
the text. To prepare the vision inputs, this method forwards the `vision_infos` and `kwrags` arguments to
|
| 986 |
+
Qwen2VLImageProcessor's [`~Qwen2VLImageProcessor.__call__`] if `vision_infos` is not `None`.
|
| 987 |
+
|
| 988 |
+
Args:
|
| 989 |
+
images (`PIL.Image.Image`, `np.ndarray`, `torch.Tensor`, `List[PIL.Image.Image]`, `List[np.ndarray]`, `List[torch.Tensor]`):
|
| 990 |
+
The image or batch of images to be prepared. Each image can be a PIL image, NumPy array or PyTorch
|
| 991 |
+
tensor. Both channels-first and channels-last formats are supported.
|
| 992 |
+
text (`str`, `List[str]`, `List[List[str]]`):
|
| 993 |
+
The sequence or batch of sequences to be encoded. Each sequence can be a string or a list of strings
|
| 994 |
+
(pretokenized string). If the sequences are provided as list of strings (pretokenized), you must set
|
| 995 |
+
`is_split_into_words=True` (to lift the ambiguity with a batch of sequences).
|
| 996 |
+
videos (`np.ndarray`, `torch.Tensor`, `List[np.ndarray]`, `List[torch.Tensor]`):
|
| 997 |
+
The image or batch of videos to be prepared. Each video can be a 4D NumPy array or PyTorch
|
| 998 |
+
tensor, or a nested list of 3D frames. Both channels-first and channels-last formats are supported.
|
| 999 |
+
return_tensors (`str` or [`~utils.TensorType`], *optional*):
|
| 1000 |
+
If set, will return tensors of a particular framework. Acceptable values are:
|
| 1001 |
+
- `'tf'`: Return TensorFlow `tf.constant` objects.
|
| 1002 |
+
- `'pt'`: Return PyTorch `torch.Tensor` objects.
|
| 1003 |
+
- `'np'`: Return NumPy `np.ndarray` objects.
|
| 1004 |
+
- `'jax'`: Return JAX `jnp.ndarray` objects.
|
| 1005 |
+
|
| 1006 |
+
Returns:
|
| 1007 |
+
[`BatchFeature`]: A [`BatchFeature`] with the following fields:
|
| 1008 |
+
|
| 1009 |
+
- **input_ids** -- List of token ids to be fed to a model. Returned when `text` is not `None`.
|
| 1010 |
+
- **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
|
| 1011 |
+
`return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not
|
| 1012 |
+
`None`).
|
| 1013 |
+
- **pixel_values** -- Pixel values to be fed to a model. Returned when `images` is not `None`.
|
| 1014 |
+
- **pixel_values_videos** -- Pixel values of videos to be fed to a model. Returned when `videos` is not `None`.
|
| 1015 |
+
- **image_grid_thw** -- List of image 3D grid in LLM. Returned when `images` is not `None`.
|
| 1016 |
+
- **video_grid_thw** -- List of video 3D grid in LLM. Returned when `videos` is not `None`.
|
| 1017 |
+
- **second_per_grid_ts** -- List of video seconds per time grid. Returned when `videos` is not `None`.
|
| 1018 |
+
"""
|
| 1019 |
+
output_kwargs = self._merge_kwargs(
|
| 1020 |
+
Qwen2_5_VLProcessorKwargs,
|
| 1021 |
+
tokenizer_init_kwargs=self.tokenizer.init_kwargs,
|
| 1022 |
+
**kwargs,
|
| 1023 |
+
)
|
| 1024 |
+
if images is not None:
|
| 1025 |
+
image_inputs = self.image_processor(images=images, videos=None, **output_kwargs["images_kwargs"])
|
| 1026 |
+
image_grid_thw = image_inputs["image_grid_thw"]
|
| 1027 |
+
else:
|
| 1028 |
+
image_inputs = {}
|
| 1029 |
+
image_grid_thw = None
|
| 1030 |
+
|
| 1031 |
+
if videos is not None:
|
| 1032 |
+
videos_inputs = self.image_processor(images=None, videos=videos, **output_kwargs["images_kwargs"])
|
| 1033 |
+
video_grid_thw = videos_inputs["video_grid_thw"]
|
| 1034 |
+
|
| 1035 |
+
fps = output_kwargs["videos_kwargs"].pop("fps", 2.0)
|
| 1036 |
+
if isinstance(fps, (int, float)):
|
| 1037 |
+
second_per_grid_ts = [self.image_processor.temporal_patch_size / fps] * len(video_grid_thw)
|
| 1038 |
+
elif hasattr(fps, "__len__") and len(fps) == len(video_grid_thw):
|
| 1039 |
+
second_per_grid_ts = [self.image_processor.temporal_patch_size / tmp for tmp in fps]
|
| 1040 |
+
else:
|
| 1041 |
+
raise ValueError(
|
| 1042 |
+
f"The length of fps ({len(fps) if hasattr(fps, '__len__') else fps}) must be equal to the length of video_grid_thw ({len(video_grid_thw)}) or fps should be a single number."
|
| 1043 |
+
)
|
| 1044 |
+
videos_inputs.update({"second_per_grid_ts": second_per_grid_ts})
|
| 1045 |
+
|
| 1046 |
+
else:
|
| 1047 |
+
videos_inputs = {}
|
| 1048 |
+
video_grid_thw = None
|
| 1049 |
+
|
| 1050 |
+
if not isinstance(text, list):
|
| 1051 |
+
text = [text]
|
| 1052 |
+
|
| 1053 |
+
if image_grid_thw is not None:
|
| 1054 |
+
merge_length = self.image_processor.merge_size**2
|
| 1055 |
+
index = 0
|
| 1056 |
+
for i in range(len(text)):
|
| 1057 |
+
while self.image_token in text[i]:
|
| 1058 |
+
text[i] = text[i].replace(
|
| 1059 |
+
self.image_token,
|
| 1060 |
+
"<|placeholder|>" * (image_grid_thw[index].prod() // merge_length),
|
| 1061 |
+
1,
|
| 1062 |
+
)
|
| 1063 |
+
index += 1
|
| 1064 |
+
text[i] = text[i].replace("<|placeholder|>", self.image_token)
|
| 1065 |
+
|
| 1066 |
+
if video_grid_thw is not None:
|
| 1067 |
+
merge_length = self.image_processor.merge_size**2
|
| 1068 |
+
index = 0
|
| 1069 |
+
for i in range(len(text)):
|
| 1070 |
+
while self.video_token in text[i]:
|
| 1071 |
+
text[i] = text[i].replace(
|
| 1072 |
+
self.video_token,
|
| 1073 |
+
"<|placeholder|>" * (video_grid_thw[index].prod() // merge_length),
|
| 1074 |
+
1,
|
| 1075 |
+
)
|
| 1076 |
+
index += 1
|
| 1077 |
+
text[i] = text[i].replace("<|placeholder|>", self.video_token)
|
| 1078 |
+
|
| 1079 |
+
text_inputs = self.tokenizer(text, **output_kwargs["text_kwargs"])
|
| 1080 |
+
|
| 1081 |
+
return BatchFeature(data={**text_inputs, **image_inputs, **videos_inputs})
|
| 1082 |
+
|
| 1083 |
+
def batch_decode(self, *args, **kwargs):
|
| 1084 |
+
"""
|
| 1085 |
+
This method forwards all its arguments to Qwen2TokenizerFast's [`~PreTrainedTokenizer.batch_decode`]. Please
|
| 1086 |
+
refer to the docstring of this method for more information.
|
| 1087 |
+
"""
|
| 1088 |
+
return self.tokenizer.batch_decode(*args, **kwargs)
|
| 1089 |
+
|
| 1090 |
+
def batch_decode_all2all(self, *args, **kwargs):
|
| 1091 |
+
"""
|
| 1092 |
+
This method forwards all its arguments to Qwen2TokenizerFast's [`~PreTrainedTokenizer.batch_decode`]. Please
|
| 1093 |
+
refer to the docstring of this method for more information.
|
| 1094 |
+
"""
|
| 1095 |
+
decoded = self.tokenizer.batch_decode(*args, **kwargs)
|
| 1096 |
+
pattern = r'<\|vision_start\|>.*?<\|vision_end\|>'
|
| 1097 |
+
decoded_with_image_tag = [re.sub(pattern, '<image>', d, flags=re.DOTALL) for d in decoded]
|
| 1098 |
+
decoded_with_image_tag = [re.sub(r'<\|im_end\|>', '', d) for d in decoded_with_image_tag]
|
| 1099 |
+
return decoded_with_image_tag
|
| 1100 |
+
|
| 1101 |
+
def decode(self, *args, **kwargs):
|
| 1102 |
+
"""
|
| 1103 |
+
This method forwards all its arguments to Qwen2TokenizerFast's [`~PreTrainedTokenizer.decode`]. Please refer to
|
| 1104 |
+
the docstring of this method for more information.
|
| 1105 |
+
"""
|
| 1106 |
+
return self.tokenizer.decode(*args, **kwargs)
|
| 1107 |
+
|
| 1108 |
+
def post_process_image_text_to_text(
|
| 1109 |
+
self, generated_outputs, skip_special_tokens=True, clean_up_tokenization_spaces=False, **kwargs
|
| 1110 |
+
):
|
| 1111 |
+
"""
|
| 1112 |
+
Post-process the output of the model to decode the text.
|
| 1113 |
+
|
| 1114 |
+
Args:
|
| 1115 |
+
generated_outputs (`torch.Tensor` or `np.ndarray`):
|
| 1116 |
+
The output of the model `generate` function. The output is expected to be a tensor of shape `(batch_size, sequence_length)`
|
| 1117 |
+
or `(sequence_length,)`.
|
| 1118 |
+
skip_special_tokens (`bool`, *optional*, defaults to `True`):
|
| 1119 |
+
Whether or not to remove special tokens in the output. Argument passed to the tokenizer's `batch_decode` method.
|
| 1120 |
+
Clean_up_tokenization_spaces (`bool`, *optional*, defaults to `False`):
|
| 1121 |
+
Whether or not to clean up the tokenization spaces. Argument passed to the tokenizer's `batch_decode` method.
|
| 1122 |
+
**kwargs:
|
| 1123 |
+
Additional arguments to be passed to the tokenizer's `batch_decode method`.
|
| 1124 |
+
|
| 1125 |
+
Returns:
|
| 1126 |
+
`List[str]`: The decoded text.
|
| 1127 |
+
"""
|
| 1128 |
+
return self.tokenizer.batch_decode(
|
| 1129 |
+
generated_outputs,
|
| 1130 |
+
skip_special_tokens=skip_special_tokens,
|
| 1131 |
+
clean_up_tokenization_spaces=clean_up_tokenization_spaces,
|
| 1132 |
+
**kwargs,
|
| 1133 |
+
)
|
| 1134 |
+
|
| 1135 |
+
@property
|
| 1136 |
+
def model_input_names(self):
|
| 1137 |
+
tokenizer_input_names = self.tokenizer.model_input_names
|
| 1138 |
+
image_processor_input_names = self.image_processor.model_input_names
|
| 1139 |
+
names_from_processor = list(dict.fromkeys(tokenizer_input_names + image_processor_input_names))
|
| 1140 |
+
return names_from_processor + ["second_per_grid_ts"]
|
| 1141 |
+
|
| 1142 |
+
|
| 1143 |
+
__all__ = ["Qwen2_5_VLProcessor"]
|
diffsynth/models/nexus_gen_projector.py
ADDED
|
@@ -0,0 +1,417 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import math
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
from typing import Optional, Tuple
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def rotate_half(x):
|
| 9 |
+
"""Rotates half the hidden dims of the input."""
|
| 10 |
+
x1 = x[..., : x.shape[-1] // 2]
|
| 11 |
+
x2 = x[..., x.shape[-1] // 2 :]
|
| 12 |
+
return torch.cat((-x2, x1), dim=-1)
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def apply_multimodal_rotary_pos_emb(q, k, cos, sin, mrope_section, unsqueeze_dim=1):
|
| 16 |
+
mrope_section = mrope_section * 2
|
| 17 |
+
cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
|
| 18 |
+
unsqueeze_dim
|
| 19 |
+
)
|
| 20 |
+
sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
|
| 21 |
+
unsqueeze_dim
|
| 22 |
+
)
|
| 23 |
+
|
| 24 |
+
q_embed = (q * cos) + (rotate_half(q) * sin)
|
| 25 |
+
k_embed = (k * cos) + (rotate_half(k) * sin)
|
| 26 |
+
return q_embed, k_embed
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
| 30 |
+
def __init__(self, config, device=None):
|
| 31 |
+
super().__init__()
|
| 32 |
+
# BC: "rope_type" was originally "type"
|
| 33 |
+
if hasattr(config, "rope_scaling") and config.rope_scaling is not None:
|
| 34 |
+
self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type"))
|
| 35 |
+
else:
|
| 36 |
+
self.rope_type = "default"
|
| 37 |
+
self.max_seq_len_cached = config.max_position_embeddings
|
| 38 |
+
self.original_max_seq_len = config.max_position_embeddings
|
| 39 |
+
|
| 40 |
+
self.config = config
|
| 41 |
+
from transformers.modeling_rope_utils import _compute_default_rope_parameters
|
| 42 |
+
self.rope_init_fn = _compute_default_rope_parameters
|
| 43 |
+
|
| 44 |
+
inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
|
| 45 |
+
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
| 46 |
+
self.original_inv_freq = self.inv_freq
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def _dynamic_frequency_update(self, position_ids, device):
|
| 50 |
+
"""
|
| 51 |
+
dynamic RoPE layers should recompute `inv_freq` in the following situations:
|
| 52 |
+
1 - growing beyond the cached sequence length (allow scaling)
|
| 53 |
+
2 - the current sequence length is in the original scale (avoid losing precision with small sequences)
|
| 54 |
+
"""
|
| 55 |
+
seq_len = torch.max(position_ids) + 1
|
| 56 |
+
if seq_len > self.max_seq_len_cached: # growth
|
| 57 |
+
inv_freq, self.attention_scaling = self.rope_init_fn(
|
| 58 |
+
self.config, device, seq_len=seq_len, **self.rope_kwargs
|
| 59 |
+
)
|
| 60 |
+
self.register_buffer("inv_freq", inv_freq, persistent=False) # TODO joao: may break with compilation
|
| 61 |
+
self.max_seq_len_cached = seq_len
|
| 62 |
+
|
| 63 |
+
if seq_len < self.original_max_seq_len and self.max_seq_len_cached > self.original_max_seq_len: # reset
|
| 64 |
+
self.register_buffer("inv_freq", self.original_inv_freq, persistent=False)
|
| 65 |
+
self.max_seq_len_cached = self.original_max_seq_len
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
@torch.no_grad()
|
| 69 |
+
def forward(self, x, position_ids):
|
| 70 |
+
if "dynamic" in self.rope_type:
|
| 71 |
+
self._dynamic_frequency_update(position_ids, device=x.device)
|
| 72 |
+
|
| 73 |
+
# Core RoPE block. In contrast to other models, Qwen2_5_VL has different position ids for the grids
|
| 74 |
+
# So we expand the inv_freq to shape (3, ...)
|
| 75 |
+
inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)
|
| 76 |
+
position_ids_expanded = position_ids[:, :, None, :].float() # shape (3, bs, 1, positions)
|
| 77 |
+
# Force float32 (see https://github.com/huggingface/transformers/pull/29285)
|
| 78 |
+
device_type = x.device.type
|
| 79 |
+
device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu"
|
| 80 |
+
with torch.autocast(device_type=device_type, enabled=False):
|
| 81 |
+
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)
|
| 82 |
+
emb = torch.cat((freqs, freqs), dim=-1)
|
| 83 |
+
cos = emb.cos()
|
| 84 |
+
sin = emb.sin()
|
| 85 |
+
|
| 86 |
+
# Advanced RoPE types (e.g. yarn) apply a post-processing scaling factor, equivalent to scaling attention
|
| 87 |
+
cos = cos * self.attention_scaling
|
| 88 |
+
sin = sin * self.attention_scaling
|
| 89 |
+
|
| 90 |
+
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
|
| 94 |
+
"""
|
| 95 |
+
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
|
| 96 |
+
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
|
| 97 |
+
"""
|
| 98 |
+
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
|
| 99 |
+
if n_rep == 1:
|
| 100 |
+
return hidden_states
|
| 101 |
+
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
|
| 102 |
+
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
class Qwen2_5_VLAttention(nn.Module):
|
| 106 |
+
def __init__(self, config, layer_idx: Optional[int] = None):
|
| 107 |
+
super().__init__()
|
| 108 |
+
self.config = config
|
| 109 |
+
self.layer_idx = layer_idx
|
| 110 |
+
|
| 111 |
+
self.hidden_size = config.hidden_size
|
| 112 |
+
self.num_heads = config.num_attention_heads
|
| 113 |
+
self.head_dim = self.hidden_size // self.num_heads
|
| 114 |
+
self.num_key_value_heads = config.num_key_value_heads
|
| 115 |
+
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
|
| 116 |
+
self.is_causal = True
|
| 117 |
+
self.attention_dropout = config.attention_dropout
|
| 118 |
+
self.rope_scaling = config.rope_scaling
|
| 119 |
+
|
| 120 |
+
if (self.head_dim * self.num_heads) != self.hidden_size:
|
| 121 |
+
raise ValueError(
|
| 122 |
+
f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}"
|
| 123 |
+
f" and `num_heads`: {self.num_heads})."
|
| 124 |
+
)
|
| 125 |
+
self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=True)
|
| 126 |
+
self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True)
|
| 127 |
+
self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True)
|
| 128 |
+
self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def forward(
|
| 132 |
+
self,
|
| 133 |
+
hidden_states: torch.Tensor,
|
| 134 |
+
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
|
| 135 |
+
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
|
| 136 |
+
bsz, q_len, _ = hidden_states.size()
|
| 137 |
+
|
| 138 |
+
query_states = self.q_proj(hidden_states)
|
| 139 |
+
key_states = self.k_proj(hidden_states)
|
| 140 |
+
value_states = self.v_proj(hidden_states)
|
| 141 |
+
|
| 142 |
+
query_states = query_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)
|
| 143 |
+
key_states = key_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)
|
| 144 |
+
value_states = value_states.view(bsz, q_len, -1, self.head_dim).transpose(1, 2)
|
| 145 |
+
|
| 146 |
+
cos, sin = position_embeddings
|
| 147 |
+
query_states, key_states = apply_multimodal_rotary_pos_emb(
|
| 148 |
+
query_states, key_states, cos, sin, self.rope_scaling["mrope_section"]
|
| 149 |
+
)
|
| 150 |
+
|
| 151 |
+
# repeat k/v heads if n_kv_heads < n_heads
|
| 152 |
+
key_states = repeat_kv(key_states, self.num_key_value_groups)
|
| 153 |
+
value_states = repeat_kv(value_states, self.num_key_value_groups)
|
| 154 |
+
|
| 155 |
+
attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)
|
| 156 |
+
|
| 157 |
+
# Fix precision issues in Qwen2-VL float16 inference
|
| 158 |
+
# Replace inf values with zeros in attention weights to prevent NaN propagation
|
| 159 |
+
if query_states.dtype == torch.float16:
|
| 160 |
+
attn_weights = torch.where(torch.isinf(attn_weights), torch.zeros_like(attn_weights), attn_weights)
|
| 161 |
+
|
| 162 |
+
# upcast attention to fp32
|
| 163 |
+
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
|
| 164 |
+
attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)
|
| 165 |
+
attn_output = torch.matmul(attn_weights, value_states)
|
| 166 |
+
|
| 167 |
+
if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):
|
| 168 |
+
raise ValueError(
|
| 169 |
+
f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is"
|
| 170 |
+
f" {attn_output.size()}"
|
| 171 |
+
)
|
| 172 |
+
|
| 173 |
+
attn_output = attn_output.transpose(1, 2).contiguous()
|
| 174 |
+
attn_output = attn_output.reshape(bsz, q_len, -1)
|
| 175 |
+
|
| 176 |
+
attn_output = self.o_proj(attn_output)
|
| 177 |
+
|
| 178 |
+
return attn_output
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
class Qwen2MLP(nn.Module):
|
| 182 |
+
def __init__(self, config):
|
| 183 |
+
super().__init__()
|
| 184 |
+
from transformers.activations import ACT2FN
|
| 185 |
+
self.config = config
|
| 186 |
+
self.hidden_size = config.hidden_size
|
| 187 |
+
self.intermediate_size = config.intermediate_size
|
| 188 |
+
self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
| 189 |
+
self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
|
| 190 |
+
self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
|
| 191 |
+
self.act_fn = ACT2FN[config.hidden_act]
|
| 192 |
+
|
| 193 |
+
def forward(self, x):
|
| 194 |
+
down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
|
| 195 |
+
return down_proj
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
class Qwen2RMSNorm(nn.Module):
|
| 199 |
+
def __init__(self, hidden_size, eps=1e-6):
|
| 200 |
+
"""
|
| 201 |
+
Qwen2RMSNorm is equivalent to T5LayerNorm
|
| 202 |
+
"""
|
| 203 |
+
super().__init__()
|
| 204 |
+
self.weight = nn.Parameter(torch.ones(hidden_size))
|
| 205 |
+
self.variance_epsilon = eps
|
| 206 |
+
|
| 207 |
+
def forward(self, hidden_states):
|
| 208 |
+
input_dtype = hidden_states.dtype
|
| 209 |
+
hidden_states = hidden_states.to(torch.float32)
|
| 210 |
+
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
| 211 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
|
| 212 |
+
return self.weight * hidden_states.to(input_dtype)
|
| 213 |
+
|
| 214 |
+
def extra_repr(self):
|
| 215 |
+
return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
|
| 216 |
+
|
| 217 |
+
|
| 218 |
+
class Qwen2_5_VLDecoderLayer(nn.Module):
|
| 219 |
+
def __init__(self, config, layer_idx):
|
| 220 |
+
super().__init__()
|
| 221 |
+
self.hidden_size = config.hidden_size
|
| 222 |
+
|
| 223 |
+
self.self_attn = Qwen2_5_VLAttention(config, layer_idx)
|
| 224 |
+
|
| 225 |
+
self.mlp = Qwen2MLP(config)
|
| 226 |
+
self.input_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
| 227 |
+
self.post_attention_layernorm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
| 228 |
+
|
| 229 |
+
def forward(
|
| 230 |
+
self,
|
| 231 |
+
hidden_states: torch.Tensor,
|
| 232 |
+
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
|
| 233 |
+
) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
|
| 234 |
+
|
| 235 |
+
residual = hidden_states
|
| 236 |
+
|
| 237 |
+
hidden_states = self.input_layernorm(hidden_states)
|
| 238 |
+
|
| 239 |
+
# Self Attention
|
| 240 |
+
hidden_states = self.self_attn(
|
| 241 |
+
hidden_states=hidden_states,
|
| 242 |
+
position_embeddings=position_embeddings,
|
| 243 |
+
)
|
| 244 |
+
hidden_states = residual + hidden_states
|
| 245 |
+
|
| 246 |
+
# Fully Connected
|
| 247 |
+
residual = hidden_states
|
| 248 |
+
hidden_states = self.post_attention_layernorm(hidden_states)
|
| 249 |
+
hidden_states = self.mlp(hidden_states)
|
| 250 |
+
hidden_states = residual + hidden_states
|
| 251 |
+
|
| 252 |
+
return hidden_states
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
class NexusGenImageEmbeddingMerger(nn.Module):
|
| 256 |
+
def __init__(self, num_layers=1, out_channel=4096, expand_ratio=4, device='cpu'):
|
| 257 |
+
super().__init__()
|
| 258 |
+
from transformers import Qwen2_5_VLConfig
|
| 259 |
+
from transformers.activations import ACT2FN
|
| 260 |
+
config = Qwen2_5_VLConfig(**{
|
| 261 |
+
"_name_or_path": "DiffSynth-Studio/Nexus-GenV2",
|
| 262 |
+
"architectures": [
|
| 263 |
+
"Qwen2_5_VLForConditionalGeneration"
|
| 264 |
+
],
|
| 265 |
+
"attention_dropout": 0.0,
|
| 266 |
+
"auto_map": {
|
| 267 |
+
"AutoConfig": "configuration_qwen2_5_vl.Qwen2_5_VLConfig",
|
| 268 |
+
"AutoModel": "modeling_qwen2_5_vl.Qwen2_5_VLModel",
|
| 269 |
+
"AutoModelForCausalLM": "modeling_qwen2_5_vl.Qwen2_5_VLForConditionalGeneration"
|
| 270 |
+
},
|
| 271 |
+
"bos_token_id": 151643,
|
| 272 |
+
"eos_token_id": 151645,
|
| 273 |
+
"hidden_act": "silu",
|
| 274 |
+
"hidden_size": 3584,
|
| 275 |
+
"image_token_id": 151655,
|
| 276 |
+
"initializer_range": 0.02,
|
| 277 |
+
"intermediate_size": 18944,
|
| 278 |
+
"max_position_embeddings": 128000,
|
| 279 |
+
"max_window_layers": 28,
|
| 280 |
+
"model_type": "qwen2_5_vl",
|
| 281 |
+
"num_attention_heads": 28,
|
| 282 |
+
"num_hidden_layers": 28,
|
| 283 |
+
"num_key_value_heads": 4,
|
| 284 |
+
"pad_token_id": 151643,
|
| 285 |
+
"rms_norm_eps": 1e-06,
|
| 286 |
+
"rope_scaling": {
|
| 287 |
+
"mrope_section": [
|
| 288 |
+
16,
|
| 289 |
+
24,
|
| 290 |
+
24
|
| 291 |
+
],
|
| 292 |
+
"rope_type": "default",
|
| 293 |
+
"type": "default"
|
| 294 |
+
},
|
| 295 |
+
"rope_theta": 1000000.0,
|
| 296 |
+
"sliding_window": 32768,
|
| 297 |
+
"tie_word_embeddings": False,
|
| 298 |
+
"torch_dtype": "bfloat16",
|
| 299 |
+
"transformers_version": "4.49.0",
|
| 300 |
+
"use_cache": False,
|
| 301 |
+
"use_sliding_window": False,
|
| 302 |
+
"video_token_id": 151656,
|
| 303 |
+
"vision_config": {
|
| 304 |
+
"hidden_size": 1280,
|
| 305 |
+
"in_chans": 3,
|
| 306 |
+
"model_type": "qwen2_5_vl",
|
| 307 |
+
"spatial_patch_size": 14,
|
| 308 |
+
"tokens_per_second": 2,
|
| 309 |
+
"torch_dtype": "bfloat16"
|
| 310 |
+
},
|
| 311 |
+
"vision_end_token_id": 151653,
|
| 312 |
+
"vision_start_token_id": 151652,
|
| 313 |
+
"vision_token_id": 151654,
|
| 314 |
+
"vocab_size": 152064
|
| 315 |
+
})
|
| 316 |
+
self.config = config
|
| 317 |
+
self.num_layers = num_layers
|
| 318 |
+
self.layers = nn.ModuleList([Qwen2_5_VLDecoderLayer(config, layer_idx) for layer_idx in range(num_layers)])
|
| 319 |
+
self.projector = nn.Sequential(Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps),
|
| 320 |
+
nn.Linear(config.hidden_size, out_channel * expand_ratio),
|
| 321 |
+
Qwen2RMSNorm(out_channel * expand_ratio, eps=config.rms_norm_eps),
|
| 322 |
+
ACT2FN[config.hidden_act], nn.Linear(out_channel * expand_ratio, out_channel),
|
| 323 |
+
Qwen2RMSNorm(out_channel, eps=config.rms_norm_eps))
|
| 324 |
+
self.base_grid = torch.tensor([[1, 72, 72]], device=device)
|
| 325 |
+
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config=config, device=device)
|
| 326 |
+
|
| 327 |
+
def get_position_ids(self, image_grid_thw):
|
| 328 |
+
"""
|
| 329 |
+
Generates position ids for the input embeddings grid.
|
| 330 |
+
modified from the qwen2_vl mrope.
|
| 331 |
+
"""
|
| 332 |
+
batch_size = image_grid_thw.shape[0]
|
| 333 |
+
spatial_merge_size = self.config.vision_config.spatial_merge_size
|
| 334 |
+
t, h, w = (
|
| 335 |
+
image_grid_thw[0][0],
|
| 336 |
+
image_grid_thw[0][1],
|
| 337 |
+
image_grid_thw[0][2],
|
| 338 |
+
)
|
| 339 |
+
llm_grid_t, llm_grid_h, llm_grid_w = (
|
| 340 |
+
t.item(),
|
| 341 |
+
h.item() // spatial_merge_size,
|
| 342 |
+
w.item() // spatial_merge_size,
|
| 343 |
+
)
|
| 344 |
+
scale_h = self.base_grid[0][1].item() / h.item()
|
| 345 |
+
scale_w = self.base_grid[0][2].item() / w.item()
|
| 346 |
+
|
| 347 |
+
range_tensor = torch.arange(llm_grid_t).view(-1, 1)
|
| 348 |
+
expanded_range = range_tensor.expand(-1, llm_grid_h * llm_grid_w)
|
| 349 |
+
time_tensor = expanded_range * self.config.vision_config.tokens_per_second
|
| 350 |
+
t_index = time_tensor.long().flatten().to(image_grid_thw.device)
|
| 351 |
+
h_index = torch.arange(llm_grid_h).view(1, -1, 1).expand(llm_grid_t, -1, llm_grid_w).flatten().to(image_grid_thw.device) * scale_h
|
| 352 |
+
w_index = torch.arange(llm_grid_w).view(1, 1, -1).expand(llm_grid_t, llm_grid_h, -1).flatten().to(image_grid_thw.device) * scale_w
|
| 353 |
+
# 3, B, L
|
| 354 |
+
position_ids = torch.stack([t_index, h_index, w_index]).unsqueeze(0).repeat(batch_size, 1, 1).permute(1, 0, 2)
|
| 355 |
+
return position_ids
|
| 356 |
+
|
| 357 |
+
def forward(self, embeds, embeds_grid, ref_embeds=None, ref_embeds_grid=None):
|
| 358 |
+
position_ids = self.get_position_ids(embeds_grid)
|
| 359 |
+
hidden_states = embeds
|
| 360 |
+
if ref_embeds is not None:
|
| 361 |
+
position_ids_ref_embeds = self.get_position_ids(ref_embeds_grid)
|
| 362 |
+
position_ids = torch.cat((position_ids, position_ids_ref_embeds), dim=-1)
|
| 363 |
+
hidden_states = torch.cat((embeds, ref_embeds), dim=1)
|
| 364 |
+
|
| 365 |
+
position_embeddings = self.rotary_emb(hidden_states, position_ids)
|
| 366 |
+
for layer in self.layers:
|
| 367 |
+
hidden_states = layer(hidden_states, position_embeddings)
|
| 368 |
+
|
| 369 |
+
hidden_states = self.projector(hidden_states)
|
| 370 |
+
return hidden_states
|
| 371 |
+
|
| 372 |
+
@staticmethod
|
| 373 |
+
def state_dict_converter():
|
| 374 |
+
return NexusGenMergerStateDictConverter()
|
| 375 |
+
|
| 376 |
+
|
| 377 |
+
class NexusGenMergerStateDictConverter:
|
| 378 |
+
def __init__(self):
|
| 379 |
+
pass
|
| 380 |
+
|
| 381 |
+
def from_diffusers(self, state_dict):
|
| 382 |
+
return state_dict
|
| 383 |
+
|
| 384 |
+
def from_civitai(self, state_dict):
|
| 385 |
+
merger_state_dict = {key.replace("embedding_merger.", ""): value for key, value in state_dict.items() if key.startswith('embedding_merger.')}
|
| 386 |
+
return merger_state_dict
|
| 387 |
+
|
| 388 |
+
|
| 389 |
+
class NexusGenAdapter(nn.Module):
|
| 390 |
+
"""
|
| 391 |
+
Adapter for Nexus-Gen generation decoder.
|
| 392 |
+
"""
|
| 393 |
+
def __init__(self, input_dim=3584, output_dim=4096):
|
| 394 |
+
super(NexusGenAdapter, self).__init__()
|
| 395 |
+
self.adapter = nn.Sequential(nn.Linear(input_dim, output_dim),
|
| 396 |
+
nn.LayerNorm(output_dim), nn.ReLU(),
|
| 397 |
+
nn.Linear(output_dim, output_dim),
|
| 398 |
+
nn.LayerNorm(output_dim))
|
| 399 |
+
|
| 400 |
+
def forward(self, x):
|
| 401 |
+
return self.adapter(x)
|
| 402 |
+
|
| 403 |
+
@staticmethod
|
| 404 |
+
def state_dict_converter():
|
| 405 |
+
return NexusGenAdapterStateDictConverter()
|
| 406 |
+
|
| 407 |
+
|
| 408 |
+
class NexusGenAdapterStateDictConverter:
|
| 409 |
+
def __init__(self):
|
| 410 |
+
pass
|
| 411 |
+
|
| 412 |
+
def from_diffusers(self, state_dict):
|
| 413 |
+
return state_dict
|
| 414 |
+
|
| 415 |
+
def from_civitai(self, state_dict):
|
| 416 |
+
adapter_state_dict = {key: value for key, value in state_dict.items() if key.startswith('adapter.')}
|
| 417 |
+
return adapter_state_dict
|
diffsynth/models/omnigen.py
ADDED
|
@@ -0,0 +1,803 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# The code is revised from DiT
|
| 2 |
+
import os
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
import numpy as np
|
| 6 |
+
import math
|
| 7 |
+
from safetensors.torch import load_file
|
| 8 |
+
from typing import List, Optional, Tuple, Union
|
| 9 |
+
import torch.utils.checkpoint
|
| 10 |
+
from huggingface_hub import snapshot_download
|
| 11 |
+
from transformers.modeling_outputs import BaseModelOutputWithPast
|
| 12 |
+
from transformers import Phi3Config, Phi3Model
|
| 13 |
+
from transformers.cache_utils import Cache, DynamicCache
|
| 14 |
+
from transformers.utils import logging
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
logger = logging.get_logger(__name__)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class Phi3Transformer(Phi3Model):
|
| 21 |
+
"""
|
| 22 |
+
Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`Phi3DecoderLayer`]
|
| 23 |
+
We only modified the attention mask
|
| 24 |
+
Args:
|
| 25 |
+
config: Phi3Config
|
| 26 |
+
"""
|
| 27 |
+
def prefetch_layer(self, layer_idx: int, device: torch.device):
|
| 28 |
+
"Starts prefetching the next layer cache"
|
| 29 |
+
with torch.cuda.stream(self.prefetch_stream):
|
| 30 |
+
# Prefetch next layer tensors to GPU
|
| 31 |
+
for name, param in self.layers[layer_idx].named_parameters():
|
| 32 |
+
param.data = param.data.to(device, non_blocking=True)
|
| 33 |
+
|
| 34 |
+
def evict_previous_layer(self, layer_idx: int):
|
| 35 |
+
"Moves the previous layer cache to the CPU"
|
| 36 |
+
prev_layer_idx = layer_idx - 1
|
| 37 |
+
for name, param in self.layers[prev_layer_idx].named_parameters():
|
| 38 |
+
param.data = param.data.to("cpu", non_blocking=True)
|
| 39 |
+
|
| 40 |
+
def get_offlaod_layer(self, layer_idx: int, device: torch.device):
|
| 41 |
+
# init stream
|
| 42 |
+
if not hasattr(self, "prefetch_stream"):
|
| 43 |
+
self.prefetch_stream = torch.cuda.Stream()
|
| 44 |
+
|
| 45 |
+
# delete previous layer
|
| 46 |
+
torch.cuda.current_stream().synchronize()
|
| 47 |
+
self.evict_previous_layer(layer_idx)
|
| 48 |
+
|
| 49 |
+
# make sure the current layer is ready
|
| 50 |
+
torch.cuda.synchronize(self.prefetch_stream)
|
| 51 |
+
|
| 52 |
+
# load next layer
|
| 53 |
+
self.prefetch_layer((layer_idx + 1) % len(self.layers), device)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def forward(
|
| 57 |
+
self,
|
| 58 |
+
input_ids: torch.LongTensor = None,
|
| 59 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 60 |
+
position_ids: Optional[torch.LongTensor] = None,
|
| 61 |
+
past_key_values: Optional[List[torch.FloatTensor]] = None,
|
| 62 |
+
inputs_embeds: Optional[torch.FloatTensor] = None,
|
| 63 |
+
use_cache: Optional[bool] = None,
|
| 64 |
+
output_attentions: Optional[bool] = None,
|
| 65 |
+
output_hidden_states: Optional[bool] = None,
|
| 66 |
+
return_dict: Optional[bool] = None,
|
| 67 |
+
cache_position: Optional[torch.LongTensor] = None,
|
| 68 |
+
offload_model: Optional[bool] = False,
|
| 69 |
+
) -> Union[Tuple, BaseModelOutputWithPast]:
|
| 70 |
+
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
| 71 |
+
output_hidden_states = (
|
| 72 |
+
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
| 73 |
+
)
|
| 74 |
+
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
| 75 |
+
|
| 76 |
+
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
| 77 |
+
|
| 78 |
+
if (input_ids is None) ^ (inputs_embeds is not None):
|
| 79 |
+
raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
|
| 80 |
+
|
| 81 |
+
if self.gradient_checkpointing and self.training:
|
| 82 |
+
if use_cache:
|
| 83 |
+
logger.warning_once(
|
| 84 |
+
"`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
|
| 85 |
+
)
|
| 86 |
+
use_cache = False
|
| 87 |
+
|
| 88 |
+
# kept for BC (non `Cache` `past_key_values` inputs)
|
| 89 |
+
return_legacy_cache = False
|
| 90 |
+
if use_cache and not isinstance(past_key_values, Cache):
|
| 91 |
+
return_legacy_cache = True
|
| 92 |
+
if past_key_values is None:
|
| 93 |
+
past_key_values = DynamicCache()
|
| 94 |
+
else:
|
| 95 |
+
past_key_values = DynamicCache.from_legacy_cache(past_key_values)
|
| 96 |
+
logger.warning_once(
|
| 97 |
+
"We detected that you are passing `past_key_values` as a tuple of tuples. This is deprecated and "
|
| 98 |
+
"will be removed in v4.47. Please convert your cache or use an appropriate `Cache` class "
|
| 99 |
+
"(https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)"
|
| 100 |
+
)
|
| 101 |
+
|
| 102 |
+
# if inputs_embeds is None:
|
| 103 |
+
# inputs_embeds = self.embed_tokens(input_ids)
|
| 104 |
+
|
| 105 |
+
# if cache_position is None:
|
| 106 |
+
# past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
|
| 107 |
+
# cache_position = torch.arange(
|
| 108 |
+
# past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
|
| 109 |
+
# )
|
| 110 |
+
# if position_ids is None:
|
| 111 |
+
# position_ids = cache_position.unsqueeze(0)
|
| 112 |
+
|
| 113 |
+
if attention_mask is not None and attention_mask.dim() == 3:
|
| 114 |
+
dtype = inputs_embeds.dtype
|
| 115 |
+
min_dtype = torch.finfo(dtype).min
|
| 116 |
+
attention_mask = (1 - attention_mask) * min_dtype
|
| 117 |
+
attention_mask = attention_mask.unsqueeze(1).to(inputs_embeds.dtype)
|
| 118 |
+
else:
|
| 119 |
+
raise Exception("attention_mask parameter was unavailable or invalid")
|
| 120 |
+
# causal_mask = self._update_causal_mask(
|
| 121 |
+
# attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions
|
| 122 |
+
# )
|
| 123 |
+
|
| 124 |
+
hidden_states = inputs_embeds
|
| 125 |
+
|
| 126 |
+
# decoder layers
|
| 127 |
+
all_hidden_states = () if output_hidden_states else None
|
| 128 |
+
all_self_attns = () if output_attentions else None
|
| 129 |
+
next_decoder_cache = None
|
| 130 |
+
|
| 131 |
+
layer_idx = -1
|
| 132 |
+
for decoder_layer in self.layers:
|
| 133 |
+
layer_idx += 1
|
| 134 |
+
|
| 135 |
+
if output_hidden_states:
|
| 136 |
+
all_hidden_states += (hidden_states,)
|
| 137 |
+
|
| 138 |
+
if self.gradient_checkpointing and self.training:
|
| 139 |
+
layer_outputs = self._gradient_checkpointing_func(
|
| 140 |
+
decoder_layer.__call__,
|
| 141 |
+
hidden_states,
|
| 142 |
+
attention_mask,
|
| 143 |
+
position_ids,
|
| 144 |
+
past_key_values,
|
| 145 |
+
output_attentions,
|
| 146 |
+
use_cache,
|
| 147 |
+
cache_position,
|
| 148 |
+
)
|
| 149 |
+
else:
|
| 150 |
+
if offload_model and not self.training:
|
| 151 |
+
self.get_offlaod_layer(layer_idx, device=inputs_embeds.device)
|
| 152 |
+
layer_outputs = decoder_layer(
|
| 153 |
+
hidden_states,
|
| 154 |
+
attention_mask=attention_mask,
|
| 155 |
+
position_ids=position_ids,
|
| 156 |
+
past_key_value=past_key_values,
|
| 157 |
+
output_attentions=output_attentions,
|
| 158 |
+
use_cache=use_cache,
|
| 159 |
+
cache_position=cache_position,
|
| 160 |
+
)
|
| 161 |
+
|
| 162 |
+
hidden_states = layer_outputs[0]
|
| 163 |
+
|
| 164 |
+
if use_cache:
|
| 165 |
+
next_decoder_cache = layer_outputs[2 if output_attentions else 1]
|
| 166 |
+
|
| 167 |
+
if output_attentions:
|
| 168 |
+
all_self_attns += (layer_outputs[1],)
|
| 169 |
+
|
| 170 |
+
hidden_states = self.norm(hidden_states)
|
| 171 |
+
|
| 172 |
+
# add hidden states from the last decoder layer
|
| 173 |
+
if output_hidden_states:
|
| 174 |
+
print('************')
|
| 175 |
+
all_hidden_states += (hidden_states,)
|
| 176 |
+
|
| 177 |
+
next_cache = next_decoder_cache if use_cache else None
|
| 178 |
+
if return_legacy_cache:
|
| 179 |
+
next_cache = next_cache.to_legacy_cache()
|
| 180 |
+
|
| 181 |
+
if not return_dict:
|
| 182 |
+
return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
|
| 183 |
+
return BaseModelOutputWithPast(
|
| 184 |
+
last_hidden_state=hidden_states,
|
| 185 |
+
past_key_values=next_cache,
|
| 186 |
+
hidden_states=all_hidden_states,
|
| 187 |
+
attentions=all_self_attns,
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
def modulate(x, shift, scale):
|
| 192 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
class TimestepEmbedder(nn.Module):
|
| 196 |
+
"""
|
| 197 |
+
Embeds scalar timesteps into vector representations.
|
| 198 |
+
"""
|
| 199 |
+
def __init__(self, hidden_size, frequency_embedding_size=256):
|
| 200 |
+
super().__init__()
|
| 201 |
+
self.mlp = nn.Sequential(
|
| 202 |
+
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
| 203 |
+
nn.SiLU(),
|
| 204 |
+
nn.Linear(hidden_size, hidden_size, bias=True),
|
| 205 |
+
)
|
| 206 |
+
self.frequency_embedding_size = frequency_embedding_size
|
| 207 |
+
|
| 208 |
+
@staticmethod
|
| 209 |
+
def timestep_embedding(t, dim, max_period=10000):
|
| 210 |
+
"""
|
| 211 |
+
Create sinusoidal timestep embeddings.
|
| 212 |
+
:param t: a 1-D Tensor of N indices, one per batch element.
|
| 213 |
+
These may be fractional.
|
| 214 |
+
:param dim: the dimension of the output.
|
| 215 |
+
:param max_period: controls the minimum frequency of the embeddings.
|
| 216 |
+
:return: an (N, D) Tensor of positional embeddings.
|
| 217 |
+
"""
|
| 218 |
+
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
| 219 |
+
half = dim // 2
|
| 220 |
+
freqs = torch.exp(
|
| 221 |
+
-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half
|
| 222 |
+
).to(device=t.device)
|
| 223 |
+
args = t[:, None].float() * freqs[None]
|
| 224 |
+
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
| 225 |
+
if dim % 2:
|
| 226 |
+
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
| 227 |
+
return embedding
|
| 228 |
+
|
| 229 |
+
def forward(self, t, dtype=torch.float32):
|
| 230 |
+
t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(dtype)
|
| 231 |
+
t_emb = self.mlp(t_freq)
|
| 232 |
+
return t_emb
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
class FinalLayer(nn.Module):
|
| 236 |
+
"""
|
| 237 |
+
The final layer of DiT.
|
| 238 |
+
"""
|
| 239 |
+
def __init__(self, hidden_size, patch_size, out_channels):
|
| 240 |
+
super().__init__()
|
| 241 |
+
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 242 |
+
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
| 243 |
+
self.adaLN_modulation = nn.Sequential(
|
| 244 |
+
nn.SiLU(),
|
| 245 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True)
|
| 246 |
+
)
|
| 247 |
+
|
| 248 |
+
def forward(self, x, c):
|
| 249 |
+
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 250 |
+
x = modulate(self.norm_final(x), shift, scale)
|
| 251 |
+
x = self.linear(x)
|
| 252 |
+
return x
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, interpolation_scale=1.0, base_size=1):
|
| 256 |
+
"""
|
| 257 |
+
grid_size: int of the grid height and width return: pos_embed: [grid_size*grid_size, embed_dim] or
|
| 258 |
+
[1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
|
| 259 |
+
"""
|
| 260 |
+
if isinstance(grid_size, int):
|
| 261 |
+
grid_size = (grid_size, grid_size)
|
| 262 |
+
|
| 263 |
+
grid_h = np.arange(grid_size[0], dtype=np.float32) / (grid_size[0] / base_size) / interpolation_scale
|
| 264 |
+
grid_w = np.arange(grid_size[1], dtype=np.float32) / (grid_size[1] / base_size) / interpolation_scale
|
| 265 |
+
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
| 266 |
+
grid = np.stack(grid, axis=0)
|
| 267 |
+
|
| 268 |
+
grid = grid.reshape([2, 1, grid_size[1], grid_size[0]])
|
| 269 |
+
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
| 270 |
+
if cls_token and extra_tokens > 0:
|
| 271 |
+
pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
| 272 |
+
return pos_embed
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
| 276 |
+
assert embed_dim % 2 == 0
|
| 277 |
+
|
| 278 |
+
# use half of dimensions to encode grid_h
|
| 279 |
+
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
| 280 |
+
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
| 281 |
+
|
| 282 |
+
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
|
| 283 |
+
return emb
|
| 284 |
+
|
| 285 |
+
|
| 286 |
+
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
| 287 |
+
"""
|
| 288 |
+
embed_dim: output dimension for each position
|
| 289 |
+
pos: a list of positions to be encoded: size (M,)
|
| 290 |
+
out: (M, D)
|
| 291 |
+
"""
|
| 292 |
+
assert embed_dim % 2 == 0
|
| 293 |
+
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
| 294 |
+
omega /= embed_dim / 2.
|
| 295 |
+
omega = 1. / 10000**omega # (D/2,)
|
| 296 |
+
|
| 297 |
+
pos = pos.reshape(-1) # (M,)
|
| 298 |
+
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
|
| 299 |
+
|
| 300 |
+
emb_sin = np.sin(out) # (M, D/2)
|
| 301 |
+
emb_cos = np.cos(out) # (M, D/2)
|
| 302 |
+
|
| 303 |
+
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
|
| 304 |
+
return emb
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
class PatchEmbedMR(nn.Module):
|
| 308 |
+
""" 2D Image to Patch Embedding
|
| 309 |
+
"""
|
| 310 |
+
def __init__(
|
| 311 |
+
self,
|
| 312 |
+
patch_size: int = 2,
|
| 313 |
+
in_chans: int = 4,
|
| 314 |
+
embed_dim: int = 768,
|
| 315 |
+
bias: bool = True,
|
| 316 |
+
):
|
| 317 |
+
super().__init__()
|
| 318 |
+
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
|
| 319 |
+
|
| 320 |
+
def forward(self, x):
|
| 321 |
+
x = self.proj(x)
|
| 322 |
+
x = x.flatten(2).transpose(1, 2) # NCHW -> NLC
|
| 323 |
+
return x
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
class OmniGenOriginalModel(nn.Module):
|
| 327 |
+
"""
|
| 328 |
+
Diffusion model with a Transformer backbone.
|
| 329 |
+
"""
|
| 330 |
+
def __init__(
|
| 331 |
+
self,
|
| 332 |
+
transformer_config: Phi3Config,
|
| 333 |
+
patch_size=2,
|
| 334 |
+
in_channels=4,
|
| 335 |
+
pe_interpolation: float = 1.0,
|
| 336 |
+
pos_embed_max_size: int = 192,
|
| 337 |
+
):
|
| 338 |
+
super().__init__()
|
| 339 |
+
self.in_channels = in_channels
|
| 340 |
+
self.out_channels = in_channels
|
| 341 |
+
self.patch_size = patch_size
|
| 342 |
+
self.pos_embed_max_size = pos_embed_max_size
|
| 343 |
+
|
| 344 |
+
hidden_size = transformer_config.hidden_size
|
| 345 |
+
|
| 346 |
+
self.x_embedder = PatchEmbedMR(patch_size, in_channels, hidden_size, bias=True)
|
| 347 |
+
self.input_x_embedder = PatchEmbedMR(patch_size, in_channels, hidden_size, bias=True)
|
| 348 |
+
|
| 349 |
+
self.time_token = TimestepEmbedder(hidden_size)
|
| 350 |
+
self.t_embedder = TimestepEmbedder(hidden_size)
|
| 351 |
+
|
| 352 |
+
self.pe_interpolation = pe_interpolation
|
| 353 |
+
pos_embed = get_2d_sincos_pos_embed(hidden_size, pos_embed_max_size, interpolation_scale=self.pe_interpolation, base_size=64)
|
| 354 |
+
self.register_buffer("pos_embed", torch.from_numpy(pos_embed).float().unsqueeze(0), persistent=True)
|
| 355 |
+
|
| 356 |
+
self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels)
|
| 357 |
+
|
| 358 |
+
self.initialize_weights()
|
| 359 |
+
|
| 360 |
+
self.llm = Phi3Transformer(config=transformer_config)
|
| 361 |
+
self.llm.config.use_cache = False
|
| 362 |
+
|
| 363 |
+
@classmethod
|
| 364 |
+
def from_pretrained(cls, model_name):
|
| 365 |
+
if not os.path.exists(model_name):
|
| 366 |
+
cache_folder = os.getenv('HF_HUB_CACHE')
|
| 367 |
+
model_name = snapshot_download(repo_id=model_name,
|
| 368 |
+
cache_dir=cache_folder,
|
| 369 |
+
ignore_patterns=['flax_model.msgpack', 'rust_model.ot', 'tf_model.h5'])
|
| 370 |
+
config = Phi3Config.from_pretrained(model_name)
|
| 371 |
+
model = cls(config)
|
| 372 |
+
if os.path.exists(os.path.join(model_name, 'model.safetensors')):
|
| 373 |
+
print("Loading safetensors")
|
| 374 |
+
ckpt = load_file(os.path.join(model_name, 'model.safetensors'))
|
| 375 |
+
else:
|
| 376 |
+
ckpt = torch.load(os.path.join(model_name, 'model.pt'), map_location='cpu')
|
| 377 |
+
model.load_state_dict(ckpt)
|
| 378 |
+
return model
|
| 379 |
+
|
| 380 |
+
def initialize_weights(self):
|
| 381 |
+
assert not hasattr(self, "llama")
|
| 382 |
+
|
| 383 |
+
# Initialize transformer layers:
|
| 384 |
+
def _basic_init(module):
|
| 385 |
+
if isinstance(module, nn.Linear):
|
| 386 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 387 |
+
if module.bias is not None:
|
| 388 |
+
nn.init.constant_(module.bias, 0)
|
| 389 |
+
self.apply(_basic_init)
|
| 390 |
+
|
| 391 |
+
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
| 392 |
+
w = self.x_embedder.proj.weight.data
|
| 393 |
+
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
| 394 |
+
nn.init.constant_(self.x_embedder.proj.bias, 0)
|
| 395 |
+
|
| 396 |
+
w = self.input_x_embedder.proj.weight.data
|
| 397 |
+
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
| 398 |
+
nn.init.constant_(self.x_embedder.proj.bias, 0)
|
| 399 |
+
|
| 400 |
+
|
| 401 |
+
# Initialize timestep embedding MLP:
|
| 402 |
+
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
| 403 |
+
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
| 404 |
+
nn.init.normal_(self.time_token.mlp[0].weight, std=0.02)
|
| 405 |
+
nn.init.normal_(self.time_token.mlp[2].weight, std=0.02)
|
| 406 |
+
|
| 407 |
+
# Zero-out output layers:
|
| 408 |
+
nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0)
|
| 409 |
+
nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0)
|
| 410 |
+
nn.init.constant_(self.final_layer.linear.weight, 0)
|
| 411 |
+
nn.init.constant_(self.final_layer.linear.bias, 0)
|
| 412 |
+
|
| 413 |
+
def unpatchify(self, x, h, w):
|
| 414 |
+
"""
|
| 415 |
+
x: (N, T, patch_size**2 * C)
|
| 416 |
+
imgs: (N, H, W, C)
|
| 417 |
+
"""
|
| 418 |
+
c = self.out_channels
|
| 419 |
+
|
| 420 |
+
x = x.reshape(shape=(x.shape[0], h//self.patch_size, w//self.patch_size, self.patch_size, self.patch_size, c))
|
| 421 |
+
x = torch.einsum('nhwpqc->nchpwq', x)
|
| 422 |
+
imgs = x.reshape(shape=(x.shape[0], c, h, w))
|
| 423 |
+
return imgs
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
def cropped_pos_embed(self, height, width):
|
| 427 |
+
"""Crops positional embeddings for SD3 compatibility."""
|
| 428 |
+
if self.pos_embed_max_size is None:
|
| 429 |
+
raise ValueError("`pos_embed_max_size` must be set for cropping.")
|
| 430 |
+
|
| 431 |
+
height = height // self.patch_size
|
| 432 |
+
width = width // self.patch_size
|
| 433 |
+
if height > self.pos_embed_max_size:
|
| 434 |
+
raise ValueError(
|
| 435 |
+
f"Height ({height}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}."
|
| 436 |
+
)
|
| 437 |
+
if width > self.pos_embed_max_size:
|
| 438 |
+
raise ValueError(
|
| 439 |
+
f"Width ({width}) cannot be greater than `pos_embed_max_size`: {self.pos_embed_max_size}."
|
| 440 |
+
)
|
| 441 |
+
|
| 442 |
+
top = (self.pos_embed_max_size - height) // 2
|
| 443 |
+
left = (self.pos_embed_max_size - width) // 2
|
| 444 |
+
spatial_pos_embed = self.pos_embed.reshape(1, self.pos_embed_max_size, self.pos_embed_max_size, -1)
|
| 445 |
+
spatial_pos_embed = spatial_pos_embed[:, top : top + height, left : left + width, :]
|
| 446 |
+
# print(top, top + height, left, left + width, spatial_pos_embed.size())
|
| 447 |
+
spatial_pos_embed = spatial_pos_embed.reshape(1, -1, spatial_pos_embed.shape[-1])
|
| 448 |
+
return spatial_pos_embed
|
| 449 |
+
|
| 450 |
+
|
| 451 |
+
def patch_multiple_resolutions(self, latents, padding_latent=None, is_input_images:bool=False):
|
| 452 |
+
if isinstance(latents, list):
|
| 453 |
+
return_list = False
|
| 454 |
+
if padding_latent is None:
|
| 455 |
+
padding_latent = [None] * len(latents)
|
| 456 |
+
return_list = True
|
| 457 |
+
patched_latents, num_tokens, shapes = [], [], []
|
| 458 |
+
for latent, padding in zip(latents, padding_latent):
|
| 459 |
+
height, width = latent.shape[-2:]
|
| 460 |
+
if is_input_images:
|
| 461 |
+
latent = self.input_x_embedder(latent)
|
| 462 |
+
else:
|
| 463 |
+
latent = self.x_embedder(latent)
|
| 464 |
+
pos_embed = self.cropped_pos_embed(height, width)
|
| 465 |
+
latent = latent + pos_embed
|
| 466 |
+
if padding is not None:
|
| 467 |
+
latent = torch.cat([latent, padding], dim=-2)
|
| 468 |
+
patched_latents.append(latent)
|
| 469 |
+
|
| 470 |
+
num_tokens.append(pos_embed.size(1))
|
| 471 |
+
shapes.append([height, width])
|
| 472 |
+
if not return_list:
|
| 473 |
+
latents = torch.cat(patched_latents, dim=0)
|
| 474 |
+
else:
|
| 475 |
+
latents = patched_latents
|
| 476 |
+
else:
|
| 477 |
+
height, width = latents.shape[-2:]
|
| 478 |
+
if is_input_images:
|
| 479 |
+
latents = self.input_x_embedder(latents)
|
| 480 |
+
else:
|
| 481 |
+
latents = self.x_embedder(latents)
|
| 482 |
+
pos_embed = self.cropped_pos_embed(height, width)
|
| 483 |
+
latents = latents + pos_embed
|
| 484 |
+
num_tokens = latents.size(1)
|
| 485 |
+
shapes = [height, width]
|
| 486 |
+
return latents, num_tokens, shapes
|
| 487 |
+
|
| 488 |
+
|
| 489 |
+
def forward(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, padding_latent=None, past_key_values=None, return_past_key_values=True, offload_model:bool=False):
|
| 490 |
+
"""
|
| 491 |
+
|
| 492 |
+
"""
|
| 493 |
+
input_is_list = isinstance(x, list)
|
| 494 |
+
x, num_tokens, shapes = self.patch_multiple_resolutions(x, padding_latent)
|
| 495 |
+
time_token = self.time_token(timestep, dtype=x[0].dtype).unsqueeze(1)
|
| 496 |
+
|
| 497 |
+
if input_img_latents is not None:
|
| 498 |
+
input_latents, _, _ = self.patch_multiple_resolutions(input_img_latents, is_input_images=True)
|
| 499 |
+
if input_ids is not None:
|
| 500 |
+
condition_embeds = self.llm.embed_tokens(input_ids).clone()
|
| 501 |
+
input_img_inx = 0
|
| 502 |
+
for b_inx in input_image_sizes.keys():
|
| 503 |
+
for start_inx, end_inx in input_image_sizes[b_inx]:
|
| 504 |
+
condition_embeds[b_inx, start_inx: end_inx] = input_latents[input_img_inx]
|
| 505 |
+
input_img_inx += 1
|
| 506 |
+
if input_img_latents is not None:
|
| 507 |
+
assert input_img_inx == len(input_latents)
|
| 508 |
+
|
| 509 |
+
input_emb = torch.cat([condition_embeds, time_token, x], dim=1)
|
| 510 |
+
else:
|
| 511 |
+
input_emb = torch.cat([time_token, x], dim=1)
|
| 512 |
+
output = self.llm(inputs_embeds=input_emb, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, offload_model=offload_model)
|
| 513 |
+
output, past_key_values = output.last_hidden_state, output.past_key_values
|
| 514 |
+
if input_is_list:
|
| 515 |
+
image_embedding = output[:, -max(num_tokens):]
|
| 516 |
+
time_emb = self.t_embedder(timestep, dtype=x.dtype)
|
| 517 |
+
x = self.final_layer(image_embedding, time_emb)
|
| 518 |
+
latents = []
|
| 519 |
+
for i in range(x.size(0)):
|
| 520 |
+
latent = x[i:i+1, :num_tokens[i]]
|
| 521 |
+
latent = self.unpatchify(latent, shapes[i][0], shapes[i][1])
|
| 522 |
+
latents.append(latent)
|
| 523 |
+
else:
|
| 524 |
+
image_embedding = output[:, -num_tokens:]
|
| 525 |
+
time_emb = self.t_embedder(timestep, dtype=x.dtype)
|
| 526 |
+
x = self.final_layer(image_embedding, time_emb)
|
| 527 |
+
latents = self.unpatchify(x, shapes[0], shapes[1])
|
| 528 |
+
|
| 529 |
+
if return_past_key_values:
|
| 530 |
+
return latents, past_key_values
|
| 531 |
+
return latents
|
| 532 |
+
|
| 533 |
+
@torch.no_grad()
|
| 534 |
+
def forward_with_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):
|
| 535 |
+
self.llm.config.use_cache = use_kv_cache
|
| 536 |
+
model_out, past_key_values = self.forward(x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, past_key_values=past_key_values, return_past_key_values=True, offload_model=offload_model)
|
| 537 |
+
if use_img_cfg:
|
| 538 |
+
cond, uncond, img_cond = torch.split(model_out, len(model_out) // 3, dim=0)
|
| 539 |
+
cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)
|
| 540 |
+
model_out = [cond, cond, cond]
|
| 541 |
+
else:
|
| 542 |
+
cond, uncond = torch.split(model_out, len(model_out) // 2, dim=0)
|
| 543 |
+
cond = uncond + cfg_scale * (cond - uncond)
|
| 544 |
+
model_out = [cond, cond]
|
| 545 |
+
|
| 546 |
+
return torch.cat(model_out, dim=0), past_key_values
|
| 547 |
+
|
| 548 |
+
|
| 549 |
+
@torch.no_grad()
|
| 550 |
+
def forward_with_separate_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):
|
| 551 |
+
self.llm.config.use_cache = use_kv_cache
|
| 552 |
+
if past_key_values is None:
|
| 553 |
+
past_key_values = [None] * len(attention_mask)
|
| 554 |
+
|
| 555 |
+
x = torch.split(x, len(x) // len(attention_mask), dim=0)
|
| 556 |
+
timestep = timestep.to(x[0].dtype)
|
| 557 |
+
timestep = torch.split(timestep, len(timestep) // len(input_ids), dim=0)
|
| 558 |
+
|
| 559 |
+
model_out, pask_key_values = [], []
|
| 560 |
+
for i in range(len(input_ids)):
|
| 561 |
+
temp_out, temp_pask_key_values = self.forward(x[i], timestep[i], input_ids[i], input_img_latents[i], input_image_sizes[i], attention_mask[i], position_ids[i], past_key_values=past_key_values[i], return_past_key_values=True, offload_model=offload_model)
|
| 562 |
+
model_out.append(temp_out)
|
| 563 |
+
pask_key_values.append(temp_pask_key_values)
|
| 564 |
+
|
| 565 |
+
if len(model_out) == 3:
|
| 566 |
+
cond, uncond, img_cond = model_out
|
| 567 |
+
cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)
|
| 568 |
+
model_out = [cond, cond, cond]
|
| 569 |
+
elif len(model_out) == 2:
|
| 570 |
+
cond, uncond = model_out
|
| 571 |
+
cond = uncond + cfg_scale * (cond - uncond)
|
| 572 |
+
model_out = [cond, cond]
|
| 573 |
+
else:
|
| 574 |
+
return model_out[0]
|
| 575 |
+
|
| 576 |
+
return torch.cat(model_out, dim=0), pask_key_values
|
| 577 |
+
|
| 578 |
+
|
| 579 |
+
|
| 580 |
+
class OmniGenTransformer(OmniGenOriginalModel):
|
| 581 |
+
def __init__(self):
|
| 582 |
+
config = {
|
| 583 |
+
"_name_or_path": "Phi-3-vision-128k-instruct",
|
| 584 |
+
"architectures": [
|
| 585 |
+
"Phi3ForCausalLM"
|
| 586 |
+
],
|
| 587 |
+
"attention_dropout": 0.0,
|
| 588 |
+
"bos_token_id": 1,
|
| 589 |
+
"eos_token_id": 2,
|
| 590 |
+
"hidden_act": "silu",
|
| 591 |
+
"hidden_size": 3072,
|
| 592 |
+
"initializer_range": 0.02,
|
| 593 |
+
"intermediate_size": 8192,
|
| 594 |
+
"max_position_embeddings": 131072,
|
| 595 |
+
"model_type": "phi3",
|
| 596 |
+
"num_attention_heads": 32,
|
| 597 |
+
"num_hidden_layers": 32,
|
| 598 |
+
"num_key_value_heads": 32,
|
| 599 |
+
"original_max_position_embeddings": 4096,
|
| 600 |
+
"rms_norm_eps": 1e-05,
|
| 601 |
+
"rope_scaling": {
|
| 602 |
+
"long_factor": [
|
| 603 |
+
1.0299999713897705,
|
| 604 |
+
1.0499999523162842,
|
| 605 |
+
1.0499999523162842,
|
| 606 |
+
1.0799999237060547,
|
| 607 |
+
1.2299998998641968,
|
| 608 |
+
1.2299998998641968,
|
| 609 |
+
1.2999999523162842,
|
| 610 |
+
1.4499999284744263,
|
| 611 |
+
1.5999999046325684,
|
| 612 |
+
1.6499998569488525,
|
| 613 |
+
1.8999998569488525,
|
| 614 |
+
2.859999895095825,
|
| 615 |
+
3.68999981880188,
|
| 616 |
+
5.419999599456787,
|
| 617 |
+
5.489999771118164,
|
| 618 |
+
5.489999771118164,
|
| 619 |
+
9.09000015258789,
|
| 620 |
+
11.579999923706055,
|
| 621 |
+
15.65999984741211,
|
| 622 |
+
15.769999504089355,
|
| 623 |
+
15.789999961853027,
|
| 624 |
+
18.360000610351562,
|
| 625 |
+
21.989999771118164,
|
| 626 |
+
23.079999923706055,
|
| 627 |
+
30.009998321533203,
|
| 628 |
+
32.35000228881836,
|
| 629 |
+
32.590003967285156,
|
| 630 |
+
35.56000518798828,
|
| 631 |
+
39.95000457763672,
|
| 632 |
+
53.840003967285156,
|
| 633 |
+
56.20000457763672,
|
| 634 |
+
57.95000457763672,
|
| 635 |
+
59.29000473022461,
|
| 636 |
+
59.77000427246094,
|
| 637 |
+
59.920005798339844,
|
| 638 |
+
61.190006256103516,
|
| 639 |
+
61.96000671386719,
|
| 640 |
+
62.50000762939453,
|
| 641 |
+
63.3700065612793,
|
| 642 |
+
63.48000717163086,
|
| 643 |
+
63.48000717163086,
|
| 644 |
+
63.66000747680664,
|
| 645 |
+
63.850006103515625,
|
| 646 |
+
64.08000946044922,
|
| 647 |
+
64.760009765625,
|
| 648 |
+
64.80001068115234,
|
| 649 |
+
64.81001281738281,
|
| 650 |
+
64.81001281738281
|
| 651 |
+
],
|
| 652 |
+
"short_factor": [
|
| 653 |
+
1.05,
|
| 654 |
+
1.05,
|
| 655 |
+
1.05,
|
| 656 |
+
1.1,
|
| 657 |
+
1.1,
|
| 658 |
+
1.1,
|
| 659 |
+
1.2500000000000002,
|
| 660 |
+
1.2500000000000002,
|
| 661 |
+
1.4000000000000004,
|
| 662 |
+
1.4500000000000004,
|
| 663 |
+
1.5500000000000005,
|
| 664 |
+
1.8500000000000008,
|
| 665 |
+
1.9000000000000008,
|
| 666 |
+
2.000000000000001,
|
| 667 |
+
2.000000000000001,
|
| 668 |
+
2.000000000000001,
|
| 669 |
+
2.000000000000001,
|
| 670 |
+
2.000000000000001,
|
| 671 |
+
2.000000000000001,
|
| 672 |
+
2.000000000000001,
|
| 673 |
+
2.000000000000001,
|
| 674 |
+
2.000000000000001,
|
| 675 |
+
2.000000000000001,
|
| 676 |
+
2.000000000000001,
|
| 677 |
+
2.000000000000001,
|
| 678 |
+
2.000000000000001,
|
| 679 |
+
2.000000000000001,
|
| 680 |
+
2.000000000000001,
|
| 681 |
+
2.000000000000001,
|
| 682 |
+
2.000000000000001,
|
| 683 |
+
2.000000000000001,
|
| 684 |
+
2.000000000000001,
|
| 685 |
+
2.1000000000000005,
|
| 686 |
+
2.1000000000000005,
|
| 687 |
+
2.2,
|
| 688 |
+
2.3499999999999996,
|
| 689 |
+
2.3499999999999996,
|
| 690 |
+
2.3499999999999996,
|
| 691 |
+
2.3499999999999996,
|
| 692 |
+
2.3999999999999995,
|
| 693 |
+
2.3999999999999995,
|
| 694 |
+
2.6499999999999986,
|
| 695 |
+
2.6999999999999984,
|
| 696 |
+
2.8999999999999977,
|
| 697 |
+
2.9499999999999975,
|
| 698 |
+
3.049999999999997,
|
| 699 |
+
3.049999999999997,
|
| 700 |
+
3.049999999999997
|
| 701 |
+
],
|
| 702 |
+
"type": "su"
|
| 703 |
+
},
|
| 704 |
+
"rope_theta": 10000.0,
|
| 705 |
+
"sliding_window": 131072,
|
| 706 |
+
"tie_word_embeddings": False,
|
| 707 |
+
"torch_dtype": "bfloat16",
|
| 708 |
+
"transformers_version": "4.38.1",
|
| 709 |
+
"use_cache": True,
|
| 710 |
+
"vocab_size": 32064,
|
| 711 |
+
"_attn_implementation": "sdpa"
|
| 712 |
+
}
|
| 713 |
+
config = Phi3Config(**config)
|
| 714 |
+
super().__init__(config)
|
| 715 |
+
|
| 716 |
+
|
| 717 |
+
def forward(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, padding_latent=None, past_key_values=None, return_past_key_values=True, offload_model:bool=False):
|
| 718 |
+
input_is_list = isinstance(x, list)
|
| 719 |
+
x, num_tokens, shapes = self.patch_multiple_resolutions(x, padding_latent)
|
| 720 |
+
time_token = self.time_token(timestep, dtype=x[0].dtype).unsqueeze(1)
|
| 721 |
+
|
| 722 |
+
if input_img_latents is not None:
|
| 723 |
+
input_latents, _, _ = self.patch_multiple_resolutions(input_img_latents, is_input_images=True)
|
| 724 |
+
if input_ids is not None:
|
| 725 |
+
condition_embeds = self.llm.embed_tokens(input_ids).clone()
|
| 726 |
+
input_img_inx = 0
|
| 727 |
+
for b_inx in input_image_sizes.keys():
|
| 728 |
+
for start_inx, end_inx in input_image_sizes[b_inx]:
|
| 729 |
+
condition_embeds[b_inx, start_inx: end_inx] = input_latents[input_img_inx]
|
| 730 |
+
input_img_inx += 1
|
| 731 |
+
if input_img_latents is not None:
|
| 732 |
+
assert input_img_inx == len(input_latents)
|
| 733 |
+
|
| 734 |
+
input_emb = torch.cat([condition_embeds, time_token, x], dim=1)
|
| 735 |
+
else:
|
| 736 |
+
input_emb = torch.cat([time_token, x], dim=1)
|
| 737 |
+
output = self.llm(inputs_embeds=input_emb, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, offload_model=offload_model)
|
| 738 |
+
output, past_key_values = output.last_hidden_state, output.past_key_values
|
| 739 |
+
if input_is_list:
|
| 740 |
+
image_embedding = output[:, -max(num_tokens):]
|
| 741 |
+
time_emb = self.t_embedder(timestep, dtype=x.dtype)
|
| 742 |
+
x = self.final_layer(image_embedding, time_emb)
|
| 743 |
+
latents = []
|
| 744 |
+
for i in range(x.size(0)):
|
| 745 |
+
latent = x[i:i+1, :num_tokens[i]]
|
| 746 |
+
latent = self.unpatchify(latent, shapes[i][0], shapes[i][1])
|
| 747 |
+
latents.append(latent)
|
| 748 |
+
else:
|
| 749 |
+
image_embedding = output[:, -num_tokens:]
|
| 750 |
+
time_emb = self.t_embedder(timestep, dtype=x.dtype)
|
| 751 |
+
x = self.final_layer(image_embedding, time_emb)
|
| 752 |
+
latents = self.unpatchify(x, shapes[0], shapes[1])
|
| 753 |
+
|
| 754 |
+
if return_past_key_values:
|
| 755 |
+
return latents, past_key_values
|
| 756 |
+
return latents
|
| 757 |
+
|
| 758 |
+
|
| 759 |
+
@torch.no_grad()
|
| 760 |
+
def forward_with_separate_cfg(self, x, timestep, input_ids, input_img_latents, input_image_sizes, attention_mask, position_ids, cfg_scale, use_img_cfg, img_cfg_scale, past_key_values, use_kv_cache, offload_model):
|
| 761 |
+
self.llm.config.use_cache = use_kv_cache
|
| 762 |
+
if past_key_values is None:
|
| 763 |
+
past_key_values = [None] * len(attention_mask)
|
| 764 |
+
|
| 765 |
+
x = torch.split(x, len(x) // len(attention_mask), dim=0)
|
| 766 |
+
timestep = timestep.to(x[0].dtype)
|
| 767 |
+
timestep = torch.split(timestep, len(timestep) // len(input_ids), dim=0)
|
| 768 |
+
|
| 769 |
+
model_out, pask_key_values = [], []
|
| 770 |
+
for i in range(len(input_ids)):
|
| 771 |
+
temp_out, temp_pask_key_values = self.forward(x[i], timestep[i], input_ids[i], input_img_latents[i], input_image_sizes[i], attention_mask[i], position_ids[i], past_key_values=past_key_values[i], return_past_key_values=True, offload_model=offload_model)
|
| 772 |
+
model_out.append(temp_out)
|
| 773 |
+
pask_key_values.append(temp_pask_key_values)
|
| 774 |
+
|
| 775 |
+
if len(model_out) == 3:
|
| 776 |
+
cond, uncond, img_cond = model_out
|
| 777 |
+
cond = uncond + img_cfg_scale * (img_cond - uncond) + cfg_scale * (cond - img_cond)
|
| 778 |
+
model_out = [cond, cond, cond]
|
| 779 |
+
elif len(model_out) == 2:
|
| 780 |
+
cond, uncond = model_out
|
| 781 |
+
cond = uncond + cfg_scale * (cond - uncond)
|
| 782 |
+
model_out = [cond, cond]
|
| 783 |
+
else:
|
| 784 |
+
return model_out[0]
|
| 785 |
+
|
| 786 |
+
return torch.cat(model_out, dim=0), pask_key_values
|
| 787 |
+
|
| 788 |
+
|
| 789 |
+
@staticmethod
|
| 790 |
+
def state_dict_converter():
|
| 791 |
+
return OmniGenTransformerStateDictConverter()
|
| 792 |
+
|
| 793 |
+
|
| 794 |
+
|
| 795 |
+
class OmniGenTransformerStateDictConverter:
|
| 796 |
+
def __init__(self):
|
| 797 |
+
pass
|
| 798 |
+
|
| 799 |
+
def from_diffusers(self, state_dict):
|
| 800 |
+
return state_dict
|
| 801 |
+
|
| 802 |
+
def from_civitai(self, state_dict):
|
| 803 |
+
return state_dict
|
diffsynth/models/qwen_image_dit.py
ADDED
|
@@ -0,0 +1,357 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from typing import Tuple, Optional, Union, List
|
| 4 |
+
from einops import rearrange
|
| 5 |
+
from .sd3_dit import TimestepEmbeddings, RMSNorm
|
| 6 |
+
from .flux_dit import AdaLayerNorm
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class ApproximateGELU(nn.Module):
|
| 10 |
+
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
|
| 11 |
+
super().__init__()
|
| 12 |
+
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
| 13 |
+
|
| 14 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 15 |
+
x = self.proj(x)
|
| 16 |
+
return x * torch.sigmoid(1.702 * x)
|
| 17 |
+
|
| 18 |
+
def apply_rotary_emb_qwen(
|
| 19 |
+
x: torch.Tensor,
|
| 20 |
+
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]]
|
| 21 |
+
):
|
| 22 |
+
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
| 23 |
+
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
|
| 24 |
+
return x_out.type_as(x)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
class QwenEmbedRope(nn.Module):
|
| 28 |
+
def __init__(self, theta: int, axes_dim: list[int], scale_rope=False):
|
| 29 |
+
super().__init__()
|
| 30 |
+
self.theta = theta
|
| 31 |
+
self.axes_dim = axes_dim
|
| 32 |
+
pos_index = torch.arange(1024)
|
| 33 |
+
neg_index = torch.arange(1024).flip(0) * -1 - 1
|
| 34 |
+
self.pos_freqs = torch.cat([
|
| 35 |
+
self.rope_params(pos_index, self.axes_dim[0], self.theta),
|
| 36 |
+
self.rope_params(pos_index, self.axes_dim[1], self.theta),
|
| 37 |
+
self.rope_params(pos_index, self.axes_dim[2], self.theta),
|
| 38 |
+
], dim=1)
|
| 39 |
+
self.neg_freqs = torch.cat([
|
| 40 |
+
self.rope_params(neg_index, self.axes_dim[0], self.theta),
|
| 41 |
+
self.rope_params(neg_index, self.axes_dim[1], self.theta),
|
| 42 |
+
self.rope_params(neg_index, self.axes_dim[2], self.theta),
|
| 43 |
+
], dim=1)
|
| 44 |
+
self.rope_cache = {}
|
| 45 |
+
self.scale_rope = scale_rope
|
| 46 |
+
|
| 47 |
+
def rope_params(self, index, dim, theta=10000):
|
| 48 |
+
"""
|
| 49 |
+
Args:
|
| 50 |
+
index: [0, 1, 2, 3] 1D Tensor representing the position index of the token
|
| 51 |
+
"""
|
| 52 |
+
assert dim % 2 == 0
|
| 53 |
+
freqs = torch.outer(
|
| 54 |
+
index,
|
| 55 |
+
1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))
|
| 56 |
+
)
|
| 57 |
+
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
| 58 |
+
return freqs
|
| 59 |
+
|
| 60 |
+
def forward(self, video_fhw, txt_seq_lens, device):
|
| 61 |
+
if self.pos_freqs.device != device:
|
| 62 |
+
self.pos_freqs = self.pos_freqs.to(device)
|
| 63 |
+
self.neg_freqs = self.neg_freqs.to(device)
|
| 64 |
+
|
| 65 |
+
if isinstance(video_fhw, list):
|
| 66 |
+
video_fhw = video_fhw[0]
|
| 67 |
+
frame, height, width = video_fhw
|
| 68 |
+
rope_key = f"{frame}_{height}_{width}"
|
| 69 |
+
|
| 70 |
+
if rope_key not in self.rope_cache:
|
| 71 |
+
seq_lens = frame * height * width
|
| 72 |
+
freqs_pos = self.pos_freqs.split([x // 2 for x in self.axes_dim], dim=1)
|
| 73 |
+
freqs_neg = self.neg_freqs.split([x // 2 for x in self.axes_dim], dim=1)
|
| 74 |
+
freqs_frame = freqs_pos[0][:frame].view(frame, 1, 1, -1).expand(frame, height, width, -1)
|
| 75 |
+
if self.scale_rope:
|
| 76 |
+
freqs_height = torch.cat(
|
| 77 |
+
[
|
| 78 |
+
freqs_neg[1][-(height - height//2):],
|
| 79 |
+
freqs_pos[1][:height//2]
|
| 80 |
+
],
|
| 81 |
+
dim=0
|
| 82 |
+
)
|
| 83 |
+
freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1)
|
| 84 |
+
freqs_width = torch.cat(
|
| 85 |
+
[
|
| 86 |
+
freqs_neg[2][-(width - width//2):],
|
| 87 |
+
freqs_pos[2][:width//2]
|
| 88 |
+
],
|
| 89 |
+
dim=0
|
| 90 |
+
)
|
| 91 |
+
freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1)
|
| 92 |
+
|
| 93 |
+
else:
|
| 94 |
+
freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1)
|
| 95 |
+
freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1)
|
| 96 |
+
|
| 97 |
+
freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1)
|
| 98 |
+
self.rope_cache[rope_key] = freqs.clone().contiguous()
|
| 99 |
+
vid_freqs = self.rope_cache[rope_key]
|
| 100 |
+
|
| 101 |
+
if self.scale_rope:
|
| 102 |
+
max_vid_index = max(height // 2, width // 2)
|
| 103 |
+
else:
|
| 104 |
+
max_vid_index = max(height, width)
|
| 105 |
+
|
| 106 |
+
max_len = max(txt_seq_lens)
|
| 107 |
+
txt_freqs = self.pos_freqs[max_vid_index: max_vid_index + max_len, ...]
|
| 108 |
+
return vid_freqs, txt_freqs
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
class QwenFeedForward(nn.Module):
|
| 112 |
+
def __init__(
|
| 113 |
+
self,
|
| 114 |
+
dim: int,
|
| 115 |
+
dim_out: Optional[int] = None,
|
| 116 |
+
dropout: float = 0.0,
|
| 117 |
+
):
|
| 118 |
+
super().__init__()
|
| 119 |
+
inner_dim = int(dim * 4)
|
| 120 |
+
self.net = nn.ModuleList([])
|
| 121 |
+
self.net.append(ApproximateGELU(dim, inner_dim))
|
| 122 |
+
self.net.append(nn.Dropout(dropout))
|
| 123 |
+
self.net.append(nn.Linear(inner_dim, dim_out))
|
| 124 |
+
|
| 125 |
+
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
| 126 |
+
for module in self.net:
|
| 127 |
+
hidden_states = module(hidden_states)
|
| 128 |
+
return hidden_states
|
| 129 |
+
|
| 130 |
+
class QwenDoubleStreamAttention(nn.Module):
|
| 131 |
+
def __init__(
|
| 132 |
+
self,
|
| 133 |
+
dim_a,
|
| 134 |
+
dim_b,
|
| 135 |
+
num_heads,
|
| 136 |
+
head_dim,
|
| 137 |
+
):
|
| 138 |
+
super().__init__()
|
| 139 |
+
self.num_heads = num_heads
|
| 140 |
+
self.head_dim = head_dim
|
| 141 |
+
|
| 142 |
+
self.to_q = nn.Linear(dim_a, dim_a)
|
| 143 |
+
self.to_k = nn.Linear(dim_a, dim_a)
|
| 144 |
+
self.to_v = nn.Linear(dim_a, dim_a)
|
| 145 |
+
self.norm_q = RMSNorm(head_dim, eps=1e-6)
|
| 146 |
+
self.norm_k = RMSNorm(head_dim, eps=1e-6)
|
| 147 |
+
|
| 148 |
+
self.add_q_proj = nn.Linear(dim_b, dim_b)
|
| 149 |
+
self.add_k_proj = nn.Linear(dim_b, dim_b)
|
| 150 |
+
self.add_v_proj = nn.Linear(dim_b, dim_b)
|
| 151 |
+
self.norm_added_q = RMSNorm(head_dim, eps=1e-6)
|
| 152 |
+
self.norm_added_k = RMSNorm(head_dim, eps=1e-6)
|
| 153 |
+
|
| 154 |
+
self.to_out = torch.nn.Sequential(nn.Linear(dim_a, dim_a))
|
| 155 |
+
self.to_add_out = nn.Linear(dim_b, dim_b)
|
| 156 |
+
|
| 157 |
+
def forward(
|
| 158 |
+
self,
|
| 159 |
+
image: torch.FloatTensor,
|
| 160 |
+
text: torch.FloatTensor,
|
| 161 |
+
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
| 162 |
+
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
|
| 163 |
+
img_q, img_k, img_v = self.to_q(image), self.to_k(image), self.to_v(image)
|
| 164 |
+
txt_q, txt_k, txt_v = self.add_q_proj(text), self.add_k_proj(text), self.add_v_proj(text)
|
| 165 |
+
seq_txt = txt_q.shape[1]
|
| 166 |
+
|
| 167 |
+
img_q = rearrange(img_q, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 168 |
+
img_k = rearrange(img_k, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 169 |
+
img_v = rearrange(img_v, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 170 |
+
|
| 171 |
+
txt_q = rearrange(txt_q, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 172 |
+
txt_k = rearrange(txt_k, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 173 |
+
txt_v = rearrange(txt_v, 'b s (h d) -> b h s d', h=self.num_heads)
|
| 174 |
+
|
| 175 |
+
img_q, img_k = self.norm_q(img_q), self.norm_k(img_k)
|
| 176 |
+
txt_q, txt_k = self.norm_added_q(txt_q), self.norm_added_k(txt_k)
|
| 177 |
+
|
| 178 |
+
if image_rotary_emb is not None:
|
| 179 |
+
img_freqs, txt_freqs = image_rotary_emb
|
| 180 |
+
img_q = apply_rotary_emb_qwen(img_q, img_freqs)
|
| 181 |
+
img_k = apply_rotary_emb_qwen(img_k, img_freqs)
|
| 182 |
+
txt_q = apply_rotary_emb_qwen(txt_q, txt_freqs)
|
| 183 |
+
txt_k = apply_rotary_emb_qwen(txt_k, txt_freqs)
|
| 184 |
+
|
| 185 |
+
joint_q = torch.cat([txt_q, img_q], dim=2)
|
| 186 |
+
joint_k = torch.cat([txt_k, img_k], dim=2)
|
| 187 |
+
joint_v = torch.cat([txt_v, img_v], dim=2)
|
| 188 |
+
|
| 189 |
+
joint_attn_out = torch.nn.functional.scaled_dot_product_attention(joint_q, joint_k, joint_v)
|
| 190 |
+
|
| 191 |
+
joint_attn_out = rearrange(joint_attn_out, 'b h s d -> b s (h d)').to(joint_q.dtype)
|
| 192 |
+
|
| 193 |
+
txt_attn_output = joint_attn_out[:, :seq_txt, :]
|
| 194 |
+
img_attn_output = joint_attn_out[:, seq_txt:, :]
|
| 195 |
+
|
| 196 |
+
img_attn_output = self.to_out(img_attn_output)
|
| 197 |
+
txt_attn_output = self.to_add_out(txt_attn_output)
|
| 198 |
+
|
| 199 |
+
return img_attn_output, txt_attn_output
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
class QwenImageTransformerBlock(nn.Module):
|
| 203 |
+
def __init__(
|
| 204 |
+
self,
|
| 205 |
+
dim: int,
|
| 206 |
+
num_attention_heads: int,
|
| 207 |
+
attention_head_dim: int,
|
| 208 |
+
eps: float = 1e-6,
|
| 209 |
+
):
|
| 210 |
+
super().__init__()
|
| 211 |
+
|
| 212 |
+
self.dim = dim
|
| 213 |
+
self.num_attention_heads = num_attention_heads
|
| 214 |
+
self.attention_head_dim = attention_head_dim
|
| 215 |
+
|
| 216 |
+
self.img_mod = nn.Sequential(
|
| 217 |
+
nn.SiLU(),
|
| 218 |
+
nn.Linear(dim, 6 * dim),
|
| 219 |
+
)
|
| 220 |
+
self.img_norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
|
| 221 |
+
self.attn = QwenDoubleStreamAttention(
|
| 222 |
+
dim_a=dim,
|
| 223 |
+
dim_b=dim,
|
| 224 |
+
num_heads=num_attention_heads,
|
| 225 |
+
head_dim=attention_head_dim,
|
| 226 |
+
)
|
| 227 |
+
self.img_norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
|
| 228 |
+
self.img_mlp = QwenFeedForward(dim=dim, dim_out=dim)
|
| 229 |
+
|
| 230 |
+
self.txt_mod = nn.Sequential(
|
| 231 |
+
nn.SiLU(),
|
| 232 |
+
nn.Linear(dim, 6 * dim, bias=True),
|
| 233 |
+
)
|
| 234 |
+
self.txt_norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
|
| 235 |
+
self.txt_norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=eps)
|
| 236 |
+
self.txt_mlp = QwenFeedForward(dim=dim, dim_out=dim)
|
| 237 |
+
|
| 238 |
+
def _modulate(self, x, mod_params):
|
| 239 |
+
shift, scale, gate = mod_params.chunk(3, dim=-1)
|
| 240 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1), gate.unsqueeze(1)
|
| 241 |
+
|
| 242 |
+
def forward(
|
| 243 |
+
self,
|
| 244 |
+
image: torch.Tensor,
|
| 245 |
+
text: torch.Tensor,
|
| 246 |
+
temb: torch.Tensor,
|
| 247 |
+
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
| 248 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 249 |
+
|
| 250 |
+
img_mod_attn, img_mod_mlp = self.img_mod(temb).chunk(2, dim=-1) # [B, 3*dim] each
|
| 251 |
+
txt_mod_attn, txt_mod_mlp = self.txt_mod(temb).chunk(2, dim=-1) # [B, 3*dim] each
|
| 252 |
+
|
| 253 |
+
img_normed = self.img_norm1(image)
|
| 254 |
+
img_modulated, img_gate = self._modulate(img_normed, img_mod_attn)
|
| 255 |
+
|
| 256 |
+
txt_normed = self.txt_norm1(text)
|
| 257 |
+
txt_modulated, txt_gate = self._modulate(txt_normed, txt_mod_attn)
|
| 258 |
+
|
| 259 |
+
img_attn_out, txt_attn_out = self.attn(
|
| 260 |
+
image=img_modulated,
|
| 261 |
+
text=txt_modulated,
|
| 262 |
+
image_rotary_emb=image_rotary_emb,
|
| 263 |
+
)
|
| 264 |
+
|
| 265 |
+
image = image + img_gate * img_attn_out
|
| 266 |
+
text = text + txt_gate * txt_attn_out
|
| 267 |
+
|
| 268 |
+
img_normed_2 = self.img_norm2(image)
|
| 269 |
+
img_modulated_2, img_gate_2 = self._modulate(img_normed_2, img_mod_mlp)
|
| 270 |
+
|
| 271 |
+
txt_normed_2 = self.txt_norm2(text)
|
| 272 |
+
txt_modulated_2, txt_gate_2 = self._modulate(txt_normed_2, txt_mod_mlp)
|
| 273 |
+
|
| 274 |
+
img_mlp_out = self.img_mlp(img_modulated_2)
|
| 275 |
+
txt_mlp_out = self.txt_mlp(txt_modulated_2)
|
| 276 |
+
|
| 277 |
+
image = image + img_gate_2 * img_mlp_out
|
| 278 |
+
text = text + txt_gate_2 * txt_mlp_out
|
| 279 |
+
|
| 280 |
+
return text, image
|
| 281 |
+
|
| 282 |
+
|
| 283 |
+
class QwenImageDiT(torch.nn.Module):
|
| 284 |
+
def __init__(
|
| 285 |
+
self,
|
| 286 |
+
num_layers: int = 60,
|
| 287 |
+
):
|
| 288 |
+
super().__init__()
|
| 289 |
+
|
| 290 |
+
self.pos_embed = QwenEmbedRope(theta=10000, axes_dim=[16,56,56], scale_rope=True)
|
| 291 |
+
|
| 292 |
+
self.time_text_embed = TimestepEmbeddings(256, 3072, diffusers_compatible_format=True, scale=1000, align_dtype_to_timestep=True)
|
| 293 |
+
self.txt_norm = RMSNorm(3584, eps=1e-6)
|
| 294 |
+
|
| 295 |
+
self.img_in = nn.Linear(64, 3072)
|
| 296 |
+
self.txt_in = nn.Linear(3584, 3072)
|
| 297 |
+
|
| 298 |
+
self.transformer_blocks = nn.ModuleList(
|
| 299 |
+
[
|
| 300 |
+
QwenImageTransformerBlock(
|
| 301 |
+
dim=3072,
|
| 302 |
+
num_attention_heads=24,
|
| 303 |
+
attention_head_dim=128,
|
| 304 |
+
)
|
| 305 |
+
for _ in range(num_layers)
|
| 306 |
+
]
|
| 307 |
+
)
|
| 308 |
+
self.norm_out = AdaLayerNorm(3072, single=True)
|
| 309 |
+
self.proj_out = nn.Linear(3072, 64)
|
| 310 |
+
|
| 311 |
+
|
| 312 |
+
def forward(
|
| 313 |
+
self,
|
| 314 |
+
latents=None,
|
| 315 |
+
timestep=None,
|
| 316 |
+
prompt_emb=None,
|
| 317 |
+
prompt_emb_mask=None,
|
| 318 |
+
height=None,
|
| 319 |
+
width=None,
|
| 320 |
+
):
|
| 321 |
+
img_shapes = [(latents.shape[0], latents.shape[2]//2, latents.shape[3]//2)]
|
| 322 |
+
txt_seq_lens = prompt_emb_mask.sum(dim=1).tolist()
|
| 323 |
+
|
| 324 |
+
image = rearrange(latents, "B C (H P) (W Q) -> B (H W) (P Q C)", H=height//16, W=width//16, P=2, Q=2)
|
| 325 |
+
image = self.img_in(image)
|
| 326 |
+
text = self.txt_in(self.txt_norm(prompt_emb))
|
| 327 |
+
|
| 328 |
+
conditioning = self.time_text_embed(timestep, image.dtype)
|
| 329 |
+
|
| 330 |
+
image_rotary_emb = self.pos_embed(img_shapes, txt_seq_lens, device=latents.device)
|
| 331 |
+
|
| 332 |
+
for block in self.transformer_blocks:
|
| 333 |
+
text, image = block(
|
| 334 |
+
image=image,
|
| 335 |
+
text=text,
|
| 336 |
+
temb=conditioning,
|
| 337 |
+
image_rotary_emb=image_rotary_emb,
|
| 338 |
+
)
|
| 339 |
+
|
| 340 |
+
image = self.norm_out(image, conditioning)
|
| 341 |
+
image = self.proj_out(image)
|
| 342 |
+
|
| 343 |
+
latents = rearrange(image, "B (H W) (P Q C) -> B C (H P) (W Q)", H=height//16, W=width//16, P=2, Q=2)
|
| 344 |
+
return image
|
| 345 |
+
|
| 346 |
+
@staticmethod
|
| 347 |
+
def state_dict_converter():
|
| 348 |
+
return QwenImageDiTStateDictConverter()
|
| 349 |
+
|
| 350 |
+
|
| 351 |
+
|
| 352 |
+
class QwenImageDiTStateDictConverter():
|
| 353 |
+
def __init__(self):
|
| 354 |
+
pass
|
| 355 |
+
|
| 356 |
+
def from_civitai(self, state_dict):
|
| 357 |
+
return state_dict
|
diffsynth/models/qwen_image_text_encoder.py
ADDED
|
@@ -0,0 +1,255 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from transformers import Qwen2_5_VLModel
|
| 2 |
+
import torch
|
| 3 |
+
from typing import Optional, Union
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class QwenImageTextEncoder(torch.nn.Module):
|
| 7 |
+
def __init__(self):
|
| 8 |
+
super().__init__()
|
| 9 |
+
from transformers import Qwen2_5_VLConfig
|
| 10 |
+
config = Qwen2_5_VLConfig(**{
|
| 11 |
+
"architectures": [
|
| 12 |
+
"Qwen2_5_VLForConditionalGeneration"
|
| 13 |
+
],
|
| 14 |
+
"attention_dropout": 0.0,
|
| 15 |
+
"bos_token_id": 151643,
|
| 16 |
+
"eos_token_id": 151645,
|
| 17 |
+
"hidden_act": "silu",
|
| 18 |
+
"hidden_size": 3584,
|
| 19 |
+
"image_token_id": 151655,
|
| 20 |
+
"initializer_range": 0.02,
|
| 21 |
+
"intermediate_size": 18944,
|
| 22 |
+
"max_position_embeddings": 128000,
|
| 23 |
+
"max_window_layers": 28,
|
| 24 |
+
"model_type": "qwen2_5_vl",
|
| 25 |
+
"num_attention_heads": 28,
|
| 26 |
+
"num_hidden_layers": 28,
|
| 27 |
+
"num_key_value_heads": 4,
|
| 28 |
+
"rms_norm_eps": 1e-06,
|
| 29 |
+
"rope_scaling": {
|
| 30 |
+
"mrope_section": [
|
| 31 |
+
16,
|
| 32 |
+
24,
|
| 33 |
+
24
|
| 34 |
+
],
|
| 35 |
+
"rope_type": "default",
|
| 36 |
+
"type": "default"
|
| 37 |
+
},
|
| 38 |
+
"rope_theta": 1000000.0,
|
| 39 |
+
"sliding_window": 32768,
|
| 40 |
+
"text_config": {
|
| 41 |
+
"architectures": [
|
| 42 |
+
"Qwen2_5_VLForConditionalGeneration"
|
| 43 |
+
],
|
| 44 |
+
"attention_dropout": 0.0,
|
| 45 |
+
"bos_token_id": 151643,
|
| 46 |
+
"eos_token_id": 151645,
|
| 47 |
+
"hidden_act": "silu",
|
| 48 |
+
"hidden_size": 3584,
|
| 49 |
+
"image_token_id": None,
|
| 50 |
+
"initializer_range": 0.02,
|
| 51 |
+
"intermediate_size": 18944,
|
| 52 |
+
"layer_types": [
|
| 53 |
+
"full_attention",
|
| 54 |
+
"full_attention",
|
| 55 |
+
"full_attention",
|
| 56 |
+
"full_attention",
|
| 57 |
+
"full_attention",
|
| 58 |
+
"full_attention",
|
| 59 |
+
"full_attention",
|
| 60 |
+
"full_attention",
|
| 61 |
+
"full_attention",
|
| 62 |
+
"full_attention",
|
| 63 |
+
"full_attention",
|
| 64 |
+
"full_attention",
|
| 65 |
+
"full_attention",
|
| 66 |
+
"full_attention",
|
| 67 |
+
"full_attention",
|
| 68 |
+
"full_attention",
|
| 69 |
+
"full_attention",
|
| 70 |
+
"full_attention",
|
| 71 |
+
"full_attention",
|
| 72 |
+
"full_attention",
|
| 73 |
+
"full_attention",
|
| 74 |
+
"full_attention",
|
| 75 |
+
"full_attention",
|
| 76 |
+
"full_attention",
|
| 77 |
+
"full_attention",
|
| 78 |
+
"full_attention",
|
| 79 |
+
"full_attention",
|
| 80 |
+
"full_attention"
|
| 81 |
+
],
|
| 82 |
+
"max_position_embeddings": 128000,
|
| 83 |
+
"max_window_layers": 28,
|
| 84 |
+
"model_type": "qwen2_5_vl_text",
|
| 85 |
+
"num_attention_heads": 28,
|
| 86 |
+
"num_hidden_layers": 28,
|
| 87 |
+
"num_key_value_heads": 4,
|
| 88 |
+
"rms_norm_eps": 1e-06,
|
| 89 |
+
"rope_scaling": {
|
| 90 |
+
"mrope_section": [
|
| 91 |
+
16,
|
| 92 |
+
24,
|
| 93 |
+
24
|
| 94 |
+
],
|
| 95 |
+
"rope_type": "default",
|
| 96 |
+
"type": "default"
|
| 97 |
+
},
|
| 98 |
+
"rope_theta": 1000000.0,
|
| 99 |
+
"sliding_window": None,
|
| 100 |
+
"torch_dtype": "float32",
|
| 101 |
+
"use_cache": True,
|
| 102 |
+
"use_sliding_window": False,
|
| 103 |
+
"video_token_id": None,
|
| 104 |
+
"vision_end_token_id": 151653,
|
| 105 |
+
"vision_start_token_id": 151652,
|
| 106 |
+
"vision_token_id": 151654,
|
| 107 |
+
"vocab_size": 152064
|
| 108 |
+
},
|
| 109 |
+
"tie_word_embeddings": False,
|
| 110 |
+
"torch_dtype": "float32",
|
| 111 |
+
"transformers_version": "4.54.0",
|
| 112 |
+
"use_cache": True,
|
| 113 |
+
"use_sliding_window": False,
|
| 114 |
+
"video_token_id": 151656,
|
| 115 |
+
"vision_config": {
|
| 116 |
+
"depth": 32,
|
| 117 |
+
"fullatt_block_indexes": [
|
| 118 |
+
7,
|
| 119 |
+
15,
|
| 120 |
+
23,
|
| 121 |
+
31
|
| 122 |
+
],
|
| 123 |
+
"hidden_act": "silu",
|
| 124 |
+
"hidden_size": 1280,
|
| 125 |
+
"in_channels": 3,
|
| 126 |
+
"in_chans": 3,
|
| 127 |
+
"initializer_range": 0.02,
|
| 128 |
+
"intermediate_size": 3420,
|
| 129 |
+
"model_type": "qwen2_5_vl",
|
| 130 |
+
"num_heads": 16,
|
| 131 |
+
"out_hidden_size": 3584,
|
| 132 |
+
"patch_size": 14,
|
| 133 |
+
"spatial_merge_size": 2,
|
| 134 |
+
"spatial_patch_size": 14,
|
| 135 |
+
"temporal_patch_size": 2,
|
| 136 |
+
"tokens_per_second": 2,
|
| 137 |
+
"torch_dtype": "float32",
|
| 138 |
+
"window_size": 112
|
| 139 |
+
},
|
| 140 |
+
"vision_end_token_id": 151653,
|
| 141 |
+
"vision_start_token_id": 151652,
|
| 142 |
+
"vision_token_id": 151654,
|
| 143 |
+
"vocab_size": 152064
|
| 144 |
+
})
|
| 145 |
+
self.model = Qwen2_5_VLModel(config)
|
| 146 |
+
self.lm_head = torch.nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
|
| 147 |
+
|
| 148 |
+
def forward(
|
| 149 |
+
self,
|
| 150 |
+
input_ids: torch.LongTensor = None,
|
| 151 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 152 |
+
position_ids: Optional[torch.LongTensor] = None,
|
| 153 |
+
past_key_values = None,
|
| 154 |
+
inputs_embeds: Optional[torch.FloatTensor] = None,
|
| 155 |
+
labels: Optional[torch.LongTensor] = None,
|
| 156 |
+
use_cache: Optional[bool] = None,
|
| 157 |
+
output_attentions: Optional[bool] = None,
|
| 158 |
+
output_hidden_states: Optional[bool] = None,
|
| 159 |
+
pixel_values: Optional[torch.Tensor] = None,
|
| 160 |
+
pixel_values_videos: Optional[torch.FloatTensor] = None,
|
| 161 |
+
image_grid_thw: Optional[torch.LongTensor] = None,
|
| 162 |
+
video_grid_thw: Optional[torch.LongTensor] = None,
|
| 163 |
+
rope_deltas: Optional[torch.LongTensor] = None,
|
| 164 |
+
cache_position: Optional[torch.LongTensor] = None,
|
| 165 |
+
second_per_grid_ts: Optional[torch.Tensor] = None,
|
| 166 |
+
logits_to_keep: Union[int, torch.Tensor] = 0,
|
| 167 |
+
**kwargs,
|
| 168 |
+
):
|
| 169 |
+
r"""
|
| 170 |
+
labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
| 171 |
+
Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
|
| 172 |
+
config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
|
| 173 |
+
(masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
|
| 174 |
+
image_grid_thw (`torch.LongTensor` of shape `(num_images, 3)`, *optional*):
|
| 175 |
+
The temporal, height and width of feature shape of each image in LLM.
|
| 176 |
+
video_grid_thw (`torch.LongTensor` of shape `(num_videos, 3)`, *optional*):
|
| 177 |
+
The temporal, height and width of feature shape of each video in LLM.
|
| 178 |
+
rope_deltas (`torch.LongTensor` of shape `(batch_size, )`, *optional*):
|
| 179 |
+
The rope index difference between sequence length and multimodal rope.
|
| 180 |
+
second_per_grid_ts (`torch.Tensor` of shape `(num_videos)`, *optional*):
|
| 181 |
+
The time interval (in seconds) for each grid along the temporal dimension in the 3D position IDs.
|
| 182 |
+
|
| 183 |
+
Example:
|
| 184 |
+
|
| 185 |
+
```python
|
| 186 |
+
>>> from PIL import Image
|
| 187 |
+
>>> import requests
|
| 188 |
+
>>> from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration
|
| 189 |
+
|
| 190 |
+
>>> model = Qwen2_5_VLForConditionalGeneration.from_pretrained("Qwen/Qwen2.5-VL-7B-Instruct")
|
| 191 |
+
>>> processor = AutoProcessor.from_pretrained("Qwen/Qwen2.5-VL-7B-Instruct")
|
| 192 |
+
|
| 193 |
+
>>> messages = [
|
| 194 |
+
{
|
| 195 |
+
"role": "user",
|
| 196 |
+
"content": [
|
| 197 |
+
{"type": "image"},
|
| 198 |
+
{"type": "text", "text": "What is shown in this image?"},
|
| 199 |
+
],
|
| 200 |
+
},
|
| 201 |
+
]
|
| 202 |
+
>>> url = "https://www.ilankelman.org/stopsigns/australia.jpg"
|
| 203 |
+
>>> image = Image.open(requests.get(url, stream=True).raw)
|
| 204 |
+
|
| 205 |
+
>>> text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 206 |
+
>>> inputs = processor(text=[text], images=[image], vision_infos=[vision_infos])
|
| 207 |
+
|
| 208 |
+
>>> # Generate
|
| 209 |
+
>>> generate_ids = model.generate(inputs.input_ids, max_length=30)
|
| 210 |
+
>>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
|
| 211 |
+
"The image shows a street scene with a red stop sign in the foreground. In the background, there is a large red gate with Chinese characters ..."
|
| 212 |
+
```"""
|
| 213 |
+
|
| 214 |
+
output_attentions = False
|
| 215 |
+
output_hidden_states = True
|
| 216 |
+
|
| 217 |
+
outputs = self.model(
|
| 218 |
+
input_ids=input_ids,
|
| 219 |
+
pixel_values=pixel_values,
|
| 220 |
+
pixel_values_videos=pixel_values_videos,
|
| 221 |
+
image_grid_thw=image_grid_thw,
|
| 222 |
+
video_grid_thw=video_grid_thw,
|
| 223 |
+
second_per_grid_ts=second_per_grid_ts,
|
| 224 |
+
position_ids=position_ids,
|
| 225 |
+
attention_mask=attention_mask,
|
| 226 |
+
past_key_values=past_key_values,
|
| 227 |
+
inputs_embeds=inputs_embeds,
|
| 228 |
+
use_cache=use_cache,
|
| 229 |
+
output_attentions=output_attentions,
|
| 230 |
+
output_hidden_states=output_hidden_states,
|
| 231 |
+
return_dict=True,
|
| 232 |
+
cache_position=cache_position,
|
| 233 |
+
**kwargs,
|
| 234 |
+
)
|
| 235 |
+
return outputs.hidden_states
|
| 236 |
+
|
| 237 |
+
@staticmethod
|
| 238 |
+
def state_dict_converter():
|
| 239 |
+
return QwenImageTextEncoderStateDictConverter()
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
class QwenImageTextEncoderStateDictConverter():
|
| 244 |
+
def __init__(self):
|
| 245 |
+
pass
|
| 246 |
+
|
| 247 |
+
def from_diffusers(self, state_dict):
|
| 248 |
+
state_dict_ = {}
|
| 249 |
+
for k, v in state_dict.items():
|
| 250 |
+
if k.startswith("visual."):
|
| 251 |
+
k = "model." + k
|
| 252 |
+
elif k.startswith("model."):
|
| 253 |
+
k = k.replace("model.", "model.language_model.")
|
| 254 |
+
state_dict_[k] = v
|
| 255 |
+
return state_dict_
|
diffsynth/models/qwen_image_vae.py
ADDED
|
@@ -0,0 +1,736 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from typing import List, Optional, Tuple, Union
|
| 3 |
+
from torch import nn
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
CACHE_T = 2
|
| 7 |
+
|
| 8 |
+
class QwenImageCausalConv3d(torch.nn.Conv3d):
|
| 9 |
+
r"""
|
| 10 |
+
A custom 3D causal convolution layer with feature caching support.
|
| 11 |
+
|
| 12 |
+
This layer extends the standard Conv3D layer by ensuring causality in the time dimension and handling feature
|
| 13 |
+
caching for efficient inference.
|
| 14 |
+
|
| 15 |
+
Args:
|
| 16 |
+
in_channels (int): Number of channels in the input image
|
| 17 |
+
out_channels (int): Number of channels produced by the convolution
|
| 18 |
+
kernel_size (int or tuple): Size of the convolving kernel
|
| 19 |
+
stride (int or tuple, optional): Stride of the convolution. Default: 1
|
| 20 |
+
padding (int or tuple, optional): Zero-padding added to all three sides of the input. Default: 0
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
def __init__(
|
| 24 |
+
self,
|
| 25 |
+
in_channels: int,
|
| 26 |
+
out_channels: int,
|
| 27 |
+
kernel_size: Union[int, Tuple[int, int, int]],
|
| 28 |
+
stride: Union[int, Tuple[int, int, int]] = 1,
|
| 29 |
+
padding: Union[int, Tuple[int, int, int]] = 0,
|
| 30 |
+
) -> None:
|
| 31 |
+
super().__init__(
|
| 32 |
+
in_channels=in_channels,
|
| 33 |
+
out_channels=out_channels,
|
| 34 |
+
kernel_size=kernel_size,
|
| 35 |
+
stride=stride,
|
| 36 |
+
padding=padding,
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
# Set up causal padding
|
| 40 |
+
self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1], 2 * self.padding[0], 0)
|
| 41 |
+
self.padding = (0, 0, 0)
|
| 42 |
+
|
| 43 |
+
def forward(self, x, cache_x=None):
|
| 44 |
+
padding = list(self._padding)
|
| 45 |
+
if cache_x is not None and self._padding[4] > 0:
|
| 46 |
+
cache_x = cache_x.to(x.device)
|
| 47 |
+
x = torch.cat([cache_x, x], dim=2)
|
| 48 |
+
padding[4] -= cache_x.shape[2]
|
| 49 |
+
x = torch.nn.functional.pad(x, padding)
|
| 50 |
+
return super().forward(x)
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class QwenImageRMS_norm(nn.Module):
|
| 55 |
+
r"""
|
| 56 |
+
A custom RMS normalization layer.
|
| 57 |
+
|
| 58 |
+
Args:
|
| 59 |
+
dim (int): The number of dimensions to normalize over.
|
| 60 |
+
channel_first (bool, optional): Whether the input tensor has channels as the first dimension.
|
| 61 |
+
Default is True.
|
| 62 |
+
images (bool, optional): Whether the input represents image data. Default is True.
|
| 63 |
+
bias (bool, optional): Whether to include a learnable bias term. Default is False.
|
| 64 |
+
"""
|
| 65 |
+
|
| 66 |
+
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None:
|
| 67 |
+
super().__init__()
|
| 68 |
+
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
| 69 |
+
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
| 70 |
+
|
| 71 |
+
self.channel_first = channel_first
|
| 72 |
+
self.scale = dim**0.5
|
| 73 |
+
self.gamma = nn.Parameter(torch.ones(shape))
|
| 74 |
+
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
|
| 75 |
+
|
| 76 |
+
def forward(self, x):
|
| 77 |
+
return torch.nn.functional.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class QwenImageResidualBlock(nn.Module):
|
| 82 |
+
r"""
|
| 83 |
+
A custom residual block module.
|
| 84 |
+
|
| 85 |
+
Args:
|
| 86 |
+
in_dim (int): Number of input channels.
|
| 87 |
+
out_dim (int): Number of output channels.
|
| 88 |
+
dropout (float, optional): Dropout rate for the dropout layer. Default is 0.0.
|
| 89 |
+
non_linearity (str, optional): Type of non-linearity to use. Default is "silu".
|
| 90 |
+
"""
|
| 91 |
+
|
| 92 |
+
def __init__(
|
| 93 |
+
self,
|
| 94 |
+
in_dim: int,
|
| 95 |
+
out_dim: int,
|
| 96 |
+
dropout: float = 0.0,
|
| 97 |
+
non_linearity: str = "silu",
|
| 98 |
+
) -> None:
|
| 99 |
+
super().__init__()
|
| 100 |
+
self.in_dim = in_dim
|
| 101 |
+
self.out_dim = out_dim
|
| 102 |
+
self.nonlinearity = torch.nn.SiLU()
|
| 103 |
+
|
| 104 |
+
# layers
|
| 105 |
+
self.norm1 = QwenImageRMS_norm(in_dim, images=False)
|
| 106 |
+
self.conv1 = QwenImageCausalConv3d(in_dim, out_dim, 3, padding=1)
|
| 107 |
+
self.norm2 = QwenImageRMS_norm(out_dim, images=False)
|
| 108 |
+
self.dropout = nn.Dropout(dropout)
|
| 109 |
+
self.conv2 = QwenImageCausalConv3d(out_dim, out_dim, 3, padding=1)
|
| 110 |
+
self.conv_shortcut = QwenImageCausalConv3d(in_dim, out_dim, 1) if in_dim != out_dim else nn.Identity()
|
| 111 |
+
|
| 112 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 113 |
+
# Apply shortcut connection
|
| 114 |
+
h = self.conv_shortcut(x)
|
| 115 |
+
|
| 116 |
+
# First normalization and activation
|
| 117 |
+
x = self.norm1(x)
|
| 118 |
+
x = self.nonlinearity(x)
|
| 119 |
+
|
| 120 |
+
if feat_cache is not None:
|
| 121 |
+
idx = feat_idx[0]
|
| 122 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 123 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 124 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 125 |
+
|
| 126 |
+
x = self.conv1(x, feat_cache[idx])
|
| 127 |
+
feat_cache[idx] = cache_x
|
| 128 |
+
feat_idx[0] += 1
|
| 129 |
+
else:
|
| 130 |
+
x = self.conv1(x)
|
| 131 |
+
|
| 132 |
+
# Second normalization and activation
|
| 133 |
+
x = self.norm2(x)
|
| 134 |
+
x = self.nonlinearity(x)
|
| 135 |
+
|
| 136 |
+
# Dropout
|
| 137 |
+
x = self.dropout(x)
|
| 138 |
+
|
| 139 |
+
if feat_cache is not None:
|
| 140 |
+
idx = feat_idx[0]
|
| 141 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 142 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 143 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 144 |
+
|
| 145 |
+
x = self.conv2(x, feat_cache[idx])
|
| 146 |
+
feat_cache[idx] = cache_x
|
| 147 |
+
feat_idx[0] += 1
|
| 148 |
+
else:
|
| 149 |
+
x = self.conv2(x)
|
| 150 |
+
|
| 151 |
+
# Add residual connection
|
| 152 |
+
return x + h
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
class QwenImageAttentionBlock(nn.Module):
|
| 157 |
+
r"""
|
| 158 |
+
Causal self-attention with a single head.
|
| 159 |
+
|
| 160 |
+
Args:
|
| 161 |
+
dim (int): The number of channels in the input tensor.
|
| 162 |
+
"""
|
| 163 |
+
|
| 164 |
+
def __init__(self, dim):
|
| 165 |
+
super().__init__()
|
| 166 |
+
self.dim = dim
|
| 167 |
+
|
| 168 |
+
# layers
|
| 169 |
+
self.norm = QwenImageRMS_norm(dim)
|
| 170 |
+
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
| 171 |
+
self.proj = nn.Conv2d(dim, dim, 1)
|
| 172 |
+
|
| 173 |
+
def forward(self, x):
|
| 174 |
+
identity = x
|
| 175 |
+
batch_size, channels, time, height, width = x.size()
|
| 176 |
+
|
| 177 |
+
x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * time, channels, height, width)
|
| 178 |
+
x = self.norm(x)
|
| 179 |
+
|
| 180 |
+
# compute query, key, value
|
| 181 |
+
qkv = self.to_qkv(x)
|
| 182 |
+
qkv = qkv.reshape(batch_size * time, 1, channels * 3, -1)
|
| 183 |
+
qkv = qkv.permute(0, 1, 3, 2).contiguous()
|
| 184 |
+
q, k, v = qkv.chunk(3, dim=-1)
|
| 185 |
+
|
| 186 |
+
# apply attention
|
| 187 |
+
x = torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
| 188 |
+
|
| 189 |
+
x = x.squeeze(1).permute(0, 2, 1).reshape(batch_size * time, channels, height, width)
|
| 190 |
+
|
| 191 |
+
# output projection
|
| 192 |
+
x = self.proj(x)
|
| 193 |
+
|
| 194 |
+
# Reshape back: [(b*t), c, h, w] -> [b, c, t, h, w]
|
| 195 |
+
x = x.view(batch_size, time, channels, height, width)
|
| 196 |
+
x = x.permute(0, 2, 1, 3, 4)
|
| 197 |
+
|
| 198 |
+
return x + identity
|
| 199 |
+
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
class QwenImageUpsample(nn.Upsample):
|
| 203 |
+
r"""
|
| 204 |
+
Perform upsampling while ensuring the output tensor has the same data type as the input.
|
| 205 |
+
|
| 206 |
+
Args:
|
| 207 |
+
x (torch.Tensor): Input tensor to be upsampled.
|
| 208 |
+
|
| 209 |
+
Returns:
|
| 210 |
+
torch.Tensor: Upsampled tensor with the same data type as the input.
|
| 211 |
+
"""
|
| 212 |
+
|
| 213 |
+
def forward(self, x):
|
| 214 |
+
return super().forward(x.float()).type_as(x)
|
| 215 |
+
|
| 216 |
+
|
| 217 |
+
|
| 218 |
+
class QwenImageResample(nn.Module):
|
| 219 |
+
r"""
|
| 220 |
+
A custom resampling module for 2D and 3D data.
|
| 221 |
+
|
| 222 |
+
Args:
|
| 223 |
+
dim (int): The number of input/output channels.
|
| 224 |
+
mode (str): The resampling mode. Must be one of:
|
| 225 |
+
- 'none': No resampling (identity operation).
|
| 226 |
+
- 'upsample2d': 2D upsampling with nearest-exact interpolation and convolution.
|
| 227 |
+
- 'upsample3d': 3D upsampling with nearest-exact interpolation, convolution, and causal 3D convolution.
|
| 228 |
+
- 'downsample2d': 2D downsampling with zero-padding and convolution.
|
| 229 |
+
- 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution.
|
| 230 |
+
"""
|
| 231 |
+
|
| 232 |
+
def __init__(self, dim: int, mode: str) -> None:
|
| 233 |
+
super().__init__()
|
| 234 |
+
self.dim = dim
|
| 235 |
+
self.mode = mode
|
| 236 |
+
|
| 237 |
+
# layers
|
| 238 |
+
if mode == "upsample2d":
|
| 239 |
+
self.resample = nn.Sequential(
|
| 240 |
+
QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), nn.Conv2d(dim, dim // 2, 3, padding=1)
|
| 241 |
+
)
|
| 242 |
+
elif mode == "upsample3d":
|
| 243 |
+
self.resample = nn.Sequential(
|
| 244 |
+
QwenImageUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"), nn.Conv2d(dim, dim // 2, 3, padding=1)
|
| 245 |
+
)
|
| 246 |
+
self.time_conv = QwenImageCausalConv3d(dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
|
| 247 |
+
|
| 248 |
+
elif mode == "downsample2d":
|
| 249 |
+
self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
| 250 |
+
elif mode == "downsample3d":
|
| 251 |
+
self.resample = nn.Sequential(nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2)))
|
| 252 |
+
self.time_conv = QwenImageCausalConv3d(dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
|
| 253 |
+
|
| 254 |
+
else:
|
| 255 |
+
self.resample = nn.Identity()
|
| 256 |
+
|
| 257 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 258 |
+
b, c, t, h, w = x.size()
|
| 259 |
+
if self.mode == "upsample3d":
|
| 260 |
+
if feat_cache is not None:
|
| 261 |
+
idx = feat_idx[0]
|
| 262 |
+
if feat_cache[idx] is None:
|
| 263 |
+
feat_cache[idx] = "Rep"
|
| 264 |
+
feat_idx[0] += 1
|
| 265 |
+
else:
|
| 266 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 267 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] != "Rep":
|
| 268 |
+
# cache last frame of last two chunk
|
| 269 |
+
cache_x = torch.cat(
|
| 270 |
+
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2
|
| 271 |
+
)
|
| 272 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx] == "Rep":
|
| 273 |
+
cache_x = torch.cat([torch.zeros_like(cache_x).to(cache_x.device), cache_x], dim=2)
|
| 274 |
+
if feat_cache[idx] == "Rep":
|
| 275 |
+
x = self.time_conv(x)
|
| 276 |
+
else:
|
| 277 |
+
x = self.time_conv(x, feat_cache[idx])
|
| 278 |
+
feat_cache[idx] = cache_x
|
| 279 |
+
feat_idx[0] += 1
|
| 280 |
+
|
| 281 |
+
x = x.reshape(b, 2, c, t, h, w)
|
| 282 |
+
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
|
| 283 |
+
x = x.reshape(b, c, t * 2, h, w)
|
| 284 |
+
t = x.shape[2]
|
| 285 |
+
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
| 286 |
+
x = self.resample(x)
|
| 287 |
+
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
|
| 288 |
+
|
| 289 |
+
if self.mode == "downsample3d":
|
| 290 |
+
if feat_cache is not None:
|
| 291 |
+
idx = feat_idx[0]
|
| 292 |
+
if feat_cache[idx] is None:
|
| 293 |
+
feat_cache[idx] = x.clone()
|
| 294 |
+
feat_idx[0] += 1
|
| 295 |
+
else:
|
| 296 |
+
cache_x = x[:, :, -1:, :, :].clone()
|
| 297 |
+
x = self.time_conv(torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
|
| 298 |
+
feat_cache[idx] = cache_x
|
| 299 |
+
feat_idx[0] += 1
|
| 300 |
+
return x
|
| 301 |
+
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
class QwenImageMidBlock(nn.Module):
|
| 305 |
+
"""
|
| 306 |
+
Middle block for WanVAE encoder and decoder.
|
| 307 |
+
|
| 308 |
+
Args:
|
| 309 |
+
dim (int): Number of input/output channels.
|
| 310 |
+
dropout (float): Dropout rate.
|
| 311 |
+
non_linearity (str): Type of non-linearity to use.
|
| 312 |
+
"""
|
| 313 |
+
|
| 314 |
+
def __init__(self, dim: int, dropout: float = 0.0, non_linearity: str = "silu", num_layers: int = 1):
|
| 315 |
+
super().__init__()
|
| 316 |
+
self.dim = dim
|
| 317 |
+
|
| 318 |
+
# Create the components
|
| 319 |
+
resnets = [QwenImageResidualBlock(dim, dim, dropout, non_linearity)]
|
| 320 |
+
attentions = []
|
| 321 |
+
for _ in range(num_layers):
|
| 322 |
+
attentions.append(QwenImageAttentionBlock(dim))
|
| 323 |
+
resnets.append(QwenImageResidualBlock(dim, dim, dropout, non_linearity))
|
| 324 |
+
self.attentions = nn.ModuleList(attentions)
|
| 325 |
+
self.resnets = nn.ModuleList(resnets)
|
| 326 |
+
|
| 327 |
+
self.gradient_checkpointing = False
|
| 328 |
+
|
| 329 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 330 |
+
# First residual block
|
| 331 |
+
x = self.resnets[0](x, feat_cache, feat_idx)
|
| 332 |
+
|
| 333 |
+
# Process through attention and residual blocks
|
| 334 |
+
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
| 335 |
+
if attn is not None:
|
| 336 |
+
x = attn(x)
|
| 337 |
+
|
| 338 |
+
x = resnet(x, feat_cache, feat_idx)
|
| 339 |
+
|
| 340 |
+
return x
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
|
| 344 |
+
class QwenImageEncoder3d(nn.Module):
|
| 345 |
+
r"""
|
| 346 |
+
A 3D encoder module.
|
| 347 |
+
|
| 348 |
+
Args:
|
| 349 |
+
dim (int): The base number of channels in the first layer.
|
| 350 |
+
z_dim (int): The dimensionality of the latent space.
|
| 351 |
+
dim_mult (list of int): Multipliers for the number of channels in each block.
|
| 352 |
+
num_res_blocks (int): Number of residual blocks in each block.
|
| 353 |
+
attn_scales (list of float): Scales at which to apply attention mechanisms.
|
| 354 |
+
temperal_downsample (list of bool): Whether to downsample temporally in each block.
|
| 355 |
+
dropout (float): Dropout rate for the dropout layers.
|
| 356 |
+
non_linearity (str): Type of non-linearity to use.
|
| 357 |
+
"""
|
| 358 |
+
|
| 359 |
+
def __init__(
|
| 360 |
+
self,
|
| 361 |
+
dim=128,
|
| 362 |
+
z_dim=4,
|
| 363 |
+
dim_mult=[1, 2, 4, 4],
|
| 364 |
+
num_res_blocks=2,
|
| 365 |
+
attn_scales=[],
|
| 366 |
+
temperal_downsample=[True, True, False],
|
| 367 |
+
dropout=0.0,
|
| 368 |
+
non_linearity: str = "silu",
|
| 369 |
+
):
|
| 370 |
+
super().__init__()
|
| 371 |
+
self.dim = dim
|
| 372 |
+
self.z_dim = z_dim
|
| 373 |
+
self.dim_mult = dim_mult
|
| 374 |
+
self.num_res_blocks = num_res_blocks
|
| 375 |
+
self.attn_scales = attn_scales
|
| 376 |
+
self.temperal_downsample = temperal_downsample
|
| 377 |
+
self.nonlinearity = torch.nn.SiLU()
|
| 378 |
+
|
| 379 |
+
# dimensions
|
| 380 |
+
dims = [dim * u for u in [1] + dim_mult]
|
| 381 |
+
scale = 1.0
|
| 382 |
+
|
| 383 |
+
# init block
|
| 384 |
+
self.conv_in = QwenImageCausalConv3d(3, dims[0], 3, padding=1)
|
| 385 |
+
|
| 386 |
+
# downsample blocks
|
| 387 |
+
self.down_blocks = torch.nn.ModuleList([])
|
| 388 |
+
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
| 389 |
+
# residual (+attention) blocks
|
| 390 |
+
for _ in range(num_res_blocks):
|
| 391 |
+
self.down_blocks.append(QwenImageResidualBlock(in_dim, out_dim, dropout))
|
| 392 |
+
if scale in attn_scales:
|
| 393 |
+
self.down_blocks.append(QwenImageAttentionBlock(out_dim))
|
| 394 |
+
in_dim = out_dim
|
| 395 |
+
|
| 396 |
+
# downsample block
|
| 397 |
+
if i != len(dim_mult) - 1:
|
| 398 |
+
mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
|
| 399 |
+
self.down_blocks.append(QwenImageResample(out_dim, mode=mode))
|
| 400 |
+
scale /= 2.0
|
| 401 |
+
|
| 402 |
+
# middle blocks
|
| 403 |
+
self.mid_block = QwenImageMidBlock(out_dim, dropout, non_linearity, num_layers=1)
|
| 404 |
+
|
| 405 |
+
# output blocks
|
| 406 |
+
self.norm_out = QwenImageRMS_norm(out_dim, images=False)
|
| 407 |
+
self.conv_out = QwenImageCausalConv3d(out_dim, z_dim, 3, padding=1)
|
| 408 |
+
|
| 409 |
+
self.gradient_checkpointing = False
|
| 410 |
+
|
| 411 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 412 |
+
if feat_cache is not None:
|
| 413 |
+
idx = feat_idx[0]
|
| 414 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 415 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 416 |
+
# cache last frame of last two chunk
|
| 417 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 418 |
+
x = self.conv_in(x, feat_cache[idx])
|
| 419 |
+
feat_cache[idx] = cache_x
|
| 420 |
+
feat_idx[0] += 1
|
| 421 |
+
else:
|
| 422 |
+
x = self.conv_in(x)
|
| 423 |
+
|
| 424 |
+
## downsamples
|
| 425 |
+
for layer in self.down_blocks:
|
| 426 |
+
if feat_cache is not None:
|
| 427 |
+
x = layer(x, feat_cache, feat_idx)
|
| 428 |
+
else:
|
| 429 |
+
x = layer(x)
|
| 430 |
+
|
| 431 |
+
## middle
|
| 432 |
+
x = self.mid_block(x, feat_cache, feat_idx)
|
| 433 |
+
|
| 434 |
+
## head
|
| 435 |
+
x = self.norm_out(x)
|
| 436 |
+
x = self.nonlinearity(x)
|
| 437 |
+
if feat_cache is not None:
|
| 438 |
+
idx = feat_idx[0]
|
| 439 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 440 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 441 |
+
# cache last frame of last two chunk
|
| 442 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 443 |
+
x = self.conv_out(x, feat_cache[idx])
|
| 444 |
+
feat_cache[idx] = cache_x
|
| 445 |
+
feat_idx[0] += 1
|
| 446 |
+
else:
|
| 447 |
+
x = self.conv_out(x)
|
| 448 |
+
return x
|
| 449 |
+
|
| 450 |
+
|
| 451 |
+
|
| 452 |
+
class QwenImageUpBlock(nn.Module):
|
| 453 |
+
"""
|
| 454 |
+
A block that handles upsampling for the WanVAE decoder.
|
| 455 |
+
|
| 456 |
+
Args:
|
| 457 |
+
in_dim (int): Input dimension
|
| 458 |
+
out_dim (int): Output dimension
|
| 459 |
+
num_res_blocks (int): Number of residual blocks
|
| 460 |
+
dropout (float): Dropout rate
|
| 461 |
+
upsample_mode (str, optional): Mode for upsampling ('upsample2d' or 'upsample3d')
|
| 462 |
+
non_linearity (str): Type of non-linearity to use
|
| 463 |
+
"""
|
| 464 |
+
|
| 465 |
+
def __init__(
|
| 466 |
+
self,
|
| 467 |
+
in_dim: int,
|
| 468 |
+
out_dim: int,
|
| 469 |
+
num_res_blocks: int,
|
| 470 |
+
dropout: float = 0.0,
|
| 471 |
+
upsample_mode: Optional[str] = None,
|
| 472 |
+
non_linearity: str = "silu",
|
| 473 |
+
):
|
| 474 |
+
super().__init__()
|
| 475 |
+
self.in_dim = in_dim
|
| 476 |
+
self.out_dim = out_dim
|
| 477 |
+
|
| 478 |
+
# Create layers list
|
| 479 |
+
resnets = []
|
| 480 |
+
# Add residual blocks and attention if needed
|
| 481 |
+
current_dim = in_dim
|
| 482 |
+
for _ in range(num_res_blocks + 1):
|
| 483 |
+
resnets.append(QwenImageResidualBlock(current_dim, out_dim, dropout, non_linearity))
|
| 484 |
+
current_dim = out_dim
|
| 485 |
+
|
| 486 |
+
self.resnets = nn.ModuleList(resnets)
|
| 487 |
+
|
| 488 |
+
# Add upsampling layer if needed
|
| 489 |
+
self.upsamplers = None
|
| 490 |
+
if upsample_mode is not None:
|
| 491 |
+
self.upsamplers = nn.ModuleList([QwenImageResample(out_dim, mode=upsample_mode)])
|
| 492 |
+
|
| 493 |
+
self.gradient_checkpointing = False
|
| 494 |
+
|
| 495 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 496 |
+
"""
|
| 497 |
+
Forward pass through the upsampling block.
|
| 498 |
+
|
| 499 |
+
Args:
|
| 500 |
+
x (torch.Tensor): Input tensor
|
| 501 |
+
feat_cache (list, optional): Feature cache for causal convolutions
|
| 502 |
+
feat_idx (list, optional): Feature index for cache management
|
| 503 |
+
|
| 504 |
+
Returns:
|
| 505 |
+
torch.Tensor: Output tensor
|
| 506 |
+
"""
|
| 507 |
+
for resnet in self.resnets:
|
| 508 |
+
if feat_cache is not None:
|
| 509 |
+
x = resnet(x, feat_cache, feat_idx)
|
| 510 |
+
else:
|
| 511 |
+
x = resnet(x)
|
| 512 |
+
|
| 513 |
+
if self.upsamplers is not None:
|
| 514 |
+
if feat_cache is not None:
|
| 515 |
+
x = self.upsamplers[0](x, feat_cache, feat_idx)
|
| 516 |
+
else:
|
| 517 |
+
x = self.upsamplers[0](x)
|
| 518 |
+
return x
|
| 519 |
+
|
| 520 |
+
|
| 521 |
+
|
| 522 |
+
class QwenImageDecoder3d(nn.Module):
|
| 523 |
+
r"""
|
| 524 |
+
A 3D decoder module.
|
| 525 |
+
|
| 526 |
+
Args:
|
| 527 |
+
dim (int): The base number of channels in the first layer.
|
| 528 |
+
z_dim (int): The dimensionality of the latent space.
|
| 529 |
+
dim_mult (list of int): Multipliers for the number of channels in each block.
|
| 530 |
+
num_res_blocks (int): Number of residual blocks in each block.
|
| 531 |
+
attn_scales (list of float): Scales at which to apply attention mechanisms.
|
| 532 |
+
temperal_upsample (list of bool): Whether to upsample temporally in each block.
|
| 533 |
+
dropout (float): Dropout rate for the dropout layers.
|
| 534 |
+
non_linearity (str): Type of non-linearity to use.
|
| 535 |
+
"""
|
| 536 |
+
|
| 537 |
+
def __init__(
|
| 538 |
+
self,
|
| 539 |
+
dim=128,
|
| 540 |
+
z_dim=4,
|
| 541 |
+
dim_mult=[1, 2, 4, 4],
|
| 542 |
+
num_res_blocks=2,
|
| 543 |
+
attn_scales=[],
|
| 544 |
+
temperal_upsample=[False, True, True],
|
| 545 |
+
dropout=0.0,
|
| 546 |
+
non_linearity: str = "silu",
|
| 547 |
+
):
|
| 548 |
+
super().__init__()
|
| 549 |
+
self.dim = dim
|
| 550 |
+
self.z_dim = z_dim
|
| 551 |
+
self.dim_mult = dim_mult
|
| 552 |
+
self.num_res_blocks = num_res_blocks
|
| 553 |
+
self.attn_scales = attn_scales
|
| 554 |
+
self.temperal_upsample = temperal_upsample
|
| 555 |
+
|
| 556 |
+
self.nonlinearity = torch.nn.SiLU()
|
| 557 |
+
|
| 558 |
+
# dimensions
|
| 559 |
+
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
| 560 |
+
scale = 1.0 / 2 ** (len(dim_mult) - 2)
|
| 561 |
+
|
| 562 |
+
# init block
|
| 563 |
+
self.conv_in = QwenImageCausalConv3d(z_dim, dims[0], 3, padding=1)
|
| 564 |
+
|
| 565 |
+
# middle blocks
|
| 566 |
+
self.mid_block = QwenImageMidBlock(dims[0], dropout, non_linearity, num_layers=1)
|
| 567 |
+
|
| 568 |
+
# upsample blocks
|
| 569 |
+
self.up_blocks = nn.ModuleList([])
|
| 570 |
+
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
| 571 |
+
# residual (+attention) blocks
|
| 572 |
+
if i > 0:
|
| 573 |
+
in_dim = in_dim // 2
|
| 574 |
+
|
| 575 |
+
# Determine if we need upsampling
|
| 576 |
+
upsample_mode = None
|
| 577 |
+
if i != len(dim_mult) - 1:
|
| 578 |
+
upsample_mode = "upsample3d" if temperal_upsample[i] else "upsample2d"
|
| 579 |
+
|
| 580 |
+
# Create and add the upsampling block
|
| 581 |
+
up_block = QwenImageUpBlock(
|
| 582 |
+
in_dim=in_dim,
|
| 583 |
+
out_dim=out_dim,
|
| 584 |
+
num_res_blocks=num_res_blocks,
|
| 585 |
+
dropout=dropout,
|
| 586 |
+
upsample_mode=upsample_mode,
|
| 587 |
+
non_linearity=non_linearity,
|
| 588 |
+
)
|
| 589 |
+
self.up_blocks.append(up_block)
|
| 590 |
+
|
| 591 |
+
# Update scale for next iteration
|
| 592 |
+
if upsample_mode is not None:
|
| 593 |
+
scale *= 2.0
|
| 594 |
+
|
| 595 |
+
# output blocks
|
| 596 |
+
self.norm_out = QwenImageRMS_norm(out_dim, images=False)
|
| 597 |
+
self.conv_out = QwenImageCausalConv3d(out_dim, 3, 3, padding=1)
|
| 598 |
+
|
| 599 |
+
self.gradient_checkpointing = False
|
| 600 |
+
|
| 601 |
+
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
| 602 |
+
## conv1
|
| 603 |
+
if feat_cache is not None:
|
| 604 |
+
idx = feat_idx[0]
|
| 605 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 606 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 607 |
+
# cache last frame of last two chunk
|
| 608 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 609 |
+
x = self.conv_in(x, feat_cache[idx])
|
| 610 |
+
feat_cache[idx] = cache_x
|
| 611 |
+
feat_idx[0] += 1
|
| 612 |
+
else:
|
| 613 |
+
x = self.conv_in(x)
|
| 614 |
+
|
| 615 |
+
## middle
|
| 616 |
+
x = self.mid_block(x, feat_cache, feat_idx)
|
| 617 |
+
|
| 618 |
+
## upsamples
|
| 619 |
+
for up_block in self.up_blocks:
|
| 620 |
+
x = up_block(x, feat_cache, feat_idx)
|
| 621 |
+
|
| 622 |
+
## head
|
| 623 |
+
x = self.norm_out(x)
|
| 624 |
+
x = self.nonlinearity(x)
|
| 625 |
+
if feat_cache is not None:
|
| 626 |
+
idx = feat_idx[0]
|
| 627 |
+
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
| 628 |
+
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
| 629 |
+
# cache last frame of last two chunk
|
| 630 |
+
cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
|
| 631 |
+
x = self.conv_out(x, feat_cache[idx])
|
| 632 |
+
feat_cache[idx] = cache_x
|
| 633 |
+
feat_idx[0] += 1
|
| 634 |
+
else:
|
| 635 |
+
x = self.conv_out(x)
|
| 636 |
+
return x
|
| 637 |
+
|
| 638 |
+
|
| 639 |
+
|
| 640 |
+
class QwenImageVAE(torch.nn.Module):
|
| 641 |
+
def __init__(
|
| 642 |
+
self,
|
| 643 |
+
base_dim: int = 96,
|
| 644 |
+
z_dim: int = 16,
|
| 645 |
+
dim_mult: Tuple[int] = [1, 2, 4, 4],
|
| 646 |
+
num_res_blocks: int = 2,
|
| 647 |
+
attn_scales: List[float] = [],
|
| 648 |
+
temperal_downsample: List[bool] = [False, True, True],
|
| 649 |
+
dropout: float = 0.0,
|
| 650 |
+
) -> None:
|
| 651 |
+
super().__init__()
|
| 652 |
+
|
| 653 |
+
self.z_dim = z_dim
|
| 654 |
+
self.temperal_downsample = temperal_downsample
|
| 655 |
+
self.temperal_upsample = temperal_downsample[::-1]
|
| 656 |
+
|
| 657 |
+
self.encoder = QwenImageEncoder3d(
|
| 658 |
+
base_dim, z_dim * 2, dim_mult, num_res_blocks, attn_scales, self.temperal_downsample, dropout
|
| 659 |
+
)
|
| 660 |
+
self.quant_conv = QwenImageCausalConv3d(z_dim * 2, z_dim * 2, 1)
|
| 661 |
+
self.post_quant_conv = QwenImageCausalConv3d(z_dim, z_dim, 1)
|
| 662 |
+
|
| 663 |
+
self.decoder = QwenImageDecoder3d(
|
| 664 |
+
base_dim, z_dim, dim_mult, num_res_blocks, attn_scales, self.temperal_upsample, dropout
|
| 665 |
+
)
|
| 666 |
+
|
| 667 |
+
mean = [
|
| 668 |
+
-0.7571,
|
| 669 |
+
-0.7089,
|
| 670 |
+
-0.9113,
|
| 671 |
+
0.1075,
|
| 672 |
+
-0.1745,
|
| 673 |
+
0.9653,
|
| 674 |
+
-0.1517,
|
| 675 |
+
1.5508,
|
| 676 |
+
0.4134,
|
| 677 |
+
-0.0715,
|
| 678 |
+
0.5517,
|
| 679 |
+
-0.3632,
|
| 680 |
+
-0.1922,
|
| 681 |
+
-0.9497,
|
| 682 |
+
0.2503,
|
| 683 |
+
-0.2921,
|
| 684 |
+
]
|
| 685 |
+
std = [
|
| 686 |
+
2.8184,
|
| 687 |
+
1.4541,
|
| 688 |
+
2.3275,
|
| 689 |
+
2.6558,
|
| 690 |
+
1.2196,
|
| 691 |
+
1.7708,
|
| 692 |
+
2.6052,
|
| 693 |
+
2.0743,
|
| 694 |
+
3.2687,
|
| 695 |
+
2.1526,
|
| 696 |
+
2.8652,
|
| 697 |
+
1.5579,
|
| 698 |
+
1.6382,
|
| 699 |
+
1.1253,
|
| 700 |
+
2.8251,
|
| 701 |
+
1.9160,
|
| 702 |
+
]
|
| 703 |
+
self.mean = torch.tensor(mean).view(1, 16, 1, 1, 1)
|
| 704 |
+
self.std = 1 / torch.tensor(std).view(1, 16, 1, 1, 1)
|
| 705 |
+
|
| 706 |
+
def encode(self, x, **kwargs):
|
| 707 |
+
x = x.unsqueeze(2)
|
| 708 |
+
x = self.encoder(x)
|
| 709 |
+
x = self.quant_conv(x)
|
| 710 |
+
x = x[:, :16]
|
| 711 |
+
mean, std = self.mean.to(dtype=x.dtype, device=x.device), self.std.to(dtype=x.dtype, device=x.device)
|
| 712 |
+
x = (x - mean) * std
|
| 713 |
+
x = x.squeeze(2)
|
| 714 |
+
return x
|
| 715 |
+
|
| 716 |
+
def decode(self, x, **kwargs):
|
| 717 |
+
x = x.unsqueeze(2)
|
| 718 |
+
mean, std = self.mean.to(dtype=x.dtype, device=x.device), self.std.to(dtype=x.dtype, device=x.device)
|
| 719 |
+
x = x / std + mean
|
| 720 |
+
x = self.post_quant_conv(x)
|
| 721 |
+
x = self.decoder(x)
|
| 722 |
+
x = x.squeeze(2)
|
| 723 |
+
return x
|
| 724 |
+
|
| 725 |
+
@staticmethod
|
| 726 |
+
def state_dict_converter():
|
| 727 |
+
return QwenImageVAEStateDictConverter()
|
| 728 |
+
|
| 729 |
+
|
| 730 |
+
|
| 731 |
+
class QwenImageVAEStateDictConverter():
|
| 732 |
+
def __init__(self):
|
| 733 |
+
pass
|
| 734 |
+
|
| 735 |
+
def from_diffusers(self, state_dict):
|
| 736 |
+
return state_dict
|
diffsynth/models/qwenvl.py
ADDED
|
@@ -0,0 +1,168 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
class Qwen25VL_7b_Embedder(torch.nn.Module):
|
| 5 |
+
def __init__(self, model_path, max_length=640, dtype=torch.bfloat16, device="cuda"):
|
| 6 |
+
super(Qwen25VL_7b_Embedder, self).__init__()
|
| 7 |
+
self.max_length = max_length
|
| 8 |
+
self.dtype = dtype
|
| 9 |
+
self.device = device
|
| 10 |
+
|
| 11 |
+
from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration
|
| 12 |
+
|
| 13 |
+
self.model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
| 14 |
+
model_path,
|
| 15 |
+
torch_dtype=dtype,
|
| 16 |
+
).to(torch.cuda.current_device())
|
| 17 |
+
|
| 18 |
+
self.model.requires_grad_(False)
|
| 19 |
+
self.processor = AutoProcessor.from_pretrained(
|
| 20 |
+
model_path, min_pixels=256 * 28 * 28, max_pixels=324 * 28 * 28
|
| 21 |
+
)
|
| 22 |
+
|
| 23 |
+
Qwen25VL_7b_PREFIX = '''Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:
|
| 24 |
+
- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.
|
| 25 |
+
- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.\n
|
| 26 |
+
Here are examples of how to transform or refine prompts:
|
| 27 |
+
- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.
|
| 28 |
+
- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.\n
|
| 29 |
+
Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:
|
| 30 |
+
User Prompt:'''
|
| 31 |
+
|
| 32 |
+
self.prefix = Qwen25VL_7b_PREFIX
|
| 33 |
+
|
| 34 |
+
@staticmethod
|
| 35 |
+
def from_pretrained(path, torch_dtype=torch.bfloat16, device="cuda"):
|
| 36 |
+
return Qwen25VL_7b_Embedder(path, dtype=torch_dtype, device=device)
|
| 37 |
+
|
| 38 |
+
def forward(self, caption, ref_images):
|
| 39 |
+
text_list = caption
|
| 40 |
+
embs = torch.zeros(
|
| 41 |
+
len(text_list),
|
| 42 |
+
self.max_length,
|
| 43 |
+
self.model.config.hidden_size,
|
| 44 |
+
dtype=torch.bfloat16,
|
| 45 |
+
device=torch.cuda.current_device(),
|
| 46 |
+
)
|
| 47 |
+
hidden_states = torch.zeros(
|
| 48 |
+
len(text_list),
|
| 49 |
+
self.max_length,
|
| 50 |
+
self.model.config.hidden_size,
|
| 51 |
+
dtype=torch.bfloat16,
|
| 52 |
+
device=torch.cuda.current_device(),
|
| 53 |
+
)
|
| 54 |
+
masks = torch.zeros(
|
| 55 |
+
len(text_list),
|
| 56 |
+
self.max_length,
|
| 57 |
+
dtype=torch.long,
|
| 58 |
+
device=torch.cuda.current_device(),
|
| 59 |
+
)
|
| 60 |
+
input_ids_list = []
|
| 61 |
+
attention_mask_list = []
|
| 62 |
+
emb_list = []
|
| 63 |
+
|
| 64 |
+
def split_string(s):
|
| 65 |
+
s = s.replace("“", '"').replace("”", '"').replace("'", '''"''') # use english quotes
|
| 66 |
+
result = []
|
| 67 |
+
in_quotes = False
|
| 68 |
+
temp = ""
|
| 69 |
+
|
| 70 |
+
for idx,char in enumerate(s):
|
| 71 |
+
if char == '"' and idx>155:
|
| 72 |
+
temp += char
|
| 73 |
+
if not in_quotes:
|
| 74 |
+
result.append(temp)
|
| 75 |
+
temp = ""
|
| 76 |
+
|
| 77 |
+
in_quotes = not in_quotes
|
| 78 |
+
continue
|
| 79 |
+
if in_quotes:
|
| 80 |
+
if char.isspace():
|
| 81 |
+
pass # have space token
|
| 82 |
+
|
| 83 |
+
result.append("“" + char + "”")
|
| 84 |
+
else:
|
| 85 |
+
temp += char
|
| 86 |
+
|
| 87 |
+
if temp:
|
| 88 |
+
result.append(temp)
|
| 89 |
+
|
| 90 |
+
return result
|
| 91 |
+
|
| 92 |
+
for idx, (txt, imgs) in enumerate(zip(text_list, ref_images)):
|
| 93 |
+
|
| 94 |
+
messages = [{"role": "user", "content": []}]
|
| 95 |
+
|
| 96 |
+
messages[0]["content"].append({"type": "text", "text": f"{self.prefix}"})
|
| 97 |
+
|
| 98 |
+
messages[0]["content"].append({"type": "image", "image": imgs})
|
| 99 |
+
|
| 100 |
+
# 再添加 text
|
| 101 |
+
messages[0]["content"].append({"type": "text", "text": f"{txt}"})
|
| 102 |
+
|
| 103 |
+
# Preparation for inference
|
| 104 |
+
text = self.processor.apply_chat_template(
|
| 105 |
+
messages, tokenize=False, add_generation_prompt=True, add_vision_id=True
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
image_inputs = [imgs]
|
| 109 |
+
|
| 110 |
+
inputs = self.processor(
|
| 111 |
+
text=[text],
|
| 112 |
+
images=image_inputs,
|
| 113 |
+
padding=True,
|
| 114 |
+
return_tensors="pt",
|
| 115 |
+
)
|
| 116 |
+
|
| 117 |
+
old_inputs_ids = inputs.input_ids
|
| 118 |
+
text_split_list = split_string(text)
|
| 119 |
+
|
| 120 |
+
token_list = []
|
| 121 |
+
for text_each in text_split_list:
|
| 122 |
+
txt_inputs = self.processor(
|
| 123 |
+
text=text_each,
|
| 124 |
+
images=None,
|
| 125 |
+
videos=None,
|
| 126 |
+
padding=True,
|
| 127 |
+
return_tensors="pt",
|
| 128 |
+
)
|
| 129 |
+
token_each = txt_inputs.input_ids
|
| 130 |
+
if token_each[0][0] == 2073 and token_each[0][-1] == 854:
|
| 131 |
+
token_each = token_each[:, 1:-1]
|
| 132 |
+
token_list.append(token_each)
|
| 133 |
+
else:
|
| 134 |
+
token_list.append(token_each)
|
| 135 |
+
|
| 136 |
+
new_txt_ids = torch.cat(token_list, dim=1).to("cuda")
|
| 137 |
+
|
| 138 |
+
new_txt_ids = new_txt_ids.to(old_inputs_ids.device)
|
| 139 |
+
|
| 140 |
+
idx1 = (old_inputs_ids == 151653).nonzero(as_tuple=True)[1][0]
|
| 141 |
+
idx2 = (new_txt_ids == 151653).nonzero(as_tuple=True)[1][0]
|
| 142 |
+
inputs.input_ids = (
|
| 143 |
+
torch.cat([old_inputs_ids[0, :idx1], new_txt_ids[0, idx2:]], dim=0)
|
| 144 |
+
.unsqueeze(0)
|
| 145 |
+
.to("cuda")
|
| 146 |
+
)
|
| 147 |
+
inputs.attention_mask = (inputs.input_ids > 0).long().to("cuda")
|
| 148 |
+
outputs = self.model(
|
| 149 |
+
input_ids=inputs.input_ids,
|
| 150 |
+
attention_mask=inputs.attention_mask,
|
| 151 |
+
pixel_values=inputs.pixel_values.to("cuda"),
|
| 152 |
+
image_grid_thw=inputs.image_grid_thw.to("cuda"),
|
| 153 |
+
output_hidden_states=True,
|
| 154 |
+
)
|
| 155 |
+
|
| 156 |
+
emb = outputs["hidden_states"][-1]
|
| 157 |
+
|
| 158 |
+
embs[idx, : min(self.max_length, emb.shape[1] - 217)] = emb[0, 217:][
|
| 159 |
+
: self.max_length
|
| 160 |
+
]
|
| 161 |
+
|
| 162 |
+
masks[idx, : min(self.max_length, emb.shape[1] - 217)] = torch.ones(
|
| 163 |
+
(min(self.max_length, emb.shape[1] - 217)),
|
| 164 |
+
dtype=torch.long,
|
| 165 |
+
device=torch.cuda.current_device(),
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
return embs, masks
|
diffsynth/models/sd3_dit.py
ADDED
|
@@ -0,0 +1,567 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from einops import rearrange
|
| 3 |
+
from .svd_unet import TemporalTimesteps
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class RMSNorm(torch.nn.Module):
|
| 9 |
+
def __init__(self, dim, eps, elementwise_affine=True):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.eps = eps
|
| 12 |
+
if elementwise_affine:
|
| 13 |
+
self.weight = torch.nn.Parameter(torch.ones((dim,)))
|
| 14 |
+
else:
|
| 15 |
+
self.weight = None
|
| 16 |
+
|
| 17 |
+
def forward(self, hidden_states):
|
| 18 |
+
input_dtype = hidden_states.dtype
|
| 19 |
+
variance = hidden_states.to(torch.float32).square().mean(-1, keepdim=True)
|
| 20 |
+
hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
|
| 21 |
+
hidden_states = hidden_states.to(input_dtype)
|
| 22 |
+
if self.weight is not None:
|
| 23 |
+
hidden_states = hidden_states * self.weight
|
| 24 |
+
return hidden_states
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class PatchEmbed(torch.nn.Module):
|
| 29 |
+
def __init__(self, patch_size=2, in_channels=16, embed_dim=1536, pos_embed_max_size=192):
|
| 30 |
+
super().__init__()
|
| 31 |
+
self.pos_embed_max_size = pos_embed_max_size
|
| 32 |
+
self.patch_size = patch_size
|
| 33 |
+
|
| 34 |
+
self.proj = torch.nn.Conv2d(in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size)
|
| 35 |
+
self.pos_embed = torch.nn.Parameter(torch.zeros(1, self.pos_embed_max_size, self.pos_embed_max_size, embed_dim))
|
| 36 |
+
|
| 37 |
+
def cropped_pos_embed(self, height, width):
|
| 38 |
+
height = height // self.patch_size
|
| 39 |
+
width = width // self.patch_size
|
| 40 |
+
top = (self.pos_embed_max_size - height) // 2
|
| 41 |
+
left = (self.pos_embed_max_size - width) // 2
|
| 42 |
+
spatial_pos_embed = self.pos_embed[:, top : top + height, left : left + width, :].flatten(1, 2)
|
| 43 |
+
return spatial_pos_embed
|
| 44 |
+
|
| 45 |
+
def forward(self, latent):
|
| 46 |
+
height, width = latent.shape[-2:]
|
| 47 |
+
latent = self.proj(latent)
|
| 48 |
+
latent = latent.flatten(2).transpose(1, 2)
|
| 49 |
+
pos_embed = self.cropped_pos_embed(height, width)
|
| 50 |
+
return latent + pos_embed
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class DiffusersCompatibleTimestepProj(torch.nn.Module):
|
| 54 |
+
def __init__(self, dim_in, dim_out):
|
| 55 |
+
super().__init__()
|
| 56 |
+
self.linear_1 = torch.nn.Linear(dim_in, dim_out)
|
| 57 |
+
self.act = torch.nn.SiLU()
|
| 58 |
+
self.linear_2 = torch.nn.Linear(dim_out, dim_out)
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
x = self.linear_1(x)
|
| 62 |
+
x = self.act(x)
|
| 63 |
+
x = self.linear_2(x)
|
| 64 |
+
return x
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class TimestepEmbeddings(torch.nn.Module):
|
| 68 |
+
def __init__(self, dim_in, dim_out, computation_device=None, diffusers_compatible_format=False, scale=1, align_dtype_to_timestep=False):
|
| 69 |
+
super().__init__()
|
| 70 |
+
self.time_proj = TemporalTimesteps(num_channels=dim_in, flip_sin_to_cos=True, downscale_freq_shift=0, computation_device=computation_device, scale=scale, align_dtype_to_timestep=align_dtype_to_timestep)
|
| 71 |
+
if diffusers_compatible_format:
|
| 72 |
+
self.timestep_embedder = DiffusersCompatibleTimestepProj(dim_in, dim_out)
|
| 73 |
+
else:
|
| 74 |
+
self.timestep_embedder = torch.nn.Sequential(
|
| 75 |
+
torch.nn.Linear(dim_in, dim_out), torch.nn.SiLU(), torch.nn.Linear(dim_out, dim_out)
|
| 76 |
+
)
|
| 77 |
+
|
| 78 |
+
def forward(self, timestep, dtype):
|
| 79 |
+
time_emb = self.time_proj(timestep).to(dtype)
|
| 80 |
+
time_emb = self.timestep_embedder(time_emb)
|
| 81 |
+
return time_emb
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
class AdaLayerNorm(torch.nn.Module):
|
| 86 |
+
def __init__(self, dim, single=False, dual=False):
|
| 87 |
+
super().__init__()
|
| 88 |
+
self.single = single
|
| 89 |
+
self.dual = dual
|
| 90 |
+
self.linear = torch.nn.Linear(dim, dim * [[6, 2][single], 9][dual])
|
| 91 |
+
self.norm = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 92 |
+
|
| 93 |
+
def forward(self, x, emb):
|
| 94 |
+
emb = self.linear(torch.nn.functional.silu(emb))
|
| 95 |
+
if self.single:
|
| 96 |
+
scale, shift = emb.unsqueeze(1).chunk(2, dim=2)
|
| 97 |
+
x = self.norm(x) * (1 + scale) + shift
|
| 98 |
+
return x
|
| 99 |
+
elif self.dual:
|
| 100 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp, shift_msa2, scale_msa2, gate_msa2 = emb.unsqueeze(1).chunk(9, dim=2)
|
| 101 |
+
norm_x = self.norm(x)
|
| 102 |
+
x = norm_x * (1 + scale_msa) + shift_msa
|
| 103 |
+
norm_x2 = norm_x * (1 + scale_msa2) + shift_msa2
|
| 104 |
+
return x, gate_msa, shift_mlp, scale_mlp, gate_mlp, norm_x2, gate_msa2
|
| 105 |
+
else:
|
| 106 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.unsqueeze(1).chunk(6, dim=2)
|
| 107 |
+
x = self.norm(x) * (1 + scale_msa) + shift_msa
|
| 108 |
+
return x, gate_msa, shift_mlp, scale_mlp, gate_mlp
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
class JointAttention(torch.nn.Module):
|
| 113 |
+
def __init__(self, dim_a, dim_b, num_heads, head_dim, only_out_a=False, use_rms_norm=False):
|
| 114 |
+
super().__init__()
|
| 115 |
+
self.num_heads = num_heads
|
| 116 |
+
self.head_dim = head_dim
|
| 117 |
+
self.only_out_a = only_out_a
|
| 118 |
+
|
| 119 |
+
self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)
|
| 120 |
+
self.b_to_qkv = torch.nn.Linear(dim_b, dim_b * 3)
|
| 121 |
+
|
| 122 |
+
self.a_to_out = torch.nn.Linear(dim_a, dim_a)
|
| 123 |
+
if not only_out_a:
|
| 124 |
+
self.b_to_out = torch.nn.Linear(dim_b, dim_b)
|
| 125 |
+
|
| 126 |
+
if use_rms_norm:
|
| 127 |
+
self.norm_q_a = RMSNorm(head_dim, eps=1e-6)
|
| 128 |
+
self.norm_k_a = RMSNorm(head_dim, eps=1e-6)
|
| 129 |
+
self.norm_q_b = RMSNorm(head_dim, eps=1e-6)
|
| 130 |
+
self.norm_k_b = RMSNorm(head_dim, eps=1e-6)
|
| 131 |
+
else:
|
| 132 |
+
self.norm_q_a = None
|
| 133 |
+
self.norm_k_a = None
|
| 134 |
+
self.norm_q_b = None
|
| 135 |
+
self.norm_k_b = None
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def process_qkv(self, hidden_states, to_qkv, norm_q, norm_k):
|
| 139 |
+
batch_size = hidden_states.shape[0]
|
| 140 |
+
qkv = to_qkv(hidden_states)
|
| 141 |
+
qkv = qkv.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 142 |
+
q, k, v = qkv.chunk(3, dim=1)
|
| 143 |
+
if norm_q is not None:
|
| 144 |
+
q = norm_q(q)
|
| 145 |
+
if norm_k is not None:
|
| 146 |
+
k = norm_k(k)
|
| 147 |
+
return q, k, v
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def forward(self, hidden_states_a, hidden_states_b):
|
| 151 |
+
batch_size = hidden_states_a.shape[0]
|
| 152 |
+
|
| 153 |
+
qa, ka, va = self.process_qkv(hidden_states_a, self.a_to_qkv, self.norm_q_a, self.norm_k_a)
|
| 154 |
+
qb, kb, vb = self.process_qkv(hidden_states_b, self.b_to_qkv, self.norm_q_b, self.norm_k_b)
|
| 155 |
+
q = torch.concat([qa, qb], dim=2)
|
| 156 |
+
k = torch.concat([ka, kb], dim=2)
|
| 157 |
+
v = torch.concat([va, vb], dim=2)
|
| 158 |
+
|
| 159 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
| 160 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 161 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 162 |
+
hidden_states_a, hidden_states_b = hidden_states[:, :hidden_states_a.shape[1]], hidden_states[:, hidden_states_a.shape[1]:]
|
| 163 |
+
hidden_states_a = self.a_to_out(hidden_states_a)
|
| 164 |
+
if self.only_out_a:
|
| 165 |
+
return hidden_states_a
|
| 166 |
+
else:
|
| 167 |
+
hidden_states_b = self.b_to_out(hidden_states_b)
|
| 168 |
+
return hidden_states_a, hidden_states_b
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
class SingleAttention(torch.nn.Module):
|
| 173 |
+
def __init__(self, dim_a, num_heads, head_dim, use_rms_norm=False):
|
| 174 |
+
super().__init__()
|
| 175 |
+
self.num_heads = num_heads
|
| 176 |
+
self.head_dim = head_dim
|
| 177 |
+
|
| 178 |
+
self.a_to_qkv = torch.nn.Linear(dim_a, dim_a * 3)
|
| 179 |
+
self.a_to_out = torch.nn.Linear(dim_a, dim_a)
|
| 180 |
+
|
| 181 |
+
if use_rms_norm:
|
| 182 |
+
self.norm_q_a = RMSNorm(head_dim, eps=1e-6)
|
| 183 |
+
self.norm_k_a = RMSNorm(head_dim, eps=1e-6)
|
| 184 |
+
else:
|
| 185 |
+
self.norm_q_a = None
|
| 186 |
+
self.norm_k_a = None
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
def process_qkv(self, hidden_states, to_qkv, norm_q, norm_k):
|
| 190 |
+
batch_size = hidden_states.shape[0]
|
| 191 |
+
qkv = to_qkv(hidden_states)
|
| 192 |
+
qkv = qkv.view(batch_size, -1, 3 * self.num_heads, self.head_dim).transpose(1, 2)
|
| 193 |
+
q, k, v = qkv.chunk(3, dim=1)
|
| 194 |
+
if norm_q is not None:
|
| 195 |
+
q = norm_q(q)
|
| 196 |
+
if norm_k is not None:
|
| 197 |
+
k = norm_k(k)
|
| 198 |
+
return q, k, v
|
| 199 |
+
|
| 200 |
+
|
| 201 |
+
def forward(self, hidden_states_a):
|
| 202 |
+
batch_size = hidden_states_a.shape[0]
|
| 203 |
+
q, k, v = self.process_qkv(hidden_states_a, self.a_to_qkv, self.norm_q_a, self.norm_k_a)
|
| 204 |
+
|
| 205 |
+
hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
| 206 |
+
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)
|
| 207 |
+
hidden_states = hidden_states.to(q.dtype)
|
| 208 |
+
hidden_states = self.a_to_out(hidden_states)
|
| 209 |
+
return hidden_states
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
class DualTransformerBlock(torch.nn.Module):
|
| 214 |
+
def __init__(self, dim, num_attention_heads, use_rms_norm=False):
|
| 215 |
+
super().__init__()
|
| 216 |
+
self.norm1_a = AdaLayerNorm(dim, dual=True)
|
| 217 |
+
self.norm1_b = AdaLayerNorm(dim)
|
| 218 |
+
|
| 219 |
+
self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)
|
| 220 |
+
self.attn2 = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)
|
| 221 |
+
|
| 222 |
+
self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 223 |
+
self.ff_a = torch.nn.Sequential(
|
| 224 |
+
torch.nn.Linear(dim, dim*4),
|
| 225 |
+
torch.nn.GELU(approximate="tanh"),
|
| 226 |
+
torch.nn.Linear(dim*4, dim)
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 230 |
+
self.ff_b = torch.nn.Sequential(
|
| 231 |
+
torch.nn.Linear(dim, dim*4),
|
| 232 |
+
torch.nn.GELU(approximate="tanh"),
|
| 233 |
+
torch.nn.Linear(dim*4, dim)
|
| 234 |
+
)
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
def forward(self, hidden_states_a, hidden_states_b, temb):
|
| 238 |
+
norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a, norm_hidden_states_a_2, gate_msa_a_2 = self.norm1_a(hidden_states_a, emb=temb)
|
| 239 |
+
norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)
|
| 240 |
+
|
| 241 |
+
# Attention
|
| 242 |
+
attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b)
|
| 243 |
+
|
| 244 |
+
# Part A
|
| 245 |
+
hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a
|
| 246 |
+
hidden_states_a = hidden_states_a + gate_msa_a_2 * self.attn2(norm_hidden_states_a_2)
|
| 247 |
+
norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a
|
| 248 |
+
hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)
|
| 249 |
+
|
| 250 |
+
# Part B
|
| 251 |
+
hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b
|
| 252 |
+
norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b
|
| 253 |
+
hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)
|
| 254 |
+
|
| 255 |
+
return hidden_states_a, hidden_states_b
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
|
| 259 |
+
class JointTransformerBlock(torch.nn.Module):
|
| 260 |
+
def __init__(self, dim, num_attention_heads, use_rms_norm=False, dual=False):
|
| 261 |
+
super().__init__()
|
| 262 |
+
self.norm1_a = AdaLayerNorm(dim, dual=dual)
|
| 263 |
+
self.norm1_b = AdaLayerNorm(dim)
|
| 264 |
+
|
| 265 |
+
self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)
|
| 266 |
+
if dual:
|
| 267 |
+
self.attn2 = SingleAttention(dim, num_attention_heads, dim // num_attention_heads, use_rms_norm=use_rms_norm)
|
| 268 |
+
|
| 269 |
+
self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 270 |
+
self.ff_a = torch.nn.Sequential(
|
| 271 |
+
torch.nn.Linear(dim, dim*4),
|
| 272 |
+
torch.nn.GELU(approximate="tanh"),
|
| 273 |
+
torch.nn.Linear(dim*4, dim)
|
| 274 |
+
)
|
| 275 |
+
|
| 276 |
+
self.norm2_b = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 277 |
+
self.ff_b = torch.nn.Sequential(
|
| 278 |
+
torch.nn.Linear(dim, dim*4),
|
| 279 |
+
torch.nn.GELU(approximate="tanh"),
|
| 280 |
+
torch.nn.Linear(dim*4, dim)
|
| 281 |
+
)
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def forward(self, hidden_states_a, hidden_states_b, temb):
|
| 285 |
+
if self.norm1_a.dual:
|
| 286 |
+
norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a, norm_hidden_states_a_2, gate_msa_a_2 = self.norm1_a(hidden_states_a, emb=temb)
|
| 287 |
+
else:
|
| 288 |
+
norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)
|
| 289 |
+
norm_hidden_states_b, gate_msa_b, shift_mlp_b, scale_mlp_b, gate_mlp_b = self.norm1_b(hidden_states_b, emb=temb)
|
| 290 |
+
|
| 291 |
+
# Attention
|
| 292 |
+
attn_output_a, attn_output_b = self.attn(norm_hidden_states_a, norm_hidden_states_b)
|
| 293 |
+
|
| 294 |
+
# Part A
|
| 295 |
+
hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a
|
| 296 |
+
if self.norm1_a.dual:
|
| 297 |
+
hidden_states_a = hidden_states_a + gate_msa_a_2 * self.attn2(norm_hidden_states_a_2)
|
| 298 |
+
norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a
|
| 299 |
+
hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)
|
| 300 |
+
|
| 301 |
+
# Part B
|
| 302 |
+
hidden_states_b = hidden_states_b + gate_msa_b * attn_output_b
|
| 303 |
+
norm_hidden_states_b = self.norm2_b(hidden_states_b) * (1 + scale_mlp_b) + shift_mlp_b
|
| 304 |
+
hidden_states_b = hidden_states_b + gate_mlp_b * self.ff_b(norm_hidden_states_b)
|
| 305 |
+
|
| 306 |
+
return hidden_states_a, hidden_states_b
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
|
| 310 |
+
class JointTransformerFinalBlock(torch.nn.Module):
|
| 311 |
+
def __init__(self, dim, num_attention_heads, use_rms_norm=False):
|
| 312 |
+
super().__init__()
|
| 313 |
+
self.norm1_a = AdaLayerNorm(dim)
|
| 314 |
+
self.norm1_b = AdaLayerNorm(dim, single=True)
|
| 315 |
+
|
| 316 |
+
self.attn = JointAttention(dim, dim, num_attention_heads, dim // num_attention_heads, only_out_a=True, use_rms_norm=use_rms_norm)
|
| 317 |
+
|
| 318 |
+
self.norm2_a = torch.nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
|
| 319 |
+
self.ff_a = torch.nn.Sequential(
|
| 320 |
+
torch.nn.Linear(dim, dim*4),
|
| 321 |
+
torch.nn.GELU(approximate="tanh"),
|
| 322 |
+
torch.nn.Linear(dim*4, dim)
|
| 323 |
+
)
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
def forward(self, hidden_states_a, hidden_states_b, temb):
|
| 327 |
+
norm_hidden_states_a, gate_msa_a, shift_mlp_a, scale_mlp_a, gate_mlp_a = self.norm1_a(hidden_states_a, emb=temb)
|
| 328 |
+
norm_hidden_states_b = self.norm1_b(hidden_states_b, emb=temb)
|
| 329 |
+
|
| 330 |
+
# Attention
|
| 331 |
+
attn_output_a = self.attn(norm_hidden_states_a, norm_hidden_states_b)
|
| 332 |
+
|
| 333 |
+
# Part A
|
| 334 |
+
hidden_states_a = hidden_states_a + gate_msa_a * attn_output_a
|
| 335 |
+
norm_hidden_states_a = self.norm2_a(hidden_states_a) * (1 + scale_mlp_a) + shift_mlp_a
|
| 336 |
+
hidden_states_a = hidden_states_a + gate_mlp_a * self.ff_a(norm_hidden_states_a)
|
| 337 |
+
|
| 338 |
+
return hidden_states_a, hidden_states_b
|
| 339 |
+
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
class SD3DiT(torch.nn.Module):
|
| 343 |
+
def __init__(self, embed_dim=1536, num_layers=24, use_rms_norm=False, num_dual_blocks=0, pos_embed_max_size=192):
|
| 344 |
+
super().__init__()
|
| 345 |
+
self.pos_embedder = PatchEmbed(patch_size=2, in_channels=16, embed_dim=embed_dim, pos_embed_max_size=pos_embed_max_size)
|
| 346 |
+
self.time_embedder = TimestepEmbeddings(256, embed_dim)
|
| 347 |
+
self.pooled_text_embedder = torch.nn.Sequential(torch.nn.Linear(2048, embed_dim), torch.nn.SiLU(), torch.nn.Linear(embed_dim, embed_dim))
|
| 348 |
+
self.context_embedder = torch.nn.Linear(4096, embed_dim)
|
| 349 |
+
self.blocks = torch.nn.ModuleList([JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm, dual=True) for _ in range(num_dual_blocks)]
|
| 350 |
+
+ [JointTransformerBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm) for _ in range(num_layers-1-num_dual_blocks)]
|
| 351 |
+
+ [JointTransformerFinalBlock(embed_dim, embed_dim//64, use_rms_norm=use_rms_norm)])
|
| 352 |
+
self.norm_out = AdaLayerNorm(embed_dim, single=True)
|
| 353 |
+
self.proj_out = torch.nn.Linear(embed_dim, 64)
|
| 354 |
+
|
| 355 |
+
def tiled_forward(self, hidden_states, timestep, prompt_emb, pooled_prompt_emb, tile_size=128, tile_stride=64):
|
| 356 |
+
# Due to the global positional embedding, we cannot implement layer-wise tiled forward.
|
| 357 |
+
hidden_states = TileWorker().tiled_forward(
|
| 358 |
+
lambda x: self.forward(x, timestep, prompt_emb, pooled_prompt_emb),
|
| 359 |
+
hidden_states,
|
| 360 |
+
tile_size,
|
| 361 |
+
tile_stride,
|
| 362 |
+
tile_device=hidden_states.device,
|
| 363 |
+
tile_dtype=hidden_states.dtype
|
| 364 |
+
)
|
| 365 |
+
return hidden_states
|
| 366 |
+
|
| 367 |
+
def forward(self, hidden_states, timestep, prompt_emb, pooled_prompt_emb, tiled=False, tile_size=128, tile_stride=64, use_gradient_checkpointing=False):
|
| 368 |
+
if tiled:
|
| 369 |
+
return self.tiled_forward(hidden_states, timestep, prompt_emb, pooled_prompt_emb, tile_size, tile_stride)
|
| 370 |
+
conditioning = self.time_embedder(timestep, hidden_states.dtype) + self.pooled_text_embedder(pooled_prompt_emb)
|
| 371 |
+
prompt_emb = self.context_embedder(prompt_emb)
|
| 372 |
+
|
| 373 |
+
height, width = hidden_states.shape[-2:]
|
| 374 |
+
hidden_states = self.pos_embedder(hidden_states)
|
| 375 |
+
|
| 376 |
+
def create_custom_forward(module):
|
| 377 |
+
def custom_forward(*inputs):
|
| 378 |
+
return module(*inputs)
|
| 379 |
+
return custom_forward
|
| 380 |
+
|
| 381 |
+
for block in self.blocks:
|
| 382 |
+
if self.training and use_gradient_checkpointing:
|
| 383 |
+
hidden_states, prompt_emb = torch.utils.checkpoint.checkpoint(
|
| 384 |
+
create_custom_forward(block),
|
| 385 |
+
hidden_states, prompt_emb, conditioning,
|
| 386 |
+
use_reentrant=False,
|
| 387 |
+
)
|
| 388 |
+
else:
|
| 389 |
+
hidden_states, prompt_emb = block(hidden_states, prompt_emb, conditioning)
|
| 390 |
+
|
| 391 |
+
hidden_states = self.norm_out(hidden_states, conditioning)
|
| 392 |
+
hidden_states = self.proj_out(hidden_states)
|
| 393 |
+
hidden_states = rearrange(hidden_states, "B (H W) (P Q C) -> B C (H P) (W Q)", P=2, Q=2, H=height//2, W=width//2)
|
| 394 |
+
return hidden_states
|
| 395 |
+
|
| 396 |
+
@staticmethod
|
| 397 |
+
def state_dict_converter():
|
| 398 |
+
return SD3DiTStateDictConverter()
|
| 399 |
+
|
| 400 |
+
|
| 401 |
+
|
| 402 |
+
class SD3DiTStateDictConverter:
|
| 403 |
+
def __init__(self):
|
| 404 |
+
pass
|
| 405 |
+
|
| 406 |
+
def infer_architecture(self, state_dict):
|
| 407 |
+
embed_dim = state_dict["blocks.0.ff_a.0.weight"].shape[1]
|
| 408 |
+
num_layers = 100
|
| 409 |
+
while num_layers > 0 and f"blocks.{num_layers-1}.ff_a.0.bias" not in state_dict:
|
| 410 |
+
num_layers -= 1
|
| 411 |
+
use_rms_norm = "blocks.0.attn.norm_q_a.weight" in state_dict
|
| 412 |
+
num_dual_blocks = 0
|
| 413 |
+
while f"blocks.{num_dual_blocks}.attn2.a_to_out.bias" in state_dict:
|
| 414 |
+
num_dual_blocks += 1
|
| 415 |
+
pos_embed_max_size = state_dict["pos_embedder.pos_embed"].shape[1]
|
| 416 |
+
return {
|
| 417 |
+
"embed_dim": embed_dim,
|
| 418 |
+
"num_layers": num_layers,
|
| 419 |
+
"use_rms_norm": use_rms_norm,
|
| 420 |
+
"num_dual_blocks": num_dual_blocks,
|
| 421 |
+
"pos_embed_max_size": pos_embed_max_size
|
| 422 |
+
}
|
| 423 |
+
|
| 424 |
+
def from_diffusers(self, state_dict):
|
| 425 |
+
rename_dict = {
|
| 426 |
+
"context_embedder": "context_embedder",
|
| 427 |
+
"pos_embed.pos_embed": "pos_embedder.pos_embed",
|
| 428 |
+
"pos_embed.proj": "pos_embedder.proj",
|
| 429 |
+
"time_text_embed.timestep_embedder.linear_1": "time_embedder.timestep_embedder.0",
|
| 430 |
+
"time_text_embed.timestep_embedder.linear_2": "time_embedder.timestep_embedder.2",
|
| 431 |
+
"time_text_embed.text_embedder.linear_1": "pooled_text_embedder.0",
|
| 432 |
+
"time_text_embed.text_embedder.linear_2": "pooled_text_embedder.2",
|
| 433 |
+
"norm_out.linear": "norm_out.linear",
|
| 434 |
+
"proj_out": "proj_out",
|
| 435 |
+
|
| 436 |
+
"norm1.linear": "norm1_a.linear",
|
| 437 |
+
"norm1_context.linear": "norm1_b.linear",
|
| 438 |
+
"attn.to_q": "attn.a_to_q",
|
| 439 |
+
"attn.to_k": "attn.a_to_k",
|
| 440 |
+
"attn.to_v": "attn.a_to_v",
|
| 441 |
+
"attn.to_out.0": "attn.a_to_out",
|
| 442 |
+
"attn.add_q_proj": "attn.b_to_q",
|
| 443 |
+
"attn.add_k_proj": "attn.b_to_k",
|
| 444 |
+
"attn.add_v_proj": "attn.b_to_v",
|
| 445 |
+
"attn.to_add_out": "attn.b_to_out",
|
| 446 |
+
"ff.net.0.proj": "ff_a.0",
|
| 447 |
+
"ff.net.2": "ff_a.2",
|
| 448 |
+
"ff_context.net.0.proj": "ff_b.0",
|
| 449 |
+
"ff_context.net.2": "ff_b.2",
|
| 450 |
+
|
| 451 |
+
"attn.norm_q": "attn.norm_q_a",
|
| 452 |
+
"attn.norm_k": "attn.norm_k_a",
|
| 453 |
+
"attn.norm_added_q": "attn.norm_q_b",
|
| 454 |
+
"attn.norm_added_k": "attn.norm_k_b",
|
| 455 |
+
}
|
| 456 |
+
state_dict_ = {}
|
| 457 |
+
for name, param in state_dict.items():
|
| 458 |
+
if name in rename_dict:
|
| 459 |
+
if name == "pos_embed.pos_embed":
|
| 460 |
+
param = param.reshape((1, 192, 192, param.shape[-1]))
|
| 461 |
+
state_dict_[rename_dict[name]] = param
|
| 462 |
+
elif name.endswith(".weight") or name.endswith(".bias"):
|
| 463 |
+
suffix = ".weight" if name.endswith(".weight") else ".bias"
|
| 464 |
+
prefix = name[:-len(suffix)]
|
| 465 |
+
if prefix in rename_dict:
|
| 466 |
+
state_dict_[rename_dict[prefix] + suffix] = param
|
| 467 |
+
elif prefix.startswith("transformer_blocks."):
|
| 468 |
+
names = prefix.split(".")
|
| 469 |
+
names[0] = "blocks"
|
| 470 |
+
middle = ".".join(names[2:])
|
| 471 |
+
if middle in rename_dict:
|
| 472 |
+
name_ = ".".join(names[:2] + [rename_dict[middle]] + [suffix[1:]])
|
| 473 |
+
state_dict_[name_] = param
|
| 474 |
+
merged_keys = [name for name in state_dict_ if ".a_to_q." in name or ".b_to_q." in name]
|
| 475 |
+
for key in merged_keys:
|
| 476 |
+
param = torch.concat([
|
| 477 |
+
state_dict_[key.replace("to_q", "to_q")],
|
| 478 |
+
state_dict_[key.replace("to_q", "to_k")],
|
| 479 |
+
state_dict_[key.replace("to_q", "to_v")],
|
| 480 |
+
], dim=0)
|
| 481 |
+
name = key.replace("to_q", "to_qkv")
|
| 482 |
+
state_dict_.pop(key.replace("to_q", "to_q"))
|
| 483 |
+
state_dict_.pop(key.replace("to_q", "to_k"))
|
| 484 |
+
state_dict_.pop(key.replace("to_q", "to_v"))
|
| 485 |
+
state_dict_[name] = param
|
| 486 |
+
return state_dict_, self.infer_architecture(state_dict_)
|
| 487 |
+
|
| 488 |
+
def from_civitai(self, state_dict):
|
| 489 |
+
rename_dict = {
|
| 490 |
+
"model.diffusion_model.context_embedder.bias": "context_embedder.bias",
|
| 491 |
+
"model.diffusion_model.context_embedder.weight": "context_embedder.weight",
|
| 492 |
+
"model.diffusion_model.final_layer.linear.bias": "proj_out.bias",
|
| 493 |
+
"model.diffusion_model.final_layer.linear.weight": "proj_out.weight",
|
| 494 |
+
|
| 495 |
+
"model.diffusion_model.pos_embed": "pos_embedder.pos_embed",
|
| 496 |
+
"model.diffusion_model.t_embedder.mlp.0.bias": "time_embedder.timestep_embedder.0.bias",
|
| 497 |
+
"model.diffusion_model.t_embedder.mlp.0.weight": "time_embedder.timestep_embedder.0.weight",
|
| 498 |
+
"model.diffusion_model.t_embedder.mlp.2.bias": "time_embedder.timestep_embedder.2.bias",
|
| 499 |
+
"model.diffusion_model.t_embedder.mlp.2.weight": "time_embedder.timestep_embedder.2.weight",
|
| 500 |
+
"model.diffusion_model.x_embedder.proj.bias": "pos_embedder.proj.bias",
|
| 501 |
+
"model.diffusion_model.x_embedder.proj.weight": "pos_embedder.proj.weight",
|
| 502 |
+
"model.diffusion_model.y_embedder.mlp.0.bias": "pooled_text_embedder.0.bias",
|
| 503 |
+
"model.diffusion_model.y_embedder.mlp.0.weight": "pooled_text_embedder.0.weight",
|
| 504 |
+
"model.diffusion_model.y_embedder.mlp.2.bias": "pooled_text_embedder.2.bias",
|
| 505 |
+
"model.diffusion_model.y_embedder.mlp.2.weight": "pooled_text_embedder.2.weight",
|
| 506 |
+
|
| 507 |
+
"model.diffusion_model.joint_blocks.23.context_block.adaLN_modulation.1.weight": "blocks.23.norm1_b.linear.weight",
|
| 508 |
+
"model.diffusion_model.joint_blocks.23.context_block.adaLN_modulation.1.bias": "blocks.23.norm1_b.linear.bias",
|
| 509 |
+
"model.diffusion_model.final_layer.adaLN_modulation.1.weight": "norm_out.linear.weight",
|
| 510 |
+
"model.diffusion_model.final_layer.adaLN_modulation.1.bias": "norm_out.linear.bias",
|
| 511 |
+
}
|
| 512 |
+
for i in range(40):
|
| 513 |
+
rename_dict.update({
|
| 514 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.adaLN_modulation.1.bias": f"blocks.{i}.norm1_b.linear.bias",
|
| 515 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.adaLN_modulation.1.weight": f"blocks.{i}.norm1_b.linear.weight",
|
| 516 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.proj.bias": f"blocks.{i}.attn.b_to_out.bias",
|
| 517 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.proj.weight": f"blocks.{i}.attn.b_to_out.weight",
|
| 518 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.qkv.bias": [f'blocks.{i}.attn.b_to_q.bias', f'blocks.{i}.attn.b_to_k.bias', f'blocks.{i}.attn.b_to_v.bias'],
|
| 519 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.qkv.weight": [f'blocks.{i}.attn.b_to_q.weight', f'blocks.{i}.attn.b_to_k.weight', f'blocks.{i}.attn.b_to_v.weight'],
|
| 520 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc1.bias": f"blocks.{i}.ff_b.0.bias",
|
| 521 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc1.weight": f"blocks.{i}.ff_b.0.weight",
|
| 522 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc2.bias": f"blocks.{i}.ff_b.2.bias",
|
| 523 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.mlp.fc2.weight": f"blocks.{i}.ff_b.2.weight",
|
| 524 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.adaLN_modulation.1.bias": f"blocks.{i}.norm1_a.linear.bias",
|
| 525 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.adaLN_modulation.1.weight": f"blocks.{i}.norm1_a.linear.weight",
|
| 526 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.proj.bias": f"blocks.{i}.attn.a_to_out.bias",
|
| 527 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.proj.weight": f"blocks.{i}.attn.a_to_out.weight",
|
| 528 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.qkv.bias": [f'blocks.{i}.attn.a_to_q.bias', f'blocks.{i}.attn.a_to_k.bias', f'blocks.{i}.attn.a_to_v.bias'],
|
| 529 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.qkv.weight": [f'blocks.{i}.attn.a_to_q.weight', f'blocks.{i}.attn.a_to_k.weight', f'blocks.{i}.attn.a_to_v.weight'],
|
| 530 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc1.bias": f"blocks.{i}.ff_a.0.bias",
|
| 531 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc1.weight": f"blocks.{i}.ff_a.0.weight",
|
| 532 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc2.bias": f"blocks.{i}.ff_a.2.bias",
|
| 533 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.mlp.fc2.weight": f"blocks.{i}.ff_a.2.weight",
|
| 534 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.ln_q.weight": f"blocks.{i}.attn.norm_q_a.weight",
|
| 535 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn.ln_k.weight": f"blocks.{i}.attn.norm_k_a.weight",
|
| 536 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.ln_q.weight": f"blocks.{i}.attn.norm_q_b.weight",
|
| 537 |
+
f"model.diffusion_model.joint_blocks.{i}.context_block.attn.ln_k.weight": f"blocks.{i}.attn.norm_k_b.weight",
|
| 538 |
+
|
| 539 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.ln_q.weight": f"blocks.{i}.attn2.norm_q_a.weight",
|
| 540 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.ln_k.weight": f"blocks.{i}.attn2.norm_k_a.weight",
|
| 541 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.qkv.weight": f"blocks.{i}.attn2.a_to_qkv.weight",
|
| 542 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.qkv.bias": f"blocks.{i}.attn2.a_to_qkv.bias",
|
| 543 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.proj.weight": f"blocks.{i}.attn2.a_to_out.weight",
|
| 544 |
+
f"model.diffusion_model.joint_blocks.{i}.x_block.attn2.proj.bias": f"blocks.{i}.attn2.a_to_out.bias",
|
| 545 |
+
})
|
| 546 |
+
state_dict_ = {}
|
| 547 |
+
for name in state_dict:
|
| 548 |
+
if name in rename_dict:
|
| 549 |
+
param = state_dict[name]
|
| 550 |
+
if name == "model.diffusion_model.pos_embed":
|
| 551 |
+
pos_embed_max_size = int(param.shape[1] ** 0.5 + 0.4)
|
| 552 |
+
param = param.reshape((1, pos_embed_max_size, pos_embed_max_size, param.shape[-1]))
|
| 553 |
+
if isinstance(rename_dict[name], str):
|
| 554 |
+
state_dict_[rename_dict[name]] = param
|
| 555 |
+
else:
|
| 556 |
+
name_ = rename_dict[name][0].replace(".a_to_q.", ".a_to_qkv.").replace(".b_to_q.", ".b_to_qkv.")
|
| 557 |
+
state_dict_[name_] = param
|
| 558 |
+
extra_kwargs = self.infer_architecture(state_dict_)
|
| 559 |
+
num_layers = extra_kwargs["num_layers"]
|
| 560 |
+
for name in [
|
| 561 |
+
f"blocks.{num_layers-1}.norm1_b.linear.weight", f"blocks.{num_layers-1}.norm1_b.linear.bias", "norm_out.linear.weight", "norm_out.linear.bias",
|
| 562 |
+
]:
|
| 563 |
+
param = state_dict_[name]
|
| 564 |
+
dim = param.shape[0] // 2
|
| 565 |
+
param = torch.concat([param[dim:], param[:dim]], axis=0)
|
| 566 |
+
state_dict_[name] = param
|
| 567 |
+
return state_dict_, self.infer_architecture(state_dict_)
|
diffsynth/models/sd3_text_encoder.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
diffsynth/models/sd3_vae_decoder.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_vae_decoder import VAEAttentionBlock, SDVAEDecoderStateDictConverter
|
| 3 |
+
from .sd_unet import ResnetBlock, UpSampler
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SD3VAEDecoder(torch.nn.Module):
|
| 9 |
+
def __init__(self):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.scaling_factor = 1.5305 # Different from SD 1.x
|
| 12 |
+
self.shift_factor = 0.0609 # Different from SD 1.x
|
| 13 |
+
self.conv_in = torch.nn.Conv2d(16, 512, kernel_size=3, padding=1) # Different from SD 1.x
|
| 14 |
+
|
| 15 |
+
self.blocks = torch.nn.ModuleList([
|
| 16 |
+
# UNetMidBlock2D
|
| 17 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 18 |
+
VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),
|
| 19 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 20 |
+
# UpDecoderBlock2D
|
| 21 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 22 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 23 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 24 |
+
UpSampler(512),
|
| 25 |
+
# UpDecoderBlock2D
|
| 26 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 27 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 28 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 29 |
+
UpSampler(512),
|
| 30 |
+
# UpDecoderBlock2D
|
| 31 |
+
ResnetBlock(512, 256, eps=1e-6),
|
| 32 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 33 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 34 |
+
UpSampler(256),
|
| 35 |
+
# UpDecoderBlock2D
|
| 36 |
+
ResnetBlock(256, 128, eps=1e-6),
|
| 37 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 38 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 39 |
+
])
|
| 40 |
+
|
| 41 |
+
self.conv_norm_out = torch.nn.GroupNorm(num_channels=128, num_groups=32, eps=1e-6)
|
| 42 |
+
self.conv_act = torch.nn.SiLU()
|
| 43 |
+
self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1)
|
| 44 |
+
|
| 45 |
+
def tiled_forward(self, sample, tile_size=64, tile_stride=32):
|
| 46 |
+
hidden_states = TileWorker().tiled_forward(
|
| 47 |
+
lambda x: self.forward(x),
|
| 48 |
+
sample,
|
| 49 |
+
tile_size,
|
| 50 |
+
tile_stride,
|
| 51 |
+
tile_device=sample.device,
|
| 52 |
+
tile_dtype=sample.dtype
|
| 53 |
+
)
|
| 54 |
+
return hidden_states
|
| 55 |
+
|
| 56 |
+
def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
|
| 57 |
+
# For VAE Decoder, we do not need to apply the tiler on each layer.
|
| 58 |
+
if tiled:
|
| 59 |
+
return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
|
| 60 |
+
|
| 61 |
+
# 1. pre-process
|
| 62 |
+
hidden_states = sample / self.scaling_factor + self.shift_factor
|
| 63 |
+
hidden_states = self.conv_in(hidden_states)
|
| 64 |
+
time_emb = None
|
| 65 |
+
text_emb = None
|
| 66 |
+
res_stack = None
|
| 67 |
+
|
| 68 |
+
# 2. blocks
|
| 69 |
+
for i, block in enumerate(self.blocks):
|
| 70 |
+
hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
|
| 71 |
+
|
| 72 |
+
# 3. output
|
| 73 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 74 |
+
hidden_states = self.conv_act(hidden_states)
|
| 75 |
+
hidden_states = self.conv_out(hidden_states)
|
| 76 |
+
|
| 77 |
+
return hidden_states
|
| 78 |
+
|
| 79 |
+
@staticmethod
|
| 80 |
+
def state_dict_converter():
|
| 81 |
+
return SDVAEDecoderStateDictConverter()
|
diffsynth/models/sd3_vae_encoder.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_unet import ResnetBlock, DownSampler
|
| 3 |
+
from .sd_vae_encoder import VAEAttentionBlock, SDVAEEncoderStateDictConverter
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
from einops import rearrange
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SD3VAEEncoder(torch.nn.Module):
|
| 9 |
+
def __init__(self):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.scaling_factor = 1.5305 # Different from SD 1.x
|
| 12 |
+
self.shift_factor = 0.0609 # Different from SD 1.x
|
| 13 |
+
self.conv_in = torch.nn.Conv2d(3, 128, kernel_size=3, padding=1)
|
| 14 |
+
|
| 15 |
+
self.blocks = torch.nn.ModuleList([
|
| 16 |
+
# DownEncoderBlock2D
|
| 17 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 18 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 19 |
+
DownSampler(128, padding=0, extra_padding=True),
|
| 20 |
+
# DownEncoderBlock2D
|
| 21 |
+
ResnetBlock(128, 256, eps=1e-6),
|
| 22 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 23 |
+
DownSampler(256, padding=0, extra_padding=True),
|
| 24 |
+
# DownEncoderBlock2D
|
| 25 |
+
ResnetBlock(256, 512, eps=1e-6),
|
| 26 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 27 |
+
DownSampler(512, padding=0, extra_padding=True),
|
| 28 |
+
# DownEncoderBlock2D
|
| 29 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 30 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 31 |
+
# UNetMidBlock2D
|
| 32 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 33 |
+
VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),
|
| 34 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 35 |
+
])
|
| 36 |
+
|
| 37 |
+
self.conv_norm_out = torch.nn.GroupNorm(num_channels=512, num_groups=32, eps=1e-6)
|
| 38 |
+
self.conv_act = torch.nn.SiLU()
|
| 39 |
+
self.conv_out = torch.nn.Conv2d(512, 32, kernel_size=3, padding=1)
|
| 40 |
+
|
| 41 |
+
def tiled_forward(self, sample, tile_size=64, tile_stride=32):
|
| 42 |
+
hidden_states = TileWorker().tiled_forward(
|
| 43 |
+
lambda x: self.forward(x),
|
| 44 |
+
sample,
|
| 45 |
+
tile_size,
|
| 46 |
+
tile_stride,
|
| 47 |
+
tile_device=sample.device,
|
| 48 |
+
tile_dtype=sample.dtype
|
| 49 |
+
)
|
| 50 |
+
return hidden_states
|
| 51 |
+
|
| 52 |
+
def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
|
| 53 |
+
# For VAE Decoder, we do not need to apply the tiler on each layer.
|
| 54 |
+
if tiled:
|
| 55 |
+
return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
|
| 56 |
+
|
| 57 |
+
# 1. pre-process
|
| 58 |
+
hidden_states = self.conv_in(sample)
|
| 59 |
+
time_emb = None
|
| 60 |
+
text_emb = None
|
| 61 |
+
res_stack = None
|
| 62 |
+
|
| 63 |
+
# 2. blocks
|
| 64 |
+
for i, block in enumerate(self.blocks):
|
| 65 |
+
hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
|
| 66 |
+
|
| 67 |
+
# 3. output
|
| 68 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 69 |
+
hidden_states = self.conv_act(hidden_states)
|
| 70 |
+
hidden_states = self.conv_out(hidden_states)
|
| 71 |
+
hidden_states = hidden_states[:, :16]
|
| 72 |
+
hidden_states = (hidden_states - self.shift_factor) * self.scaling_factor
|
| 73 |
+
|
| 74 |
+
return hidden_states
|
| 75 |
+
|
| 76 |
+
def encode_video(self, sample, batch_size=8):
|
| 77 |
+
B = sample.shape[0]
|
| 78 |
+
hidden_states = []
|
| 79 |
+
|
| 80 |
+
for i in range(0, sample.shape[2], batch_size):
|
| 81 |
+
|
| 82 |
+
j = min(i + batch_size, sample.shape[2])
|
| 83 |
+
sample_batch = rearrange(sample[:,:,i:j], "B C T H W -> (B T) C H W")
|
| 84 |
+
|
| 85 |
+
hidden_states_batch = self(sample_batch)
|
| 86 |
+
hidden_states_batch = rearrange(hidden_states_batch, "(B T) C H W -> B C T H W", B=B)
|
| 87 |
+
|
| 88 |
+
hidden_states.append(hidden_states_batch)
|
| 89 |
+
|
| 90 |
+
hidden_states = torch.concat(hidden_states, dim=2)
|
| 91 |
+
return hidden_states
|
| 92 |
+
|
| 93 |
+
@staticmethod
|
| 94 |
+
def state_dict_converter():
|
| 95 |
+
return SDVAEEncoderStateDictConverter()
|
diffsynth/models/sd_controlnet.py
ADDED
|
@@ -0,0 +1,589 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_unet import Timesteps, ResnetBlock, AttentionBlock, PushBlock, DownSampler
|
| 3 |
+
from .tiler import TileWorker
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class ControlNetConditioningLayer(torch.nn.Module):
|
| 7 |
+
def __init__(self, channels = (3, 16, 32, 96, 256, 320)):
|
| 8 |
+
super().__init__()
|
| 9 |
+
self.blocks = torch.nn.ModuleList([])
|
| 10 |
+
self.blocks.append(torch.nn.Conv2d(channels[0], channels[1], kernel_size=3, padding=1))
|
| 11 |
+
self.blocks.append(torch.nn.SiLU())
|
| 12 |
+
for i in range(1, len(channels) - 2):
|
| 13 |
+
self.blocks.append(torch.nn.Conv2d(channels[i], channels[i], kernel_size=3, padding=1))
|
| 14 |
+
self.blocks.append(torch.nn.SiLU())
|
| 15 |
+
self.blocks.append(torch.nn.Conv2d(channels[i], channels[i+1], kernel_size=3, padding=1, stride=2))
|
| 16 |
+
self.blocks.append(torch.nn.SiLU())
|
| 17 |
+
self.blocks.append(torch.nn.Conv2d(channels[-2], channels[-1], kernel_size=3, padding=1))
|
| 18 |
+
|
| 19 |
+
def forward(self, conditioning):
|
| 20 |
+
for block in self.blocks:
|
| 21 |
+
conditioning = block(conditioning)
|
| 22 |
+
return conditioning
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class SDControlNet(torch.nn.Module):
|
| 26 |
+
def __init__(self, global_pool=False):
|
| 27 |
+
super().__init__()
|
| 28 |
+
self.time_proj = Timesteps(320)
|
| 29 |
+
self.time_embedding = torch.nn.Sequential(
|
| 30 |
+
torch.nn.Linear(320, 1280),
|
| 31 |
+
torch.nn.SiLU(),
|
| 32 |
+
torch.nn.Linear(1280, 1280)
|
| 33 |
+
)
|
| 34 |
+
self.conv_in = torch.nn.Conv2d(4, 320, kernel_size=3, padding=1)
|
| 35 |
+
|
| 36 |
+
self.controlnet_conv_in = ControlNetConditioningLayer(channels=(3, 16, 32, 96, 256, 320))
|
| 37 |
+
|
| 38 |
+
self.blocks = torch.nn.ModuleList([
|
| 39 |
+
# CrossAttnDownBlock2D
|
| 40 |
+
ResnetBlock(320, 320, 1280),
|
| 41 |
+
AttentionBlock(8, 40, 320, 1, 768),
|
| 42 |
+
PushBlock(),
|
| 43 |
+
ResnetBlock(320, 320, 1280),
|
| 44 |
+
AttentionBlock(8, 40, 320, 1, 768),
|
| 45 |
+
PushBlock(),
|
| 46 |
+
DownSampler(320),
|
| 47 |
+
PushBlock(),
|
| 48 |
+
# CrossAttnDownBlock2D
|
| 49 |
+
ResnetBlock(320, 640, 1280),
|
| 50 |
+
AttentionBlock(8, 80, 640, 1, 768),
|
| 51 |
+
PushBlock(),
|
| 52 |
+
ResnetBlock(640, 640, 1280),
|
| 53 |
+
AttentionBlock(8, 80, 640, 1, 768),
|
| 54 |
+
PushBlock(),
|
| 55 |
+
DownSampler(640),
|
| 56 |
+
PushBlock(),
|
| 57 |
+
# CrossAttnDownBlock2D
|
| 58 |
+
ResnetBlock(640, 1280, 1280),
|
| 59 |
+
AttentionBlock(8, 160, 1280, 1, 768),
|
| 60 |
+
PushBlock(),
|
| 61 |
+
ResnetBlock(1280, 1280, 1280),
|
| 62 |
+
AttentionBlock(8, 160, 1280, 1, 768),
|
| 63 |
+
PushBlock(),
|
| 64 |
+
DownSampler(1280),
|
| 65 |
+
PushBlock(),
|
| 66 |
+
# DownBlock2D
|
| 67 |
+
ResnetBlock(1280, 1280, 1280),
|
| 68 |
+
PushBlock(),
|
| 69 |
+
ResnetBlock(1280, 1280, 1280),
|
| 70 |
+
PushBlock(),
|
| 71 |
+
# UNetMidBlock2DCrossAttn
|
| 72 |
+
ResnetBlock(1280, 1280, 1280),
|
| 73 |
+
AttentionBlock(8, 160, 1280, 1, 768),
|
| 74 |
+
ResnetBlock(1280, 1280, 1280),
|
| 75 |
+
PushBlock()
|
| 76 |
+
])
|
| 77 |
+
|
| 78 |
+
self.controlnet_blocks = torch.nn.ModuleList([
|
| 79 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),
|
| 80 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False),
|
| 81 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False),
|
| 82 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1), bias=False),
|
| 83 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),
|
| 84 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1), bias=False),
|
| 85 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1), bias=False),
|
| 86 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),
|
| 87 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1), bias=False),
|
| 88 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1), bias=False),
|
| 89 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1), bias=False),
|
| 90 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1), bias=False),
|
| 91 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1), bias=False),
|
| 92 |
+
])
|
| 93 |
+
|
| 94 |
+
self.global_pool = global_pool
|
| 95 |
+
|
| 96 |
+
def forward(
|
| 97 |
+
self,
|
| 98 |
+
sample, timestep, encoder_hidden_states, conditioning,
|
| 99 |
+
tiled=False, tile_size=64, tile_stride=32,
|
| 100 |
+
**kwargs
|
| 101 |
+
):
|
| 102 |
+
# 1. time
|
| 103 |
+
time_emb = self.time_proj(timestep).to(sample.dtype)
|
| 104 |
+
time_emb = self.time_embedding(time_emb)
|
| 105 |
+
time_emb = time_emb.repeat(sample.shape[0], 1)
|
| 106 |
+
|
| 107 |
+
# 2. pre-process
|
| 108 |
+
height, width = sample.shape[2], sample.shape[3]
|
| 109 |
+
hidden_states = self.conv_in(sample) + self.controlnet_conv_in(conditioning)
|
| 110 |
+
text_emb = encoder_hidden_states
|
| 111 |
+
res_stack = [hidden_states]
|
| 112 |
+
|
| 113 |
+
# 3. blocks
|
| 114 |
+
for i, block in enumerate(self.blocks):
|
| 115 |
+
if tiled and not isinstance(block, PushBlock):
|
| 116 |
+
_, _, inter_height, _ = hidden_states.shape
|
| 117 |
+
resize_scale = inter_height / height
|
| 118 |
+
hidden_states = TileWorker().tiled_forward(
|
| 119 |
+
lambda x: block(x, time_emb, text_emb, res_stack)[0],
|
| 120 |
+
hidden_states,
|
| 121 |
+
int(tile_size * resize_scale),
|
| 122 |
+
int(tile_stride * resize_scale),
|
| 123 |
+
tile_device=hidden_states.device,
|
| 124 |
+
tile_dtype=hidden_states.dtype
|
| 125 |
+
)
|
| 126 |
+
else:
|
| 127 |
+
hidden_states, _, _, _ = block(hidden_states, time_emb, text_emb, res_stack)
|
| 128 |
+
|
| 129 |
+
# 4. ControlNet blocks
|
| 130 |
+
controlnet_res_stack = [block(res) for block, res in zip(self.controlnet_blocks, res_stack)]
|
| 131 |
+
|
| 132 |
+
# pool
|
| 133 |
+
if self.global_pool:
|
| 134 |
+
controlnet_res_stack = [res.mean(dim=(2, 3), keepdim=True) for res in controlnet_res_stack]
|
| 135 |
+
|
| 136 |
+
return controlnet_res_stack
|
| 137 |
+
|
| 138 |
+
@staticmethod
|
| 139 |
+
def state_dict_converter():
|
| 140 |
+
return SDControlNetStateDictConverter()
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
class SDControlNetStateDictConverter:
|
| 144 |
+
def __init__(self):
|
| 145 |
+
pass
|
| 146 |
+
|
| 147 |
+
def from_diffusers(self, state_dict):
|
| 148 |
+
# architecture
|
| 149 |
+
block_types = [
|
| 150 |
+
'ResnetBlock', 'AttentionBlock', 'PushBlock', 'ResnetBlock', 'AttentionBlock', 'PushBlock', 'DownSampler', 'PushBlock',
|
| 151 |
+
'ResnetBlock', 'AttentionBlock', 'PushBlock', 'ResnetBlock', 'AttentionBlock', 'PushBlock', 'DownSampler', 'PushBlock',
|
| 152 |
+
'ResnetBlock', 'AttentionBlock', 'PushBlock', 'ResnetBlock', 'AttentionBlock', 'PushBlock', 'DownSampler', 'PushBlock',
|
| 153 |
+
'ResnetBlock', 'PushBlock', 'ResnetBlock', 'PushBlock',
|
| 154 |
+
'ResnetBlock', 'AttentionBlock', 'ResnetBlock',
|
| 155 |
+
'PopBlock', 'ResnetBlock', 'PopBlock', 'ResnetBlock', 'PopBlock', 'ResnetBlock', 'UpSampler',
|
| 156 |
+
'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock', 'UpSampler',
|
| 157 |
+
'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock', 'UpSampler',
|
| 158 |
+
'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock', 'PopBlock', 'ResnetBlock', 'AttentionBlock'
|
| 159 |
+
]
|
| 160 |
+
|
| 161 |
+
# controlnet_rename_dict
|
| 162 |
+
controlnet_rename_dict = {
|
| 163 |
+
"controlnet_cond_embedding.conv_in.weight": "controlnet_conv_in.blocks.0.weight",
|
| 164 |
+
"controlnet_cond_embedding.conv_in.bias": "controlnet_conv_in.blocks.0.bias",
|
| 165 |
+
"controlnet_cond_embedding.blocks.0.weight": "controlnet_conv_in.blocks.2.weight",
|
| 166 |
+
"controlnet_cond_embedding.blocks.0.bias": "controlnet_conv_in.blocks.2.bias",
|
| 167 |
+
"controlnet_cond_embedding.blocks.1.weight": "controlnet_conv_in.blocks.4.weight",
|
| 168 |
+
"controlnet_cond_embedding.blocks.1.bias": "controlnet_conv_in.blocks.4.bias",
|
| 169 |
+
"controlnet_cond_embedding.blocks.2.weight": "controlnet_conv_in.blocks.6.weight",
|
| 170 |
+
"controlnet_cond_embedding.blocks.2.bias": "controlnet_conv_in.blocks.6.bias",
|
| 171 |
+
"controlnet_cond_embedding.blocks.3.weight": "controlnet_conv_in.blocks.8.weight",
|
| 172 |
+
"controlnet_cond_embedding.blocks.3.bias": "controlnet_conv_in.blocks.8.bias",
|
| 173 |
+
"controlnet_cond_embedding.blocks.4.weight": "controlnet_conv_in.blocks.10.weight",
|
| 174 |
+
"controlnet_cond_embedding.blocks.4.bias": "controlnet_conv_in.blocks.10.bias",
|
| 175 |
+
"controlnet_cond_embedding.blocks.5.weight": "controlnet_conv_in.blocks.12.weight",
|
| 176 |
+
"controlnet_cond_embedding.blocks.5.bias": "controlnet_conv_in.blocks.12.bias",
|
| 177 |
+
"controlnet_cond_embedding.conv_out.weight": "controlnet_conv_in.blocks.14.weight",
|
| 178 |
+
"controlnet_cond_embedding.conv_out.bias": "controlnet_conv_in.blocks.14.bias",
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
# Rename each parameter
|
| 182 |
+
name_list = sorted([name for name in state_dict])
|
| 183 |
+
rename_dict = {}
|
| 184 |
+
block_id = {"ResnetBlock": -1, "AttentionBlock": -1, "DownSampler": -1, "UpSampler": -1}
|
| 185 |
+
last_block_type_with_id = {"ResnetBlock": "", "AttentionBlock": "", "DownSampler": "", "UpSampler": ""}
|
| 186 |
+
for name in name_list:
|
| 187 |
+
names = name.split(".")
|
| 188 |
+
if names[0] in ["conv_in", "conv_norm_out", "conv_out"]:
|
| 189 |
+
pass
|
| 190 |
+
elif name in controlnet_rename_dict:
|
| 191 |
+
names = controlnet_rename_dict[name].split(".")
|
| 192 |
+
elif names[0] == "controlnet_down_blocks":
|
| 193 |
+
names[0] = "controlnet_blocks"
|
| 194 |
+
elif names[0] == "controlnet_mid_block":
|
| 195 |
+
names = ["controlnet_blocks", "12", names[-1]]
|
| 196 |
+
elif names[0] in ["time_embedding", "add_embedding"]:
|
| 197 |
+
if names[0] == "add_embedding":
|
| 198 |
+
names[0] = "add_time_embedding"
|
| 199 |
+
names[1] = {"linear_1": "0", "linear_2": "2"}[names[1]]
|
| 200 |
+
elif names[0] in ["down_blocks", "mid_block", "up_blocks"]:
|
| 201 |
+
if names[0] == "mid_block":
|
| 202 |
+
names.insert(1, "0")
|
| 203 |
+
block_type = {"resnets": "ResnetBlock", "attentions": "AttentionBlock", "downsamplers": "DownSampler", "upsamplers": "UpSampler"}[names[2]]
|
| 204 |
+
block_type_with_id = ".".join(names[:4])
|
| 205 |
+
if block_type_with_id != last_block_type_with_id[block_type]:
|
| 206 |
+
block_id[block_type] += 1
|
| 207 |
+
last_block_type_with_id[block_type] = block_type_with_id
|
| 208 |
+
while block_id[block_type] < len(block_types) and block_types[block_id[block_type]] != block_type:
|
| 209 |
+
block_id[block_type] += 1
|
| 210 |
+
block_type_with_id = ".".join(names[:4])
|
| 211 |
+
names = ["blocks", str(block_id[block_type])] + names[4:]
|
| 212 |
+
if "ff" in names:
|
| 213 |
+
ff_index = names.index("ff")
|
| 214 |
+
component = ".".join(names[ff_index:ff_index+3])
|
| 215 |
+
component = {"ff.net.0": "act_fn", "ff.net.2": "ff"}[component]
|
| 216 |
+
names = names[:ff_index] + [component] + names[ff_index+3:]
|
| 217 |
+
if "to_out" in names:
|
| 218 |
+
names.pop(names.index("to_out") + 1)
|
| 219 |
+
else:
|
| 220 |
+
raise ValueError(f"Unknown parameters: {name}")
|
| 221 |
+
rename_dict[name] = ".".join(names)
|
| 222 |
+
|
| 223 |
+
# Convert state_dict
|
| 224 |
+
state_dict_ = {}
|
| 225 |
+
for name, param in state_dict.items():
|
| 226 |
+
if ".proj_in." in name or ".proj_out." in name:
|
| 227 |
+
param = param.squeeze()
|
| 228 |
+
if rename_dict[name] in [
|
| 229 |
+
"controlnet_blocks.1.bias", "controlnet_blocks.2.bias", "controlnet_blocks.3.bias", "controlnet_blocks.5.bias", "controlnet_blocks.6.bias",
|
| 230 |
+
"controlnet_blocks.8.bias", "controlnet_blocks.9.bias", "controlnet_blocks.10.bias", "controlnet_blocks.11.bias", "controlnet_blocks.12.bias"
|
| 231 |
+
]:
|
| 232 |
+
continue
|
| 233 |
+
state_dict_[rename_dict[name]] = param
|
| 234 |
+
return state_dict_
|
| 235 |
+
|
| 236 |
+
def from_civitai(self, state_dict):
|
| 237 |
+
if "mid_block.resnets.1.time_emb_proj.weight" in state_dict:
|
| 238 |
+
# For controlnets in diffusers format
|
| 239 |
+
return self.from_diffusers(state_dict)
|
| 240 |
+
rename_dict = {
|
| 241 |
+
"control_model.time_embed.0.weight": "time_embedding.0.weight",
|
| 242 |
+
"control_model.time_embed.0.bias": "time_embedding.0.bias",
|
| 243 |
+
"control_model.time_embed.2.weight": "time_embedding.2.weight",
|
| 244 |
+
"control_model.time_embed.2.bias": "time_embedding.2.bias",
|
| 245 |
+
"control_model.input_blocks.0.0.weight": "conv_in.weight",
|
| 246 |
+
"control_model.input_blocks.0.0.bias": "conv_in.bias",
|
| 247 |
+
"control_model.input_blocks.1.0.in_layers.0.weight": "blocks.0.norm1.weight",
|
| 248 |
+
"control_model.input_blocks.1.0.in_layers.0.bias": "blocks.0.norm1.bias",
|
| 249 |
+
"control_model.input_blocks.1.0.in_layers.2.weight": "blocks.0.conv1.weight",
|
| 250 |
+
"control_model.input_blocks.1.0.in_layers.2.bias": "blocks.0.conv1.bias",
|
| 251 |
+
"control_model.input_blocks.1.0.emb_layers.1.weight": "blocks.0.time_emb_proj.weight",
|
| 252 |
+
"control_model.input_blocks.1.0.emb_layers.1.bias": "blocks.0.time_emb_proj.bias",
|
| 253 |
+
"control_model.input_blocks.1.0.out_layers.0.weight": "blocks.0.norm2.weight",
|
| 254 |
+
"control_model.input_blocks.1.0.out_layers.0.bias": "blocks.0.norm2.bias",
|
| 255 |
+
"control_model.input_blocks.1.0.out_layers.3.weight": "blocks.0.conv2.weight",
|
| 256 |
+
"control_model.input_blocks.1.0.out_layers.3.bias": "blocks.0.conv2.bias",
|
| 257 |
+
"control_model.input_blocks.1.1.norm.weight": "blocks.1.norm.weight",
|
| 258 |
+
"control_model.input_blocks.1.1.norm.bias": "blocks.1.norm.bias",
|
| 259 |
+
"control_model.input_blocks.1.1.proj_in.weight": "blocks.1.proj_in.weight",
|
| 260 |
+
"control_model.input_blocks.1.1.proj_in.bias": "blocks.1.proj_in.bias",
|
| 261 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn1.to_q.weight": "blocks.1.transformer_blocks.0.attn1.to_q.weight",
|
| 262 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn1.to_k.weight": "blocks.1.transformer_blocks.0.attn1.to_k.weight",
|
| 263 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn1.to_v.weight": "blocks.1.transformer_blocks.0.attn1.to_v.weight",
|
| 264 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.1.transformer_blocks.0.attn1.to_out.weight",
|
| 265 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.1.transformer_blocks.0.attn1.to_out.bias",
|
| 266 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.1.transformer_blocks.0.act_fn.proj.weight",
|
| 267 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.1.transformer_blocks.0.act_fn.proj.bias",
|
| 268 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.ff.net.2.weight": "blocks.1.transformer_blocks.0.ff.weight",
|
| 269 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.ff.net.2.bias": "blocks.1.transformer_blocks.0.ff.bias",
|
| 270 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn2.to_q.weight": "blocks.1.transformer_blocks.0.attn2.to_q.weight",
|
| 271 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn2.to_k.weight": "blocks.1.transformer_blocks.0.attn2.to_k.weight",
|
| 272 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn2.to_v.weight": "blocks.1.transformer_blocks.0.attn2.to_v.weight",
|
| 273 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.1.transformer_blocks.0.attn2.to_out.weight",
|
| 274 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.1.transformer_blocks.0.attn2.to_out.bias",
|
| 275 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm1.weight": "blocks.1.transformer_blocks.0.norm1.weight",
|
| 276 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm1.bias": "blocks.1.transformer_blocks.0.norm1.bias",
|
| 277 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm2.weight": "blocks.1.transformer_blocks.0.norm2.weight",
|
| 278 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm2.bias": "blocks.1.transformer_blocks.0.norm2.bias",
|
| 279 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm3.weight": "blocks.1.transformer_blocks.0.norm3.weight",
|
| 280 |
+
"control_model.input_blocks.1.1.transformer_blocks.0.norm3.bias": "blocks.1.transformer_blocks.0.norm3.bias",
|
| 281 |
+
"control_model.input_blocks.1.1.proj_out.weight": "blocks.1.proj_out.weight",
|
| 282 |
+
"control_model.input_blocks.1.1.proj_out.bias": "blocks.1.proj_out.bias",
|
| 283 |
+
"control_model.input_blocks.2.0.in_layers.0.weight": "blocks.3.norm1.weight",
|
| 284 |
+
"control_model.input_blocks.2.0.in_layers.0.bias": "blocks.3.norm1.bias",
|
| 285 |
+
"control_model.input_blocks.2.0.in_layers.2.weight": "blocks.3.conv1.weight",
|
| 286 |
+
"control_model.input_blocks.2.0.in_layers.2.bias": "blocks.3.conv1.bias",
|
| 287 |
+
"control_model.input_blocks.2.0.emb_layers.1.weight": "blocks.3.time_emb_proj.weight",
|
| 288 |
+
"control_model.input_blocks.2.0.emb_layers.1.bias": "blocks.3.time_emb_proj.bias",
|
| 289 |
+
"control_model.input_blocks.2.0.out_layers.0.weight": "blocks.3.norm2.weight",
|
| 290 |
+
"control_model.input_blocks.2.0.out_layers.0.bias": "blocks.3.norm2.bias",
|
| 291 |
+
"control_model.input_blocks.2.0.out_layers.3.weight": "blocks.3.conv2.weight",
|
| 292 |
+
"control_model.input_blocks.2.0.out_layers.3.bias": "blocks.3.conv2.bias",
|
| 293 |
+
"control_model.input_blocks.2.1.norm.weight": "blocks.4.norm.weight",
|
| 294 |
+
"control_model.input_blocks.2.1.norm.bias": "blocks.4.norm.bias",
|
| 295 |
+
"control_model.input_blocks.2.1.proj_in.weight": "blocks.4.proj_in.weight",
|
| 296 |
+
"control_model.input_blocks.2.1.proj_in.bias": "blocks.4.proj_in.bias",
|
| 297 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn1.to_q.weight": "blocks.4.transformer_blocks.0.attn1.to_q.weight",
|
| 298 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn1.to_k.weight": "blocks.4.transformer_blocks.0.attn1.to_k.weight",
|
| 299 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn1.to_v.weight": "blocks.4.transformer_blocks.0.attn1.to_v.weight",
|
| 300 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.4.transformer_blocks.0.attn1.to_out.weight",
|
| 301 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.4.transformer_blocks.0.attn1.to_out.bias",
|
| 302 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.4.transformer_blocks.0.act_fn.proj.weight",
|
| 303 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.4.transformer_blocks.0.act_fn.proj.bias",
|
| 304 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.ff.net.2.weight": "blocks.4.transformer_blocks.0.ff.weight",
|
| 305 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.ff.net.2.bias": "blocks.4.transformer_blocks.0.ff.bias",
|
| 306 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn2.to_q.weight": "blocks.4.transformer_blocks.0.attn2.to_q.weight",
|
| 307 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn2.to_k.weight": "blocks.4.transformer_blocks.0.attn2.to_k.weight",
|
| 308 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn2.to_v.weight": "blocks.4.transformer_blocks.0.attn2.to_v.weight",
|
| 309 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.4.transformer_blocks.0.attn2.to_out.weight",
|
| 310 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.4.transformer_blocks.0.attn2.to_out.bias",
|
| 311 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm1.weight": "blocks.4.transformer_blocks.0.norm1.weight",
|
| 312 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm1.bias": "blocks.4.transformer_blocks.0.norm1.bias",
|
| 313 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm2.weight": "blocks.4.transformer_blocks.0.norm2.weight",
|
| 314 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm2.bias": "blocks.4.transformer_blocks.0.norm2.bias",
|
| 315 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm3.weight": "blocks.4.transformer_blocks.0.norm3.weight",
|
| 316 |
+
"control_model.input_blocks.2.1.transformer_blocks.0.norm3.bias": "blocks.4.transformer_blocks.0.norm3.bias",
|
| 317 |
+
"control_model.input_blocks.2.1.proj_out.weight": "blocks.4.proj_out.weight",
|
| 318 |
+
"control_model.input_blocks.2.1.proj_out.bias": "blocks.4.proj_out.bias",
|
| 319 |
+
"control_model.input_blocks.3.0.op.weight": "blocks.6.conv.weight",
|
| 320 |
+
"control_model.input_blocks.3.0.op.bias": "blocks.6.conv.bias",
|
| 321 |
+
"control_model.input_blocks.4.0.in_layers.0.weight": "blocks.8.norm1.weight",
|
| 322 |
+
"control_model.input_blocks.4.0.in_layers.0.bias": "blocks.8.norm1.bias",
|
| 323 |
+
"control_model.input_blocks.4.0.in_layers.2.weight": "blocks.8.conv1.weight",
|
| 324 |
+
"control_model.input_blocks.4.0.in_layers.2.bias": "blocks.8.conv1.bias",
|
| 325 |
+
"control_model.input_blocks.4.0.emb_layers.1.weight": "blocks.8.time_emb_proj.weight",
|
| 326 |
+
"control_model.input_blocks.4.0.emb_layers.1.bias": "blocks.8.time_emb_proj.bias",
|
| 327 |
+
"control_model.input_blocks.4.0.out_layers.0.weight": "blocks.8.norm2.weight",
|
| 328 |
+
"control_model.input_blocks.4.0.out_layers.0.bias": "blocks.8.norm2.bias",
|
| 329 |
+
"control_model.input_blocks.4.0.out_layers.3.weight": "blocks.8.conv2.weight",
|
| 330 |
+
"control_model.input_blocks.4.0.out_layers.3.bias": "blocks.8.conv2.bias",
|
| 331 |
+
"control_model.input_blocks.4.0.skip_connection.weight": "blocks.8.conv_shortcut.weight",
|
| 332 |
+
"control_model.input_blocks.4.0.skip_connection.bias": "blocks.8.conv_shortcut.bias",
|
| 333 |
+
"control_model.input_blocks.4.1.norm.weight": "blocks.9.norm.weight",
|
| 334 |
+
"control_model.input_blocks.4.1.norm.bias": "blocks.9.norm.bias",
|
| 335 |
+
"control_model.input_blocks.4.1.proj_in.weight": "blocks.9.proj_in.weight",
|
| 336 |
+
"control_model.input_blocks.4.1.proj_in.bias": "blocks.9.proj_in.bias",
|
| 337 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn1.to_q.weight": "blocks.9.transformer_blocks.0.attn1.to_q.weight",
|
| 338 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn1.to_k.weight": "blocks.9.transformer_blocks.0.attn1.to_k.weight",
|
| 339 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn1.to_v.weight": "blocks.9.transformer_blocks.0.attn1.to_v.weight",
|
| 340 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.9.transformer_blocks.0.attn1.to_out.weight",
|
| 341 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.9.transformer_blocks.0.attn1.to_out.bias",
|
| 342 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.9.transformer_blocks.0.act_fn.proj.weight",
|
| 343 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.9.transformer_blocks.0.act_fn.proj.bias",
|
| 344 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.ff.net.2.weight": "blocks.9.transformer_blocks.0.ff.weight",
|
| 345 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.ff.net.2.bias": "blocks.9.transformer_blocks.0.ff.bias",
|
| 346 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn2.to_q.weight": "blocks.9.transformer_blocks.0.attn2.to_q.weight",
|
| 347 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn2.to_k.weight": "blocks.9.transformer_blocks.0.attn2.to_k.weight",
|
| 348 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn2.to_v.weight": "blocks.9.transformer_blocks.0.attn2.to_v.weight",
|
| 349 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.9.transformer_blocks.0.attn2.to_out.weight",
|
| 350 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.9.transformer_blocks.0.attn2.to_out.bias",
|
| 351 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm1.weight": "blocks.9.transformer_blocks.0.norm1.weight",
|
| 352 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm1.bias": "blocks.9.transformer_blocks.0.norm1.bias",
|
| 353 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm2.weight": "blocks.9.transformer_blocks.0.norm2.weight",
|
| 354 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm2.bias": "blocks.9.transformer_blocks.0.norm2.bias",
|
| 355 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm3.weight": "blocks.9.transformer_blocks.0.norm3.weight",
|
| 356 |
+
"control_model.input_blocks.4.1.transformer_blocks.0.norm3.bias": "blocks.9.transformer_blocks.0.norm3.bias",
|
| 357 |
+
"control_model.input_blocks.4.1.proj_out.weight": "blocks.9.proj_out.weight",
|
| 358 |
+
"control_model.input_blocks.4.1.proj_out.bias": "blocks.9.proj_out.bias",
|
| 359 |
+
"control_model.input_blocks.5.0.in_layers.0.weight": "blocks.11.norm1.weight",
|
| 360 |
+
"control_model.input_blocks.5.0.in_layers.0.bias": "blocks.11.norm1.bias",
|
| 361 |
+
"control_model.input_blocks.5.0.in_layers.2.weight": "blocks.11.conv1.weight",
|
| 362 |
+
"control_model.input_blocks.5.0.in_layers.2.bias": "blocks.11.conv1.bias",
|
| 363 |
+
"control_model.input_blocks.5.0.emb_layers.1.weight": "blocks.11.time_emb_proj.weight",
|
| 364 |
+
"control_model.input_blocks.5.0.emb_layers.1.bias": "blocks.11.time_emb_proj.bias",
|
| 365 |
+
"control_model.input_blocks.5.0.out_layers.0.weight": "blocks.11.norm2.weight",
|
| 366 |
+
"control_model.input_blocks.5.0.out_layers.0.bias": "blocks.11.norm2.bias",
|
| 367 |
+
"control_model.input_blocks.5.0.out_layers.3.weight": "blocks.11.conv2.weight",
|
| 368 |
+
"control_model.input_blocks.5.0.out_layers.3.bias": "blocks.11.conv2.bias",
|
| 369 |
+
"control_model.input_blocks.5.1.norm.weight": "blocks.12.norm.weight",
|
| 370 |
+
"control_model.input_blocks.5.1.norm.bias": "blocks.12.norm.bias",
|
| 371 |
+
"control_model.input_blocks.5.1.proj_in.weight": "blocks.12.proj_in.weight",
|
| 372 |
+
"control_model.input_blocks.5.1.proj_in.bias": "blocks.12.proj_in.bias",
|
| 373 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn1.to_q.weight": "blocks.12.transformer_blocks.0.attn1.to_q.weight",
|
| 374 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn1.to_k.weight": "blocks.12.transformer_blocks.0.attn1.to_k.weight",
|
| 375 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn1.to_v.weight": "blocks.12.transformer_blocks.0.attn1.to_v.weight",
|
| 376 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.12.transformer_blocks.0.attn1.to_out.weight",
|
| 377 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.12.transformer_blocks.0.attn1.to_out.bias",
|
| 378 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.12.transformer_blocks.0.act_fn.proj.weight",
|
| 379 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.12.transformer_blocks.0.act_fn.proj.bias",
|
| 380 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.ff.net.2.weight": "blocks.12.transformer_blocks.0.ff.weight",
|
| 381 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.ff.net.2.bias": "blocks.12.transformer_blocks.0.ff.bias",
|
| 382 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn2.to_q.weight": "blocks.12.transformer_blocks.0.attn2.to_q.weight",
|
| 383 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn2.to_k.weight": "blocks.12.transformer_blocks.0.attn2.to_k.weight",
|
| 384 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn2.to_v.weight": "blocks.12.transformer_blocks.0.attn2.to_v.weight",
|
| 385 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.12.transformer_blocks.0.attn2.to_out.weight",
|
| 386 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.12.transformer_blocks.0.attn2.to_out.bias",
|
| 387 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm1.weight": "blocks.12.transformer_blocks.0.norm1.weight",
|
| 388 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm1.bias": "blocks.12.transformer_blocks.0.norm1.bias",
|
| 389 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm2.weight": "blocks.12.transformer_blocks.0.norm2.weight",
|
| 390 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm2.bias": "blocks.12.transformer_blocks.0.norm2.bias",
|
| 391 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm3.weight": "blocks.12.transformer_blocks.0.norm3.weight",
|
| 392 |
+
"control_model.input_blocks.5.1.transformer_blocks.0.norm3.bias": "blocks.12.transformer_blocks.0.norm3.bias",
|
| 393 |
+
"control_model.input_blocks.5.1.proj_out.weight": "blocks.12.proj_out.weight",
|
| 394 |
+
"control_model.input_blocks.5.1.proj_out.bias": "blocks.12.proj_out.bias",
|
| 395 |
+
"control_model.input_blocks.6.0.op.weight": "blocks.14.conv.weight",
|
| 396 |
+
"control_model.input_blocks.6.0.op.bias": "blocks.14.conv.bias",
|
| 397 |
+
"control_model.input_blocks.7.0.in_layers.0.weight": "blocks.16.norm1.weight",
|
| 398 |
+
"control_model.input_blocks.7.0.in_layers.0.bias": "blocks.16.norm1.bias",
|
| 399 |
+
"control_model.input_blocks.7.0.in_layers.2.weight": "blocks.16.conv1.weight",
|
| 400 |
+
"control_model.input_blocks.7.0.in_layers.2.bias": "blocks.16.conv1.bias",
|
| 401 |
+
"control_model.input_blocks.7.0.emb_layers.1.weight": "blocks.16.time_emb_proj.weight",
|
| 402 |
+
"control_model.input_blocks.7.0.emb_layers.1.bias": "blocks.16.time_emb_proj.bias",
|
| 403 |
+
"control_model.input_blocks.7.0.out_layers.0.weight": "blocks.16.norm2.weight",
|
| 404 |
+
"control_model.input_blocks.7.0.out_layers.0.bias": "blocks.16.norm2.bias",
|
| 405 |
+
"control_model.input_blocks.7.0.out_layers.3.weight": "blocks.16.conv2.weight",
|
| 406 |
+
"control_model.input_blocks.7.0.out_layers.3.bias": "blocks.16.conv2.bias",
|
| 407 |
+
"control_model.input_blocks.7.0.skip_connection.weight": "blocks.16.conv_shortcut.weight",
|
| 408 |
+
"control_model.input_blocks.7.0.skip_connection.bias": "blocks.16.conv_shortcut.bias",
|
| 409 |
+
"control_model.input_blocks.7.1.norm.weight": "blocks.17.norm.weight",
|
| 410 |
+
"control_model.input_blocks.7.1.norm.bias": "blocks.17.norm.bias",
|
| 411 |
+
"control_model.input_blocks.7.1.proj_in.weight": "blocks.17.proj_in.weight",
|
| 412 |
+
"control_model.input_blocks.7.1.proj_in.bias": "blocks.17.proj_in.bias",
|
| 413 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn1.to_q.weight": "blocks.17.transformer_blocks.0.attn1.to_q.weight",
|
| 414 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn1.to_k.weight": "blocks.17.transformer_blocks.0.attn1.to_k.weight",
|
| 415 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn1.to_v.weight": "blocks.17.transformer_blocks.0.attn1.to_v.weight",
|
| 416 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.17.transformer_blocks.0.attn1.to_out.weight",
|
| 417 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.17.transformer_blocks.0.attn1.to_out.bias",
|
| 418 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.17.transformer_blocks.0.act_fn.proj.weight",
|
| 419 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.17.transformer_blocks.0.act_fn.proj.bias",
|
| 420 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.ff.net.2.weight": "blocks.17.transformer_blocks.0.ff.weight",
|
| 421 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.ff.net.2.bias": "blocks.17.transformer_blocks.0.ff.bias",
|
| 422 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn2.to_q.weight": "blocks.17.transformer_blocks.0.attn2.to_q.weight",
|
| 423 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn2.to_k.weight": "blocks.17.transformer_blocks.0.attn2.to_k.weight",
|
| 424 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn2.to_v.weight": "blocks.17.transformer_blocks.0.attn2.to_v.weight",
|
| 425 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.17.transformer_blocks.0.attn2.to_out.weight",
|
| 426 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.17.transformer_blocks.0.attn2.to_out.bias",
|
| 427 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm1.weight": "blocks.17.transformer_blocks.0.norm1.weight",
|
| 428 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm1.bias": "blocks.17.transformer_blocks.0.norm1.bias",
|
| 429 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm2.weight": "blocks.17.transformer_blocks.0.norm2.weight",
|
| 430 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm2.bias": "blocks.17.transformer_blocks.0.norm2.bias",
|
| 431 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm3.weight": "blocks.17.transformer_blocks.0.norm3.weight",
|
| 432 |
+
"control_model.input_blocks.7.1.transformer_blocks.0.norm3.bias": "blocks.17.transformer_blocks.0.norm3.bias",
|
| 433 |
+
"control_model.input_blocks.7.1.proj_out.weight": "blocks.17.proj_out.weight",
|
| 434 |
+
"control_model.input_blocks.7.1.proj_out.bias": "blocks.17.proj_out.bias",
|
| 435 |
+
"control_model.input_blocks.8.0.in_layers.0.weight": "blocks.19.norm1.weight",
|
| 436 |
+
"control_model.input_blocks.8.0.in_layers.0.bias": "blocks.19.norm1.bias",
|
| 437 |
+
"control_model.input_blocks.8.0.in_layers.2.weight": "blocks.19.conv1.weight",
|
| 438 |
+
"control_model.input_blocks.8.0.in_layers.2.bias": "blocks.19.conv1.bias",
|
| 439 |
+
"control_model.input_blocks.8.0.emb_layers.1.weight": "blocks.19.time_emb_proj.weight",
|
| 440 |
+
"control_model.input_blocks.8.0.emb_layers.1.bias": "blocks.19.time_emb_proj.bias",
|
| 441 |
+
"control_model.input_blocks.8.0.out_layers.0.weight": "blocks.19.norm2.weight",
|
| 442 |
+
"control_model.input_blocks.8.0.out_layers.0.bias": "blocks.19.norm2.bias",
|
| 443 |
+
"control_model.input_blocks.8.0.out_layers.3.weight": "blocks.19.conv2.weight",
|
| 444 |
+
"control_model.input_blocks.8.0.out_layers.3.bias": "blocks.19.conv2.bias",
|
| 445 |
+
"control_model.input_blocks.8.1.norm.weight": "blocks.20.norm.weight",
|
| 446 |
+
"control_model.input_blocks.8.1.norm.bias": "blocks.20.norm.bias",
|
| 447 |
+
"control_model.input_blocks.8.1.proj_in.weight": "blocks.20.proj_in.weight",
|
| 448 |
+
"control_model.input_blocks.8.1.proj_in.bias": "blocks.20.proj_in.bias",
|
| 449 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn1.to_q.weight": "blocks.20.transformer_blocks.0.attn1.to_q.weight",
|
| 450 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn1.to_k.weight": "blocks.20.transformer_blocks.0.attn1.to_k.weight",
|
| 451 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn1.to_v.weight": "blocks.20.transformer_blocks.0.attn1.to_v.weight",
|
| 452 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.20.transformer_blocks.0.attn1.to_out.weight",
|
| 453 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.20.transformer_blocks.0.attn1.to_out.bias",
|
| 454 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.20.transformer_blocks.0.act_fn.proj.weight",
|
| 455 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.20.transformer_blocks.0.act_fn.proj.bias",
|
| 456 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.ff.net.2.weight": "blocks.20.transformer_blocks.0.ff.weight",
|
| 457 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.ff.net.2.bias": "blocks.20.transformer_blocks.0.ff.bias",
|
| 458 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn2.to_q.weight": "blocks.20.transformer_blocks.0.attn2.to_q.weight",
|
| 459 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn2.to_k.weight": "blocks.20.transformer_blocks.0.attn2.to_k.weight",
|
| 460 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn2.to_v.weight": "blocks.20.transformer_blocks.0.attn2.to_v.weight",
|
| 461 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.20.transformer_blocks.0.attn2.to_out.weight",
|
| 462 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.20.transformer_blocks.0.attn2.to_out.bias",
|
| 463 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm1.weight": "blocks.20.transformer_blocks.0.norm1.weight",
|
| 464 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm1.bias": "blocks.20.transformer_blocks.0.norm1.bias",
|
| 465 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm2.weight": "blocks.20.transformer_blocks.0.norm2.weight",
|
| 466 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm2.bias": "blocks.20.transformer_blocks.0.norm2.bias",
|
| 467 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm3.weight": "blocks.20.transformer_blocks.0.norm3.weight",
|
| 468 |
+
"control_model.input_blocks.8.1.transformer_blocks.0.norm3.bias": "blocks.20.transformer_blocks.0.norm3.bias",
|
| 469 |
+
"control_model.input_blocks.8.1.proj_out.weight": "blocks.20.proj_out.weight",
|
| 470 |
+
"control_model.input_blocks.8.1.proj_out.bias": "blocks.20.proj_out.bias",
|
| 471 |
+
"control_model.input_blocks.9.0.op.weight": "blocks.22.conv.weight",
|
| 472 |
+
"control_model.input_blocks.9.0.op.bias": "blocks.22.conv.bias",
|
| 473 |
+
"control_model.input_blocks.10.0.in_layers.0.weight": "blocks.24.norm1.weight",
|
| 474 |
+
"control_model.input_blocks.10.0.in_layers.0.bias": "blocks.24.norm1.bias",
|
| 475 |
+
"control_model.input_blocks.10.0.in_layers.2.weight": "blocks.24.conv1.weight",
|
| 476 |
+
"control_model.input_blocks.10.0.in_layers.2.bias": "blocks.24.conv1.bias",
|
| 477 |
+
"control_model.input_blocks.10.0.emb_layers.1.weight": "blocks.24.time_emb_proj.weight",
|
| 478 |
+
"control_model.input_blocks.10.0.emb_layers.1.bias": "blocks.24.time_emb_proj.bias",
|
| 479 |
+
"control_model.input_blocks.10.0.out_layers.0.weight": "blocks.24.norm2.weight",
|
| 480 |
+
"control_model.input_blocks.10.0.out_layers.0.bias": "blocks.24.norm2.bias",
|
| 481 |
+
"control_model.input_blocks.10.0.out_layers.3.weight": "blocks.24.conv2.weight",
|
| 482 |
+
"control_model.input_blocks.10.0.out_layers.3.bias": "blocks.24.conv2.bias",
|
| 483 |
+
"control_model.input_blocks.11.0.in_layers.0.weight": "blocks.26.norm1.weight",
|
| 484 |
+
"control_model.input_blocks.11.0.in_layers.0.bias": "blocks.26.norm1.bias",
|
| 485 |
+
"control_model.input_blocks.11.0.in_layers.2.weight": "blocks.26.conv1.weight",
|
| 486 |
+
"control_model.input_blocks.11.0.in_layers.2.bias": "blocks.26.conv1.bias",
|
| 487 |
+
"control_model.input_blocks.11.0.emb_layers.1.weight": "blocks.26.time_emb_proj.weight",
|
| 488 |
+
"control_model.input_blocks.11.0.emb_layers.1.bias": "blocks.26.time_emb_proj.bias",
|
| 489 |
+
"control_model.input_blocks.11.0.out_layers.0.weight": "blocks.26.norm2.weight",
|
| 490 |
+
"control_model.input_blocks.11.0.out_layers.0.bias": "blocks.26.norm2.bias",
|
| 491 |
+
"control_model.input_blocks.11.0.out_layers.3.weight": "blocks.26.conv2.weight",
|
| 492 |
+
"control_model.input_blocks.11.0.out_layers.3.bias": "blocks.26.conv2.bias",
|
| 493 |
+
"control_model.zero_convs.0.0.weight": "controlnet_blocks.0.weight",
|
| 494 |
+
"control_model.zero_convs.0.0.bias": "controlnet_blocks.0.bias",
|
| 495 |
+
"control_model.zero_convs.1.0.weight": "controlnet_blocks.1.weight",
|
| 496 |
+
"control_model.zero_convs.1.0.bias": "controlnet_blocks.0.bias",
|
| 497 |
+
"control_model.zero_convs.2.0.weight": "controlnet_blocks.2.weight",
|
| 498 |
+
"control_model.zero_convs.2.0.bias": "controlnet_blocks.0.bias",
|
| 499 |
+
"control_model.zero_convs.3.0.weight": "controlnet_blocks.3.weight",
|
| 500 |
+
"control_model.zero_convs.3.0.bias": "controlnet_blocks.0.bias",
|
| 501 |
+
"control_model.zero_convs.4.0.weight": "controlnet_blocks.4.weight",
|
| 502 |
+
"control_model.zero_convs.4.0.bias": "controlnet_blocks.4.bias",
|
| 503 |
+
"control_model.zero_convs.5.0.weight": "controlnet_blocks.5.weight",
|
| 504 |
+
"control_model.zero_convs.5.0.bias": "controlnet_blocks.4.bias",
|
| 505 |
+
"control_model.zero_convs.6.0.weight": "controlnet_blocks.6.weight",
|
| 506 |
+
"control_model.zero_convs.6.0.bias": "controlnet_blocks.4.bias",
|
| 507 |
+
"control_model.zero_convs.7.0.weight": "controlnet_blocks.7.weight",
|
| 508 |
+
"control_model.zero_convs.7.0.bias": "controlnet_blocks.7.bias",
|
| 509 |
+
"control_model.zero_convs.8.0.weight": "controlnet_blocks.8.weight",
|
| 510 |
+
"control_model.zero_convs.8.0.bias": "controlnet_blocks.7.bias",
|
| 511 |
+
"control_model.zero_convs.9.0.weight": "controlnet_blocks.9.weight",
|
| 512 |
+
"control_model.zero_convs.9.0.bias": "controlnet_blocks.7.bias",
|
| 513 |
+
"control_model.zero_convs.10.0.weight": "controlnet_blocks.10.weight",
|
| 514 |
+
"control_model.zero_convs.10.0.bias": "controlnet_blocks.7.bias",
|
| 515 |
+
"control_model.zero_convs.11.0.weight": "controlnet_blocks.11.weight",
|
| 516 |
+
"control_model.zero_convs.11.0.bias": "controlnet_blocks.7.bias",
|
| 517 |
+
"control_model.input_hint_block.0.weight": "controlnet_conv_in.blocks.0.weight",
|
| 518 |
+
"control_model.input_hint_block.0.bias": "controlnet_conv_in.blocks.0.bias",
|
| 519 |
+
"control_model.input_hint_block.2.weight": "controlnet_conv_in.blocks.2.weight",
|
| 520 |
+
"control_model.input_hint_block.2.bias": "controlnet_conv_in.blocks.2.bias",
|
| 521 |
+
"control_model.input_hint_block.4.weight": "controlnet_conv_in.blocks.4.weight",
|
| 522 |
+
"control_model.input_hint_block.4.bias": "controlnet_conv_in.blocks.4.bias",
|
| 523 |
+
"control_model.input_hint_block.6.weight": "controlnet_conv_in.blocks.6.weight",
|
| 524 |
+
"control_model.input_hint_block.6.bias": "controlnet_conv_in.blocks.6.bias",
|
| 525 |
+
"control_model.input_hint_block.8.weight": "controlnet_conv_in.blocks.8.weight",
|
| 526 |
+
"control_model.input_hint_block.8.bias": "controlnet_conv_in.blocks.8.bias",
|
| 527 |
+
"control_model.input_hint_block.10.weight": "controlnet_conv_in.blocks.10.weight",
|
| 528 |
+
"control_model.input_hint_block.10.bias": "controlnet_conv_in.blocks.10.bias",
|
| 529 |
+
"control_model.input_hint_block.12.weight": "controlnet_conv_in.blocks.12.weight",
|
| 530 |
+
"control_model.input_hint_block.12.bias": "controlnet_conv_in.blocks.12.bias",
|
| 531 |
+
"control_model.input_hint_block.14.weight": "controlnet_conv_in.blocks.14.weight",
|
| 532 |
+
"control_model.input_hint_block.14.bias": "controlnet_conv_in.blocks.14.bias",
|
| 533 |
+
"control_model.middle_block.0.in_layers.0.weight": "blocks.28.norm1.weight",
|
| 534 |
+
"control_model.middle_block.0.in_layers.0.bias": "blocks.28.norm1.bias",
|
| 535 |
+
"control_model.middle_block.0.in_layers.2.weight": "blocks.28.conv1.weight",
|
| 536 |
+
"control_model.middle_block.0.in_layers.2.bias": "blocks.28.conv1.bias",
|
| 537 |
+
"control_model.middle_block.0.emb_layers.1.weight": "blocks.28.time_emb_proj.weight",
|
| 538 |
+
"control_model.middle_block.0.emb_layers.1.bias": "blocks.28.time_emb_proj.bias",
|
| 539 |
+
"control_model.middle_block.0.out_layers.0.weight": "blocks.28.norm2.weight",
|
| 540 |
+
"control_model.middle_block.0.out_layers.0.bias": "blocks.28.norm2.bias",
|
| 541 |
+
"control_model.middle_block.0.out_layers.3.weight": "blocks.28.conv2.weight",
|
| 542 |
+
"control_model.middle_block.0.out_layers.3.bias": "blocks.28.conv2.bias",
|
| 543 |
+
"control_model.middle_block.1.norm.weight": "blocks.29.norm.weight",
|
| 544 |
+
"control_model.middle_block.1.norm.bias": "blocks.29.norm.bias",
|
| 545 |
+
"control_model.middle_block.1.proj_in.weight": "blocks.29.proj_in.weight",
|
| 546 |
+
"control_model.middle_block.1.proj_in.bias": "blocks.29.proj_in.bias",
|
| 547 |
+
"control_model.middle_block.1.transformer_blocks.0.attn1.to_q.weight": "blocks.29.transformer_blocks.0.attn1.to_q.weight",
|
| 548 |
+
"control_model.middle_block.1.transformer_blocks.0.attn1.to_k.weight": "blocks.29.transformer_blocks.0.attn1.to_k.weight",
|
| 549 |
+
"control_model.middle_block.1.transformer_blocks.0.attn1.to_v.weight": "blocks.29.transformer_blocks.0.attn1.to_v.weight",
|
| 550 |
+
"control_model.middle_block.1.transformer_blocks.0.attn1.to_out.0.weight": "blocks.29.transformer_blocks.0.attn1.to_out.weight",
|
| 551 |
+
"control_model.middle_block.1.transformer_blocks.0.attn1.to_out.0.bias": "blocks.29.transformer_blocks.0.attn1.to_out.bias",
|
| 552 |
+
"control_model.middle_block.1.transformer_blocks.0.ff.net.0.proj.weight": "blocks.29.transformer_blocks.0.act_fn.proj.weight",
|
| 553 |
+
"control_model.middle_block.1.transformer_blocks.0.ff.net.0.proj.bias": "blocks.29.transformer_blocks.0.act_fn.proj.bias",
|
| 554 |
+
"control_model.middle_block.1.transformer_blocks.0.ff.net.2.weight": "blocks.29.transformer_blocks.0.ff.weight",
|
| 555 |
+
"control_model.middle_block.1.transformer_blocks.0.ff.net.2.bias": "blocks.29.transformer_blocks.0.ff.bias",
|
| 556 |
+
"control_model.middle_block.1.transformer_blocks.0.attn2.to_q.weight": "blocks.29.transformer_blocks.0.attn2.to_q.weight",
|
| 557 |
+
"control_model.middle_block.1.transformer_blocks.0.attn2.to_k.weight": "blocks.29.transformer_blocks.0.attn2.to_k.weight",
|
| 558 |
+
"control_model.middle_block.1.transformer_blocks.0.attn2.to_v.weight": "blocks.29.transformer_blocks.0.attn2.to_v.weight",
|
| 559 |
+
"control_model.middle_block.1.transformer_blocks.0.attn2.to_out.0.weight": "blocks.29.transformer_blocks.0.attn2.to_out.weight",
|
| 560 |
+
"control_model.middle_block.1.transformer_blocks.0.attn2.to_out.0.bias": "blocks.29.transformer_blocks.0.attn2.to_out.bias",
|
| 561 |
+
"control_model.middle_block.1.transformer_blocks.0.norm1.weight": "blocks.29.transformer_blocks.0.norm1.weight",
|
| 562 |
+
"control_model.middle_block.1.transformer_blocks.0.norm1.bias": "blocks.29.transformer_blocks.0.norm1.bias",
|
| 563 |
+
"control_model.middle_block.1.transformer_blocks.0.norm2.weight": "blocks.29.transformer_blocks.0.norm2.weight",
|
| 564 |
+
"control_model.middle_block.1.transformer_blocks.0.norm2.bias": "blocks.29.transformer_blocks.0.norm2.bias",
|
| 565 |
+
"control_model.middle_block.1.transformer_blocks.0.norm3.weight": "blocks.29.transformer_blocks.0.norm3.weight",
|
| 566 |
+
"control_model.middle_block.1.transformer_blocks.0.norm3.bias": "blocks.29.transformer_blocks.0.norm3.bias",
|
| 567 |
+
"control_model.middle_block.1.proj_out.weight": "blocks.29.proj_out.weight",
|
| 568 |
+
"control_model.middle_block.1.proj_out.bias": "blocks.29.proj_out.bias",
|
| 569 |
+
"control_model.middle_block.2.in_layers.0.weight": "blocks.30.norm1.weight",
|
| 570 |
+
"control_model.middle_block.2.in_layers.0.bias": "blocks.30.norm1.bias",
|
| 571 |
+
"control_model.middle_block.2.in_layers.2.weight": "blocks.30.conv1.weight",
|
| 572 |
+
"control_model.middle_block.2.in_layers.2.bias": "blocks.30.conv1.bias",
|
| 573 |
+
"control_model.middle_block.2.emb_layers.1.weight": "blocks.30.time_emb_proj.weight",
|
| 574 |
+
"control_model.middle_block.2.emb_layers.1.bias": "blocks.30.time_emb_proj.bias",
|
| 575 |
+
"control_model.middle_block.2.out_layers.0.weight": "blocks.30.norm2.weight",
|
| 576 |
+
"control_model.middle_block.2.out_layers.0.bias": "blocks.30.norm2.bias",
|
| 577 |
+
"control_model.middle_block.2.out_layers.3.weight": "blocks.30.conv2.weight",
|
| 578 |
+
"control_model.middle_block.2.out_layers.3.bias": "blocks.30.conv2.bias",
|
| 579 |
+
"control_model.middle_block_out.0.weight": "controlnet_blocks.12.weight",
|
| 580 |
+
"control_model.middle_block_out.0.bias": "controlnet_blocks.7.bias",
|
| 581 |
+
}
|
| 582 |
+
state_dict_ = {}
|
| 583 |
+
for name in state_dict:
|
| 584 |
+
if name in rename_dict:
|
| 585 |
+
param = state_dict[name]
|
| 586 |
+
if ".proj_in." in name or ".proj_out." in name:
|
| 587 |
+
param = param.squeeze()
|
| 588 |
+
state_dict_[rename_dict[name]] = param
|
| 589 |
+
return state_dict_
|
diffsynth/models/sd_ipadapter.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .svd_image_encoder import SVDImageEncoder
|
| 2 |
+
from .sdxl_ipadapter import IpAdapterImageProjModel, IpAdapterModule, SDXLIpAdapterStateDictConverter
|
| 3 |
+
from transformers import CLIPImageProcessor
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class IpAdapterCLIPImageEmbedder(SVDImageEncoder):
|
| 8 |
+
def __init__(self):
|
| 9 |
+
super().__init__()
|
| 10 |
+
self.image_processor = CLIPImageProcessor()
|
| 11 |
+
|
| 12 |
+
def forward(self, image):
|
| 13 |
+
pixel_values = self.image_processor(images=image, return_tensors="pt").pixel_values
|
| 14 |
+
pixel_values = pixel_values.to(device=self.embeddings.class_embedding.device, dtype=self.embeddings.class_embedding.dtype)
|
| 15 |
+
return super().forward(pixel_values)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class SDIpAdapter(torch.nn.Module):
|
| 19 |
+
def __init__(self):
|
| 20 |
+
super().__init__()
|
| 21 |
+
shape_list = [(768, 320)] * 2 + [(768, 640)] * 2 + [(768, 1280)] * 5 + [(768, 640)] * 3 + [(768, 320)] * 3 + [(768, 1280)] * 1
|
| 22 |
+
self.ipadapter_modules = torch.nn.ModuleList([IpAdapterModule(*shape) for shape in shape_list])
|
| 23 |
+
self.image_proj = IpAdapterImageProjModel(cross_attention_dim=768, clip_embeddings_dim=1024, clip_extra_context_tokens=4)
|
| 24 |
+
self.set_full_adapter()
|
| 25 |
+
|
| 26 |
+
def set_full_adapter(self):
|
| 27 |
+
block_ids = [1, 4, 9, 12, 17, 20, 40, 43, 46, 50, 53, 56, 60, 63, 66, 29]
|
| 28 |
+
self.call_block_id = {(i, 0): j for j, i in enumerate(block_ids)}
|
| 29 |
+
|
| 30 |
+
def set_less_adapter(self):
|
| 31 |
+
# IP-Adapter for SD v1.5 doesn't support this feature.
|
| 32 |
+
self.set_full_adapter()
|
| 33 |
+
|
| 34 |
+
def forward(self, hidden_states, scale=1.0):
|
| 35 |
+
hidden_states = self.image_proj(hidden_states)
|
| 36 |
+
hidden_states = hidden_states.view(1, -1, hidden_states.shape[-1])
|
| 37 |
+
ip_kv_dict = {}
|
| 38 |
+
for (block_id, transformer_id) in self.call_block_id:
|
| 39 |
+
ipadapter_id = self.call_block_id[(block_id, transformer_id)]
|
| 40 |
+
ip_k, ip_v = self.ipadapter_modules[ipadapter_id](hidden_states)
|
| 41 |
+
if block_id not in ip_kv_dict:
|
| 42 |
+
ip_kv_dict[block_id] = {}
|
| 43 |
+
ip_kv_dict[block_id][transformer_id] = {
|
| 44 |
+
"ip_k": ip_k,
|
| 45 |
+
"ip_v": ip_v,
|
| 46 |
+
"scale": scale
|
| 47 |
+
}
|
| 48 |
+
return ip_kv_dict
|
| 49 |
+
|
| 50 |
+
@staticmethod
|
| 51 |
+
def state_dict_converter():
|
| 52 |
+
return SDIpAdapterStateDictConverter()
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class SDIpAdapterStateDictConverter(SDXLIpAdapterStateDictConverter):
|
| 56 |
+
def __init__(self):
|
| 57 |
+
pass
|
diffsynth/models/sd_motion.py
ADDED
|
@@ -0,0 +1,199 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sd_unet import SDUNet, Attention, GEGLU
|
| 2 |
+
import torch
|
| 3 |
+
from einops import rearrange, repeat
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class TemporalTransformerBlock(torch.nn.Module):
|
| 7 |
+
|
| 8 |
+
def __init__(self, dim, num_attention_heads, attention_head_dim, max_position_embeddings=32):
|
| 9 |
+
super().__init__()
|
| 10 |
+
|
| 11 |
+
# 1. Self-Attn
|
| 12 |
+
self.pe1 = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, dim))
|
| 13 |
+
self.norm1 = torch.nn.LayerNorm(dim, elementwise_affine=True)
|
| 14 |
+
self.attn1 = Attention(q_dim=dim, num_heads=num_attention_heads, head_dim=attention_head_dim, bias_out=True)
|
| 15 |
+
|
| 16 |
+
# 2. Cross-Attn
|
| 17 |
+
self.pe2 = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, dim))
|
| 18 |
+
self.norm2 = torch.nn.LayerNorm(dim, elementwise_affine=True)
|
| 19 |
+
self.attn2 = Attention(q_dim=dim, num_heads=num_attention_heads, head_dim=attention_head_dim, bias_out=True)
|
| 20 |
+
|
| 21 |
+
# 3. Feed-forward
|
| 22 |
+
self.norm3 = torch.nn.LayerNorm(dim, elementwise_affine=True)
|
| 23 |
+
self.act_fn = GEGLU(dim, dim * 4)
|
| 24 |
+
self.ff = torch.nn.Linear(dim * 4, dim)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def forward(self, hidden_states, batch_size=1):
|
| 28 |
+
|
| 29 |
+
# 1. Self-Attention
|
| 30 |
+
norm_hidden_states = self.norm1(hidden_states)
|
| 31 |
+
norm_hidden_states = rearrange(norm_hidden_states, "(b f) h c -> (b h) f c", b=batch_size)
|
| 32 |
+
attn_output = self.attn1(norm_hidden_states + self.pe1[:, :norm_hidden_states.shape[1]])
|
| 33 |
+
attn_output = rearrange(attn_output, "(b h) f c -> (b f) h c", b=batch_size)
|
| 34 |
+
hidden_states = attn_output + hidden_states
|
| 35 |
+
|
| 36 |
+
# 2. Cross-Attention
|
| 37 |
+
norm_hidden_states = self.norm2(hidden_states)
|
| 38 |
+
norm_hidden_states = rearrange(norm_hidden_states, "(b f) h c -> (b h) f c", b=batch_size)
|
| 39 |
+
attn_output = self.attn2(norm_hidden_states + self.pe2[:, :norm_hidden_states.shape[1]])
|
| 40 |
+
attn_output = rearrange(attn_output, "(b h) f c -> (b f) h c", b=batch_size)
|
| 41 |
+
hidden_states = attn_output + hidden_states
|
| 42 |
+
|
| 43 |
+
# 3. Feed-forward
|
| 44 |
+
norm_hidden_states = self.norm3(hidden_states)
|
| 45 |
+
ff_output = self.act_fn(norm_hidden_states)
|
| 46 |
+
ff_output = self.ff(ff_output)
|
| 47 |
+
hidden_states = ff_output + hidden_states
|
| 48 |
+
|
| 49 |
+
return hidden_states
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class TemporalBlock(torch.nn.Module):
|
| 53 |
+
|
| 54 |
+
def __init__(self, num_attention_heads, attention_head_dim, in_channels, num_layers=1, norm_num_groups=32, eps=1e-5):
|
| 55 |
+
super().__init__()
|
| 56 |
+
inner_dim = num_attention_heads * attention_head_dim
|
| 57 |
+
|
| 58 |
+
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=eps, affine=True)
|
| 59 |
+
self.proj_in = torch.nn.Linear(in_channels, inner_dim)
|
| 60 |
+
|
| 61 |
+
self.transformer_blocks = torch.nn.ModuleList([
|
| 62 |
+
TemporalTransformerBlock(
|
| 63 |
+
inner_dim,
|
| 64 |
+
num_attention_heads,
|
| 65 |
+
attention_head_dim
|
| 66 |
+
)
|
| 67 |
+
for d in range(num_layers)
|
| 68 |
+
])
|
| 69 |
+
|
| 70 |
+
self.proj_out = torch.nn.Linear(inner_dim, in_channels)
|
| 71 |
+
|
| 72 |
+
def forward(self, hidden_states, time_emb, text_emb, res_stack, batch_size=1):
|
| 73 |
+
batch, _, height, width = hidden_states.shape
|
| 74 |
+
residual = hidden_states
|
| 75 |
+
|
| 76 |
+
hidden_states = self.norm(hidden_states)
|
| 77 |
+
inner_dim = hidden_states.shape[1]
|
| 78 |
+
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim)
|
| 79 |
+
hidden_states = self.proj_in(hidden_states)
|
| 80 |
+
|
| 81 |
+
for block in self.transformer_blocks:
|
| 82 |
+
hidden_states = block(
|
| 83 |
+
hidden_states,
|
| 84 |
+
batch_size=batch_size
|
| 85 |
+
)
|
| 86 |
+
|
| 87 |
+
hidden_states = self.proj_out(hidden_states)
|
| 88 |
+
hidden_states = hidden_states.reshape(batch, height, width, inner_dim).permute(0, 3, 1, 2).contiguous()
|
| 89 |
+
hidden_states = hidden_states + residual
|
| 90 |
+
|
| 91 |
+
return hidden_states, time_emb, text_emb, res_stack
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
class SDMotionModel(torch.nn.Module):
|
| 95 |
+
def __init__(self):
|
| 96 |
+
super().__init__()
|
| 97 |
+
self.motion_modules = torch.nn.ModuleList([
|
| 98 |
+
TemporalBlock(8, 40, 320, eps=1e-6),
|
| 99 |
+
TemporalBlock(8, 40, 320, eps=1e-6),
|
| 100 |
+
TemporalBlock(8, 80, 640, eps=1e-6),
|
| 101 |
+
TemporalBlock(8, 80, 640, eps=1e-6),
|
| 102 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 103 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 104 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 105 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 106 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 107 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 108 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 109 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 110 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 111 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 112 |
+
TemporalBlock(8, 160, 1280, eps=1e-6),
|
| 113 |
+
TemporalBlock(8, 80, 640, eps=1e-6),
|
| 114 |
+
TemporalBlock(8, 80, 640, eps=1e-6),
|
| 115 |
+
TemporalBlock(8, 80, 640, eps=1e-6),
|
| 116 |
+
TemporalBlock(8, 40, 320, eps=1e-6),
|
| 117 |
+
TemporalBlock(8, 40, 320, eps=1e-6),
|
| 118 |
+
TemporalBlock(8, 40, 320, eps=1e-6),
|
| 119 |
+
])
|
| 120 |
+
self.call_block_id = {
|
| 121 |
+
1: 0,
|
| 122 |
+
4: 1,
|
| 123 |
+
9: 2,
|
| 124 |
+
12: 3,
|
| 125 |
+
17: 4,
|
| 126 |
+
20: 5,
|
| 127 |
+
24: 6,
|
| 128 |
+
26: 7,
|
| 129 |
+
29: 8,
|
| 130 |
+
32: 9,
|
| 131 |
+
34: 10,
|
| 132 |
+
36: 11,
|
| 133 |
+
40: 12,
|
| 134 |
+
43: 13,
|
| 135 |
+
46: 14,
|
| 136 |
+
50: 15,
|
| 137 |
+
53: 16,
|
| 138 |
+
56: 17,
|
| 139 |
+
60: 18,
|
| 140 |
+
63: 19,
|
| 141 |
+
66: 20
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
def forward(self):
|
| 145 |
+
pass
|
| 146 |
+
|
| 147 |
+
@staticmethod
|
| 148 |
+
def state_dict_converter():
|
| 149 |
+
return SDMotionModelStateDictConverter()
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
class SDMotionModelStateDictConverter:
|
| 153 |
+
def __init__(self):
|
| 154 |
+
pass
|
| 155 |
+
|
| 156 |
+
def from_diffusers(self, state_dict):
|
| 157 |
+
rename_dict = {
|
| 158 |
+
"norm": "norm",
|
| 159 |
+
"proj_in": "proj_in",
|
| 160 |
+
"transformer_blocks.0.attention_blocks.0.to_q": "transformer_blocks.0.attn1.to_q",
|
| 161 |
+
"transformer_blocks.0.attention_blocks.0.to_k": "transformer_blocks.0.attn1.to_k",
|
| 162 |
+
"transformer_blocks.0.attention_blocks.0.to_v": "transformer_blocks.0.attn1.to_v",
|
| 163 |
+
"transformer_blocks.0.attention_blocks.0.to_out.0": "transformer_blocks.0.attn1.to_out",
|
| 164 |
+
"transformer_blocks.0.attention_blocks.0.pos_encoder": "transformer_blocks.0.pe1",
|
| 165 |
+
"transformer_blocks.0.attention_blocks.1.to_q": "transformer_blocks.0.attn2.to_q",
|
| 166 |
+
"transformer_blocks.0.attention_blocks.1.to_k": "transformer_blocks.0.attn2.to_k",
|
| 167 |
+
"transformer_blocks.0.attention_blocks.1.to_v": "transformer_blocks.0.attn2.to_v",
|
| 168 |
+
"transformer_blocks.0.attention_blocks.1.to_out.0": "transformer_blocks.0.attn2.to_out",
|
| 169 |
+
"transformer_blocks.0.attention_blocks.1.pos_encoder": "transformer_blocks.0.pe2",
|
| 170 |
+
"transformer_blocks.0.norms.0": "transformer_blocks.0.norm1",
|
| 171 |
+
"transformer_blocks.0.norms.1": "transformer_blocks.0.norm2",
|
| 172 |
+
"transformer_blocks.0.ff.net.0.proj": "transformer_blocks.0.act_fn.proj",
|
| 173 |
+
"transformer_blocks.0.ff.net.2": "transformer_blocks.0.ff",
|
| 174 |
+
"transformer_blocks.0.ff_norm": "transformer_blocks.0.norm3",
|
| 175 |
+
"proj_out": "proj_out",
|
| 176 |
+
}
|
| 177 |
+
name_list = sorted([i for i in state_dict if i.startswith("down_blocks.")])
|
| 178 |
+
name_list += sorted([i for i in state_dict if i.startswith("mid_block.")])
|
| 179 |
+
name_list += sorted([i for i in state_dict if i.startswith("up_blocks.")])
|
| 180 |
+
state_dict_ = {}
|
| 181 |
+
last_prefix, module_id = "", -1
|
| 182 |
+
for name in name_list:
|
| 183 |
+
names = name.split(".")
|
| 184 |
+
prefix_index = names.index("temporal_transformer") + 1
|
| 185 |
+
prefix = ".".join(names[:prefix_index])
|
| 186 |
+
if prefix != last_prefix:
|
| 187 |
+
last_prefix = prefix
|
| 188 |
+
module_id += 1
|
| 189 |
+
middle_name = ".".join(names[prefix_index:-1])
|
| 190 |
+
suffix = names[-1]
|
| 191 |
+
if "pos_encoder" in names:
|
| 192 |
+
rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name]])
|
| 193 |
+
else:
|
| 194 |
+
rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name], suffix])
|
| 195 |
+
state_dict_[rename] = state_dict[name]
|
| 196 |
+
return state_dict_
|
| 197 |
+
|
| 198 |
+
def from_civitai(self, state_dict):
|
| 199 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/sd_text_encoder.py
ADDED
|
@@ -0,0 +1,321 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .attention import Attention
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class CLIPEncoderLayer(torch.nn.Module):
|
| 6 |
+
def __init__(self, embed_dim, intermediate_size, num_heads=12, head_dim=64, use_quick_gelu=True):
|
| 7 |
+
super().__init__()
|
| 8 |
+
self.attn = Attention(q_dim=embed_dim, num_heads=num_heads, head_dim=head_dim, bias_q=True, bias_kv=True, bias_out=True)
|
| 9 |
+
self.layer_norm1 = torch.nn.LayerNorm(embed_dim)
|
| 10 |
+
self.layer_norm2 = torch.nn.LayerNorm(embed_dim)
|
| 11 |
+
self.fc1 = torch.nn.Linear(embed_dim, intermediate_size)
|
| 12 |
+
self.fc2 = torch.nn.Linear(intermediate_size, embed_dim)
|
| 13 |
+
|
| 14 |
+
self.use_quick_gelu = use_quick_gelu
|
| 15 |
+
|
| 16 |
+
def quickGELU(self, x):
|
| 17 |
+
return x * torch.sigmoid(1.702 * x)
|
| 18 |
+
|
| 19 |
+
def forward(self, hidden_states, attn_mask=None):
|
| 20 |
+
residual = hidden_states
|
| 21 |
+
|
| 22 |
+
hidden_states = self.layer_norm1(hidden_states)
|
| 23 |
+
hidden_states = self.attn(hidden_states, attn_mask=attn_mask)
|
| 24 |
+
hidden_states = residual + hidden_states
|
| 25 |
+
|
| 26 |
+
residual = hidden_states
|
| 27 |
+
hidden_states = self.layer_norm2(hidden_states)
|
| 28 |
+
hidden_states = self.fc1(hidden_states)
|
| 29 |
+
if self.use_quick_gelu:
|
| 30 |
+
hidden_states = self.quickGELU(hidden_states)
|
| 31 |
+
else:
|
| 32 |
+
hidden_states = torch.nn.functional.gelu(hidden_states)
|
| 33 |
+
hidden_states = self.fc2(hidden_states)
|
| 34 |
+
hidden_states = residual + hidden_states
|
| 35 |
+
|
| 36 |
+
return hidden_states
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
class SDTextEncoder(torch.nn.Module):
|
| 40 |
+
def __init__(self, embed_dim=768, vocab_size=49408, max_position_embeddings=77, num_encoder_layers=12, encoder_intermediate_size=3072):
|
| 41 |
+
super().__init__()
|
| 42 |
+
|
| 43 |
+
# token_embedding
|
| 44 |
+
self.token_embedding = torch.nn.Embedding(vocab_size, embed_dim)
|
| 45 |
+
|
| 46 |
+
# position_embeds (This is a fixed tensor)
|
| 47 |
+
self.position_embeds = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, embed_dim))
|
| 48 |
+
|
| 49 |
+
# encoders
|
| 50 |
+
self.encoders = torch.nn.ModuleList([CLIPEncoderLayer(embed_dim, encoder_intermediate_size) for _ in range(num_encoder_layers)])
|
| 51 |
+
|
| 52 |
+
# attn_mask
|
| 53 |
+
self.attn_mask = self.attention_mask(max_position_embeddings)
|
| 54 |
+
|
| 55 |
+
# final_layer_norm
|
| 56 |
+
self.final_layer_norm = torch.nn.LayerNorm(embed_dim)
|
| 57 |
+
|
| 58 |
+
def attention_mask(self, length):
|
| 59 |
+
mask = torch.empty(length, length)
|
| 60 |
+
mask.fill_(float("-inf"))
|
| 61 |
+
mask.triu_(1)
|
| 62 |
+
return mask
|
| 63 |
+
|
| 64 |
+
def forward(self, input_ids, clip_skip=1):
|
| 65 |
+
embeds = self.token_embedding(input_ids) + self.position_embeds
|
| 66 |
+
attn_mask = self.attn_mask.to(device=embeds.device, dtype=embeds.dtype)
|
| 67 |
+
for encoder_id, encoder in enumerate(self.encoders):
|
| 68 |
+
embeds = encoder(embeds, attn_mask=attn_mask)
|
| 69 |
+
if encoder_id + clip_skip == len(self.encoders):
|
| 70 |
+
break
|
| 71 |
+
embeds = self.final_layer_norm(embeds)
|
| 72 |
+
return embeds
|
| 73 |
+
|
| 74 |
+
@staticmethod
|
| 75 |
+
def state_dict_converter():
|
| 76 |
+
return SDTextEncoderStateDictConverter()
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class SDTextEncoderStateDictConverter:
|
| 80 |
+
def __init__(self):
|
| 81 |
+
pass
|
| 82 |
+
|
| 83 |
+
def from_diffusers(self, state_dict):
|
| 84 |
+
rename_dict = {
|
| 85 |
+
"text_model.embeddings.token_embedding.weight": "token_embedding.weight",
|
| 86 |
+
"text_model.embeddings.position_embedding.weight": "position_embeds",
|
| 87 |
+
"text_model.final_layer_norm.weight": "final_layer_norm.weight",
|
| 88 |
+
"text_model.final_layer_norm.bias": "final_layer_norm.bias"
|
| 89 |
+
}
|
| 90 |
+
attn_rename_dict = {
|
| 91 |
+
"self_attn.q_proj": "attn.to_q",
|
| 92 |
+
"self_attn.k_proj": "attn.to_k",
|
| 93 |
+
"self_attn.v_proj": "attn.to_v",
|
| 94 |
+
"self_attn.out_proj": "attn.to_out",
|
| 95 |
+
"layer_norm1": "layer_norm1",
|
| 96 |
+
"layer_norm2": "layer_norm2",
|
| 97 |
+
"mlp.fc1": "fc1",
|
| 98 |
+
"mlp.fc2": "fc2",
|
| 99 |
+
}
|
| 100 |
+
state_dict_ = {}
|
| 101 |
+
for name in state_dict:
|
| 102 |
+
if name in rename_dict:
|
| 103 |
+
param = state_dict[name]
|
| 104 |
+
if name == "text_model.embeddings.position_embedding.weight":
|
| 105 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 106 |
+
state_dict_[rename_dict[name]] = param
|
| 107 |
+
elif name.startswith("text_model.encoder.layers."):
|
| 108 |
+
param = state_dict[name]
|
| 109 |
+
names = name.split(".")
|
| 110 |
+
layer_id, layer_type, tail = names[3], ".".join(names[4:-1]), names[-1]
|
| 111 |
+
name_ = ".".join(["encoders", layer_id, attn_rename_dict[layer_type], tail])
|
| 112 |
+
state_dict_[name_] = param
|
| 113 |
+
return state_dict_
|
| 114 |
+
|
| 115 |
+
def from_civitai(self, state_dict):
|
| 116 |
+
rename_dict = {
|
| 117 |
+
"cond_stage_model.transformer.text_model.embeddings.token_embedding.weight": "token_embedding.weight",
|
| 118 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm1.bias": "encoders.0.layer_norm1.bias",
|
| 119 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm1.weight": "encoders.0.layer_norm1.weight",
|
| 120 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm2.bias": "encoders.0.layer_norm2.bias",
|
| 121 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm2.weight": "encoders.0.layer_norm2.weight",
|
| 122 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.mlp.fc1.bias": "encoders.0.fc1.bias",
|
| 123 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.mlp.fc1.weight": "encoders.0.fc1.weight",
|
| 124 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.mlp.fc2.bias": "encoders.0.fc2.bias",
|
| 125 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.mlp.fc2.weight": "encoders.0.fc2.weight",
|
| 126 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.k_proj.bias": "encoders.0.attn.to_k.bias",
|
| 127 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.k_proj.weight": "encoders.0.attn.to_k.weight",
|
| 128 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.out_proj.bias": "encoders.0.attn.to_out.bias",
|
| 129 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.out_proj.weight": "encoders.0.attn.to_out.weight",
|
| 130 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.q_proj.bias": "encoders.0.attn.to_q.bias",
|
| 131 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.q_proj.weight": "encoders.0.attn.to_q.weight",
|
| 132 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.v_proj.bias": "encoders.0.attn.to_v.bias",
|
| 133 |
+
"cond_stage_model.transformer.text_model.encoder.layers.0.self_attn.v_proj.weight": "encoders.0.attn.to_v.weight",
|
| 134 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.layer_norm1.bias": "encoders.1.layer_norm1.bias",
|
| 135 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.layer_norm1.weight": "encoders.1.layer_norm1.weight",
|
| 136 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.layer_norm2.bias": "encoders.1.layer_norm2.bias",
|
| 137 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.layer_norm2.weight": "encoders.1.layer_norm2.weight",
|
| 138 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.mlp.fc1.bias": "encoders.1.fc1.bias",
|
| 139 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.mlp.fc1.weight": "encoders.1.fc1.weight",
|
| 140 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.mlp.fc2.bias": "encoders.1.fc2.bias",
|
| 141 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.mlp.fc2.weight": "encoders.1.fc2.weight",
|
| 142 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.k_proj.bias": "encoders.1.attn.to_k.bias",
|
| 143 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.k_proj.weight": "encoders.1.attn.to_k.weight",
|
| 144 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.out_proj.bias": "encoders.1.attn.to_out.bias",
|
| 145 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.out_proj.weight": "encoders.1.attn.to_out.weight",
|
| 146 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.q_proj.bias": "encoders.1.attn.to_q.bias",
|
| 147 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.q_proj.weight": "encoders.1.attn.to_q.weight",
|
| 148 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.v_proj.bias": "encoders.1.attn.to_v.bias",
|
| 149 |
+
"cond_stage_model.transformer.text_model.encoder.layers.1.self_attn.v_proj.weight": "encoders.1.attn.to_v.weight",
|
| 150 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.layer_norm1.bias": "encoders.10.layer_norm1.bias",
|
| 151 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.layer_norm1.weight": "encoders.10.layer_norm1.weight",
|
| 152 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.layer_norm2.bias": "encoders.10.layer_norm2.bias",
|
| 153 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.layer_norm2.weight": "encoders.10.layer_norm2.weight",
|
| 154 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.mlp.fc1.bias": "encoders.10.fc1.bias",
|
| 155 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.mlp.fc1.weight": "encoders.10.fc1.weight",
|
| 156 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.mlp.fc2.bias": "encoders.10.fc2.bias",
|
| 157 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.mlp.fc2.weight": "encoders.10.fc2.weight",
|
| 158 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.k_proj.bias": "encoders.10.attn.to_k.bias",
|
| 159 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.k_proj.weight": "encoders.10.attn.to_k.weight",
|
| 160 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.out_proj.bias": "encoders.10.attn.to_out.bias",
|
| 161 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.out_proj.weight": "encoders.10.attn.to_out.weight",
|
| 162 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.q_proj.bias": "encoders.10.attn.to_q.bias",
|
| 163 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.q_proj.weight": "encoders.10.attn.to_q.weight",
|
| 164 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.v_proj.bias": "encoders.10.attn.to_v.bias",
|
| 165 |
+
"cond_stage_model.transformer.text_model.encoder.layers.10.self_attn.v_proj.weight": "encoders.10.attn.to_v.weight",
|
| 166 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.layer_norm1.bias": "encoders.11.layer_norm1.bias",
|
| 167 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.layer_norm1.weight": "encoders.11.layer_norm1.weight",
|
| 168 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.layer_norm2.bias": "encoders.11.layer_norm2.bias",
|
| 169 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.layer_norm2.weight": "encoders.11.layer_norm2.weight",
|
| 170 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.mlp.fc1.bias": "encoders.11.fc1.bias",
|
| 171 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.mlp.fc1.weight": "encoders.11.fc1.weight",
|
| 172 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.mlp.fc2.bias": "encoders.11.fc2.bias",
|
| 173 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.mlp.fc2.weight": "encoders.11.fc2.weight",
|
| 174 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.k_proj.bias": "encoders.11.attn.to_k.bias",
|
| 175 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.k_proj.weight": "encoders.11.attn.to_k.weight",
|
| 176 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.out_proj.bias": "encoders.11.attn.to_out.bias",
|
| 177 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.out_proj.weight": "encoders.11.attn.to_out.weight",
|
| 178 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.q_proj.bias": "encoders.11.attn.to_q.bias",
|
| 179 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.q_proj.weight": "encoders.11.attn.to_q.weight",
|
| 180 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.v_proj.bias": "encoders.11.attn.to_v.bias",
|
| 181 |
+
"cond_stage_model.transformer.text_model.encoder.layers.11.self_attn.v_proj.weight": "encoders.11.attn.to_v.weight",
|
| 182 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.layer_norm1.bias": "encoders.2.layer_norm1.bias",
|
| 183 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.layer_norm1.weight": "encoders.2.layer_norm1.weight",
|
| 184 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.layer_norm2.bias": "encoders.2.layer_norm2.bias",
|
| 185 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.layer_norm2.weight": "encoders.2.layer_norm2.weight",
|
| 186 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.mlp.fc1.bias": "encoders.2.fc1.bias",
|
| 187 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.mlp.fc1.weight": "encoders.2.fc1.weight",
|
| 188 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.mlp.fc2.bias": "encoders.2.fc2.bias",
|
| 189 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.mlp.fc2.weight": "encoders.2.fc2.weight",
|
| 190 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.k_proj.bias": "encoders.2.attn.to_k.bias",
|
| 191 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.k_proj.weight": "encoders.2.attn.to_k.weight",
|
| 192 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.out_proj.bias": "encoders.2.attn.to_out.bias",
|
| 193 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.out_proj.weight": "encoders.2.attn.to_out.weight",
|
| 194 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.q_proj.bias": "encoders.2.attn.to_q.bias",
|
| 195 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.q_proj.weight": "encoders.2.attn.to_q.weight",
|
| 196 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.v_proj.bias": "encoders.2.attn.to_v.bias",
|
| 197 |
+
"cond_stage_model.transformer.text_model.encoder.layers.2.self_attn.v_proj.weight": "encoders.2.attn.to_v.weight",
|
| 198 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.layer_norm1.bias": "encoders.3.layer_norm1.bias",
|
| 199 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.layer_norm1.weight": "encoders.3.layer_norm1.weight",
|
| 200 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.layer_norm2.bias": "encoders.3.layer_norm2.bias",
|
| 201 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.layer_norm2.weight": "encoders.3.layer_norm2.weight",
|
| 202 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.mlp.fc1.bias": "encoders.3.fc1.bias",
|
| 203 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.mlp.fc1.weight": "encoders.3.fc1.weight",
|
| 204 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.mlp.fc2.bias": "encoders.3.fc2.bias",
|
| 205 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.mlp.fc2.weight": "encoders.3.fc2.weight",
|
| 206 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.k_proj.bias": "encoders.3.attn.to_k.bias",
|
| 207 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.k_proj.weight": "encoders.3.attn.to_k.weight",
|
| 208 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.out_proj.bias": "encoders.3.attn.to_out.bias",
|
| 209 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.out_proj.weight": "encoders.3.attn.to_out.weight",
|
| 210 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.q_proj.bias": "encoders.3.attn.to_q.bias",
|
| 211 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.q_proj.weight": "encoders.3.attn.to_q.weight",
|
| 212 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.v_proj.bias": "encoders.3.attn.to_v.bias",
|
| 213 |
+
"cond_stage_model.transformer.text_model.encoder.layers.3.self_attn.v_proj.weight": "encoders.3.attn.to_v.weight",
|
| 214 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.layer_norm1.bias": "encoders.4.layer_norm1.bias",
|
| 215 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.layer_norm1.weight": "encoders.4.layer_norm1.weight",
|
| 216 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.layer_norm2.bias": "encoders.4.layer_norm2.bias",
|
| 217 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.layer_norm2.weight": "encoders.4.layer_norm2.weight",
|
| 218 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.mlp.fc1.bias": "encoders.4.fc1.bias",
|
| 219 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.mlp.fc1.weight": "encoders.4.fc1.weight",
|
| 220 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.mlp.fc2.bias": "encoders.4.fc2.bias",
|
| 221 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.mlp.fc2.weight": "encoders.4.fc2.weight",
|
| 222 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.k_proj.bias": "encoders.4.attn.to_k.bias",
|
| 223 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.k_proj.weight": "encoders.4.attn.to_k.weight",
|
| 224 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.out_proj.bias": "encoders.4.attn.to_out.bias",
|
| 225 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.out_proj.weight": "encoders.4.attn.to_out.weight",
|
| 226 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.q_proj.bias": "encoders.4.attn.to_q.bias",
|
| 227 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.q_proj.weight": "encoders.4.attn.to_q.weight",
|
| 228 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.v_proj.bias": "encoders.4.attn.to_v.bias",
|
| 229 |
+
"cond_stage_model.transformer.text_model.encoder.layers.4.self_attn.v_proj.weight": "encoders.4.attn.to_v.weight",
|
| 230 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.layer_norm1.bias": "encoders.5.layer_norm1.bias",
|
| 231 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.layer_norm1.weight": "encoders.5.layer_norm1.weight",
|
| 232 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.layer_norm2.bias": "encoders.5.layer_norm2.bias",
|
| 233 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.layer_norm2.weight": "encoders.5.layer_norm2.weight",
|
| 234 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.mlp.fc1.bias": "encoders.5.fc1.bias",
|
| 235 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.mlp.fc1.weight": "encoders.5.fc1.weight",
|
| 236 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.mlp.fc2.bias": "encoders.5.fc2.bias",
|
| 237 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.mlp.fc2.weight": "encoders.5.fc2.weight",
|
| 238 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.k_proj.bias": "encoders.5.attn.to_k.bias",
|
| 239 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.k_proj.weight": "encoders.5.attn.to_k.weight",
|
| 240 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.out_proj.bias": "encoders.5.attn.to_out.bias",
|
| 241 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.out_proj.weight": "encoders.5.attn.to_out.weight",
|
| 242 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.q_proj.bias": "encoders.5.attn.to_q.bias",
|
| 243 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.q_proj.weight": "encoders.5.attn.to_q.weight",
|
| 244 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.v_proj.bias": "encoders.5.attn.to_v.bias",
|
| 245 |
+
"cond_stage_model.transformer.text_model.encoder.layers.5.self_attn.v_proj.weight": "encoders.5.attn.to_v.weight",
|
| 246 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.layer_norm1.bias": "encoders.6.layer_norm1.bias",
|
| 247 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.layer_norm1.weight": "encoders.6.layer_norm1.weight",
|
| 248 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.layer_norm2.bias": "encoders.6.layer_norm2.bias",
|
| 249 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.layer_norm2.weight": "encoders.6.layer_norm2.weight",
|
| 250 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.mlp.fc1.bias": "encoders.6.fc1.bias",
|
| 251 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.mlp.fc1.weight": "encoders.6.fc1.weight",
|
| 252 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.mlp.fc2.bias": "encoders.6.fc2.bias",
|
| 253 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.mlp.fc2.weight": "encoders.6.fc2.weight",
|
| 254 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.k_proj.bias": "encoders.6.attn.to_k.bias",
|
| 255 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.k_proj.weight": "encoders.6.attn.to_k.weight",
|
| 256 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.out_proj.bias": "encoders.6.attn.to_out.bias",
|
| 257 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.out_proj.weight": "encoders.6.attn.to_out.weight",
|
| 258 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.q_proj.bias": "encoders.6.attn.to_q.bias",
|
| 259 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.q_proj.weight": "encoders.6.attn.to_q.weight",
|
| 260 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.v_proj.bias": "encoders.6.attn.to_v.bias",
|
| 261 |
+
"cond_stage_model.transformer.text_model.encoder.layers.6.self_attn.v_proj.weight": "encoders.6.attn.to_v.weight",
|
| 262 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.layer_norm1.bias": "encoders.7.layer_norm1.bias",
|
| 263 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.layer_norm1.weight": "encoders.7.layer_norm1.weight",
|
| 264 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.layer_norm2.bias": "encoders.7.layer_norm2.bias",
|
| 265 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.layer_norm2.weight": "encoders.7.layer_norm2.weight",
|
| 266 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.mlp.fc1.bias": "encoders.7.fc1.bias",
|
| 267 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.mlp.fc1.weight": "encoders.7.fc1.weight",
|
| 268 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.mlp.fc2.bias": "encoders.7.fc2.bias",
|
| 269 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.mlp.fc2.weight": "encoders.7.fc2.weight",
|
| 270 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.k_proj.bias": "encoders.7.attn.to_k.bias",
|
| 271 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.k_proj.weight": "encoders.7.attn.to_k.weight",
|
| 272 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.out_proj.bias": "encoders.7.attn.to_out.bias",
|
| 273 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.out_proj.weight": "encoders.7.attn.to_out.weight",
|
| 274 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.q_proj.bias": "encoders.7.attn.to_q.bias",
|
| 275 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.q_proj.weight": "encoders.7.attn.to_q.weight",
|
| 276 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.v_proj.bias": "encoders.7.attn.to_v.bias",
|
| 277 |
+
"cond_stage_model.transformer.text_model.encoder.layers.7.self_attn.v_proj.weight": "encoders.7.attn.to_v.weight",
|
| 278 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.layer_norm1.bias": "encoders.8.layer_norm1.bias",
|
| 279 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.layer_norm1.weight": "encoders.8.layer_norm1.weight",
|
| 280 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.layer_norm2.bias": "encoders.8.layer_norm2.bias",
|
| 281 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.layer_norm2.weight": "encoders.8.layer_norm2.weight",
|
| 282 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.mlp.fc1.bias": "encoders.8.fc1.bias",
|
| 283 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.mlp.fc1.weight": "encoders.8.fc1.weight",
|
| 284 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.mlp.fc2.bias": "encoders.8.fc2.bias",
|
| 285 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.mlp.fc2.weight": "encoders.8.fc2.weight",
|
| 286 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.k_proj.bias": "encoders.8.attn.to_k.bias",
|
| 287 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.k_proj.weight": "encoders.8.attn.to_k.weight",
|
| 288 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.out_proj.bias": "encoders.8.attn.to_out.bias",
|
| 289 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.out_proj.weight": "encoders.8.attn.to_out.weight",
|
| 290 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.q_proj.bias": "encoders.8.attn.to_q.bias",
|
| 291 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.q_proj.weight": "encoders.8.attn.to_q.weight",
|
| 292 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.v_proj.bias": "encoders.8.attn.to_v.bias",
|
| 293 |
+
"cond_stage_model.transformer.text_model.encoder.layers.8.self_attn.v_proj.weight": "encoders.8.attn.to_v.weight",
|
| 294 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.layer_norm1.bias": "encoders.9.layer_norm1.bias",
|
| 295 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.layer_norm1.weight": "encoders.9.layer_norm1.weight",
|
| 296 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.layer_norm2.bias": "encoders.9.layer_norm2.bias",
|
| 297 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.layer_norm2.weight": "encoders.9.layer_norm2.weight",
|
| 298 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.mlp.fc1.bias": "encoders.9.fc1.bias",
|
| 299 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.mlp.fc1.weight": "encoders.9.fc1.weight",
|
| 300 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.mlp.fc2.bias": "encoders.9.fc2.bias",
|
| 301 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.mlp.fc2.weight": "encoders.9.fc2.weight",
|
| 302 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.k_proj.bias": "encoders.9.attn.to_k.bias",
|
| 303 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.k_proj.weight": "encoders.9.attn.to_k.weight",
|
| 304 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.out_proj.bias": "encoders.9.attn.to_out.bias",
|
| 305 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.out_proj.weight": "encoders.9.attn.to_out.weight",
|
| 306 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.q_proj.bias": "encoders.9.attn.to_q.bias",
|
| 307 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.q_proj.weight": "encoders.9.attn.to_q.weight",
|
| 308 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.v_proj.bias": "encoders.9.attn.to_v.bias",
|
| 309 |
+
"cond_stage_model.transformer.text_model.encoder.layers.9.self_attn.v_proj.weight": "encoders.9.attn.to_v.weight",
|
| 310 |
+
"cond_stage_model.transformer.text_model.final_layer_norm.bias": "final_layer_norm.bias",
|
| 311 |
+
"cond_stage_model.transformer.text_model.final_layer_norm.weight": "final_layer_norm.weight",
|
| 312 |
+
"cond_stage_model.transformer.text_model.embeddings.position_embedding.weight": "position_embeds"
|
| 313 |
+
}
|
| 314 |
+
state_dict_ = {}
|
| 315 |
+
for name in state_dict:
|
| 316 |
+
if name in rename_dict:
|
| 317 |
+
param = state_dict[name]
|
| 318 |
+
if name == "cond_stage_model.transformer.text_model.embeddings.position_embedding.weight":
|
| 319 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 320 |
+
state_dict_[rename_dict[name]] = param
|
| 321 |
+
return state_dict_
|
diffsynth/models/sd_unet.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
diffsynth/models/sd_vae_decoder.py
ADDED
|
@@ -0,0 +1,336 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .attention import Attention
|
| 3 |
+
from .sd_unet import ResnetBlock, UpSampler
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class VAEAttentionBlock(torch.nn.Module):
|
| 8 |
+
|
| 9 |
+
def __init__(self, num_attention_heads, attention_head_dim, in_channels, num_layers=1, norm_num_groups=32, eps=1e-5):
|
| 10 |
+
super().__init__()
|
| 11 |
+
inner_dim = num_attention_heads * attention_head_dim
|
| 12 |
+
|
| 13 |
+
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=eps, affine=True)
|
| 14 |
+
|
| 15 |
+
self.transformer_blocks = torch.nn.ModuleList([
|
| 16 |
+
Attention(
|
| 17 |
+
inner_dim,
|
| 18 |
+
num_attention_heads,
|
| 19 |
+
attention_head_dim,
|
| 20 |
+
bias_q=True,
|
| 21 |
+
bias_kv=True,
|
| 22 |
+
bias_out=True
|
| 23 |
+
)
|
| 24 |
+
for d in range(num_layers)
|
| 25 |
+
])
|
| 26 |
+
|
| 27 |
+
def forward(self, hidden_states, time_emb, text_emb, res_stack):
|
| 28 |
+
batch, _, height, width = hidden_states.shape
|
| 29 |
+
residual = hidden_states
|
| 30 |
+
|
| 31 |
+
hidden_states = self.norm(hidden_states)
|
| 32 |
+
inner_dim = hidden_states.shape[1]
|
| 33 |
+
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * width, inner_dim)
|
| 34 |
+
|
| 35 |
+
for block in self.transformer_blocks:
|
| 36 |
+
hidden_states = block(hidden_states)
|
| 37 |
+
|
| 38 |
+
hidden_states = hidden_states.reshape(batch, height, width, inner_dim).permute(0, 3, 1, 2).contiguous()
|
| 39 |
+
hidden_states = hidden_states + residual
|
| 40 |
+
|
| 41 |
+
return hidden_states, time_emb, text_emb, res_stack
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class SDVAEDecoder(torch.nn.Module):
|
| 45 |
+
def __init__(self):
|
| 46 |
+
super().__init__()
|
| 47 |
+
self.scaling_factor = 0.18215
|
| 48 |
+
self.post_quant_conv = torch.nn.Conv2d(4, 4, kernel_size=1)
|
| 49 |
+
self.conv_in = torch.nn.Conv2d(4, 512, kernel_size=3, padding=1)
|
| 50 |
+
|
| 51 |
+
self.blocks = torch.nn.ModuleList([
|
| 52 |
+
# UNetMidBlock2D
|
| 53 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 54 |
+
VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),
|
| 55 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 56 |
+
# UpDecoderBlock2D
|
| 57 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 58 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 59 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 60 |
+
UpSampler(512),
|
| 61 |
+
# UpDecoderBlock2D
|
| 62 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 63 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 64 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 65 |
+
UpSampler(512),
|
| 66 |
+
# UpDecoderBlock2D
|
| 67 |
+
ResnetBlock(512, 256, eps=1e-6),
|
| 68 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 69 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 70 |
+
UpSampler(256),
|
| 71 |
+
# UpDecoderBlock2D
|
| 72 |
+
ResnetBlock(256, 128, eps=1e-6),
|
| 73 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 74 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 75 |
+
])
|
| 76 |
+
|
| 77 |
+
self.conv_norm_out = torch.nn.GroupNorm(num_channels=128, num_groups=32, eps=1e-5)
|
| 78 |
+
self.conv_act = torch.nn.SiLU()
|
| 79 |
+
self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1)
|
| 80 |
+
|
| 81 |
+
def tiled_forward(self, sample, tile_size=64, tile_stride=32):
|
| 82 |
+
hidden_states = TileWorker().tiled_forward(
|
| 83 |
+
lambda x: self.forward(x),
|
| 84 |
+
sample,
|
| 85 |
+
tile_size,
|
| 86 |
+
tile_stride,
|
| 87 |
+
tile_device=sample.device,
|
| 88 |
+
tile_dtype=sample.dtype
|
| 89 |
+
)
|
| 90 |
+
return hidden_states
|
| 91 |
+
|
| 92 |
+
def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
|
| 93 |
+
original_dtype = sample.dtype
|
| 94 |
+
sample = sample.to(dtype=next(iter(self.parameters())).dtype)
|
| 95 |
+
# For VAE Decoder, we do not need to apply the tiler on each layer.
|
| 96 |
+
if tiled:
|
| 97 |
+
return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
|
| 98 |
+
|
| 99 |
+
# 1. pre-process
|
| 100 |
+
sample = sample / self.scaling_factor
|
| 101 |
+
hidden_states = self.post_quant_conv(sample)
|
| 102 |
+
hidden_states = self.conv_in(hidden_states)
|
| 103 |
+
time_emb = None
|
| 104 |
+
text_emb = None
|
| 105 |
+
res_stack = None
|
| 106 |
+
|
| 107 |
+
# 2. blocks
|
| 108 |
+
for i, block in enumerate(self.blocks):
|
| 109 |
+
hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
|
| 110 |
+
|
| 111 |
+
# 3. output
|
| 112 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 113 |
+
hidden_states = self.conv_act(hidden_states)
|
| 114 |
+
hidden_states = self.conv_out(hidden_states)
|
| 115 |
+
hidden_states = hidden_states.to(original_dtype)
|
| 116 |
+
|
| 117 |
+
return hidden_states
|
| 118 |
+
|
| 119 |
+
@staticmethod
|
| 120 |
+
def state_dict_converter():
|
| 121 |
+
return SDVAEDecoderStateDictConverter()
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
class SDVAEDecoderStateDictConverter:
|
| 125 |
+
def __init__(self):
|
| 126 |
+
pass
|
| 127 |
+
|
| 128 |
+
def from_diffusers(self, state_dict):
|
| 129 |
+
# architecture
|
| 130 |
+
block_types = [
|
| 131 |
+
'ResnetBlock', 'VAEAttentionBlock', 'ResnetBlock',
|
| 132 |
+
'ResnetBlock', 'ResnetBlock', 'ResnetBlock', 'UpSampler',
|
| 133 |
+
'ResnetBlock', 'ResnetBlock', 'ResnetBlock', 'UpSampler',
|
| 134 |
+
'ResnetBlock', 'ResnetBlock', 'ResnetBlock', 'UpSampler',
|
| 135 |
+
'ResnetBlock', 'ResnetBlock', 'ResnetBlock'
|
| 136 |
+
]
|
| 137 |
+
|
| 138 |
+
# Rename each parameter
|
| 139 |
+
local_rename_dict = {
|
| 140 |
+
"post_quant_conv": "post_quant_conv",
|
| 141 |
+
"decoder.conv_in": "conv_in",
|
| 142 |
+
"decoder.mid_block.attentions.0.group_norm": "blocks.1.norm",
|
| 143 |
+
"decoder.mid_block.attentions.0.to_q": "blocks.1.transformer_blocks.0.to_q",
|
| 144 |
+
"decoder.mid_block.attentions.0.to_k": "blocks.1.transformer_blocks.0.to_k",
|
| 145 |
+
"decoder.mid_block.attentions.0.to_v": "blocks.1.transformer_blocks.0.to_v",
|
| 146 |
+
"decoder.mid_block.attentions.0.to_out.0": "blocks.1.transformer_blocks.0.to_out",
|
| 147 |
+
"decoder.mid_block.resnets.0.norm1": "blocks.0.norm1",
|
| 148 |
+
"decoder.mid_block.resnets.0.conv1": "blocks.0.conv1",
|
| 149 |
+
"decoder.mid_block.resnets.0.norm2": "blocks.0.norm2",
|
| 150 |
+
"decoder.mid_block.resnets.0.conv2": "blocks.0.conv2",
|
| 151 |
+
"decoder.mid_block.resnets.1.norm1": "blocks.2.norm1",
|
| 152 |
+
"decoder.mid_block.resnets.1.conv1": "blocks.2.conv1",
|
| 153 |
+
"decoder.mid_block.resnets.1.norm2": "blocks.2.norm2",
|
| 154 |
+
"decoder.mid_block.resnets.1.conv2": "blocks.2.conv2",
|
| 155 |
+
"decoder.conv_norm_out": "conv_norm_out",
|
| 156 |
+
"decoder.conv_out": "conv_out",
|
| 157 |
+
}
|
| 158 |
+
name_list = sorted([name for name in state_dict])
|
| 159 |
+
rename_dict = {}
|
| 160 |
+
block_id = {"ResnetBlock": 2, "DownSampler": 2, "UpSampler": 2}
|
| 161 |
+
last_block_type_with_id = {"ResnetBlock": "", "DownSampler": "", "UpSampler": ""}
|
| 162 |
+
for name in name_list:
|
| 163 |
+
names = name.split(".")
|
| 164 |
+
name_prefix = ".".join(names[:-1])
|
| 165 |
+
if name_prefix in local_rename_dict:
|
| 166 |
+
rename_dict[name] = local_rename_dict[name_prefix] + "." + names[-1]
|
| 167 |
+
elif name.startswith("decoder.up_blocks"):
|
| 168 |
+
block_type = {"resnets": "ResnetBlock", "downsamplers": "DownSampler", "upsamplers": "UpSampler"}[names[3]]
|
| 169 |
+
block_type_with_id = ".".join(names[:5])
|
| 170 |
+
if block_type_with_id != last_block_type_with_id[block_type]:
|
| 171 |
+
block_id[block_type] += 1
|
| 172 |
+
last_block_type_with_id[block_type] = block_type_with_id
|
| 173 |
+
while block_id[block_type] < len(block_types) and block_types[block_id[block_type]] != block_type:
|
| 174 |
+
block_id[block_type] += 1
|
| 175 |
+
block_type_with_id = ".".join(names[:5])
|
| 176 |
+
names = ["blocks", str(block_id[block_type])] + names[5:]
|
| 177 |
+
rename_dict[name] = ".".join(names)
|
| 178 |
+
|
| 179 |
+
# Convert state_dict
|
| 180 |
+
state_dict_ = {}
|
| 181 |
+
for name, param in state_dict.items():
|
| 182 |
+
if name in rename_dict:
|
| 183 |
+
state_dict_[rename_dict[name]] = param
|
| 184 |
+
return state_dict_
|
| 185 |
+
|
| 186 |
+
def from_civitai(self, state_dict):
|
| 187 |
+
rename_dict = {
|
| 188 |
+
"first_stage_model.decoder.conv_in.bias": "conv_in.bias",
|
| 189 |
+
"first_stage_model.decoder.conv_in.weight": "conv_in.weight",
|
| 190 |
+
"first_stage_model.decoder.conv_out.bias": "conv_out.bias",
|
| 191 |
+
"first_stage_model.decoder.conv_out.weight": "conv_out.weight",
|
| 192 |
+
"first_stage_model.decoder.mid.attn_1.k.bias": "blocks.1.transformer_blocks.0.to_k.bias",
|
| 193 |
+
"first_stage_model.decoder.mid.attn_1.k.weight": "blocks.1.transformer_blocks.0.to_k.weight",
|
| 194 |
+
"first_stage_model.decoder.mid.attn_1.norm.bias": "blocks.1.norm.bias",
|
| 195 |
+
"first_stage_model.decoder.mid.attn_1.norm.weight": "blocks.1.norm.weight",
|
| 196 |
+
"first_stage_model.decoder.mid.attn_1.proj_out.bias": "blocks.1.transformer_blocks.0.to_out.bias",
|
| 197 |
+
"first_stage_model.decoder.mid.attn_1.proj_out.weight": "blocks.1.transformer_blocks.0.to_out.weight",
|
| 198 |
+
"first_stage_model.decoder.mid.attn_1.q.bias": "blocks.1.transformer_blocks.0.to_q.bias",
|
| 199 |
+
"first_stage_model.decoder.mid.attn_1.q.weight": "blocks.1.transformer_blocks.0.to_q.weight",
|
| 200 |
+
"first_stage_model.decoder.mid.attn_1.v.bias": "blocks.1.transformer_blocks.0.to_v.bias",
|
| 201 |
+
"first_stage_model.decoder.mid.attn_1.v.weight": "blocks.1.transformer_blocks.0.to_v.weight",
|
| 202 |
+
"first_stage_model.decoder.mid.block_1.conv1.bias": "blocks.0.conv1.bias",
|
| 203 |
+
"first_stage_model.decoder.mid.block_1.conv1.weight": "blocks.0.conv1.weight",
|
| 204 |
+
"first_stage_model.decoder.mid.block_1.conv2.bias": "blocks.0.conv2.bias",
|
| 205 |
+
"first_stage_model.decoder.mid.block_1.conv2.weight": "blocks.0.conv2.weight",
|
| 206 |
+
"first_stage_model.decoder.mid.block_1.norm1.bias": "blocks.0.norm1.bias",
|
| 207 |
+
"first_stage_model.decoder.mid.block_1.norm1.weight": "blocks.0.norm1.weight",
|
| 208 |
+
"first_stage_model.decoder.mid.block_1.norm2.bias": "blocks.0.norm2.bias",
|
| 209 |
+
"first_stage_model.decoder.mid.block_1.norm2.weight": "blocks.0.norm2.weight",
|
| 210 |
+
"first_stage_model.decoder.mid.block_2.conv1.bias": "blocks.2.conv1.bias",
|
| 211 |
+
"first_stage_model.decoder.mid.block_2.conv1.weight": "blocks.2.conv1.weight",
|
| 212 |
+
"first_stage_model.decoder.mid.block_2.conv2.bias": "blocks.2.conv2.bias",
|
| 213 |
+
"first_stage_model.decoder.mid.block_2.conv2.weight": "blocks.2.conv2.weight",
|
| 214 |
+
"first_stage_model.decoder.mid.block_2.norm1.bias": "blocks.2.norm1.bias",
|
| 215 |
+
"first_stage_model.decoder.mid.block_2.norm1.weight": "blocks.2.norm1.weight",
|
| 216 |
+
"first_stage_model.decoder.mid.block_2.norm2.bias": "blocks.2.norm2.bias",
|
| 217 |
+
"first_stage_model.decoder.mid.block_2.norm2.weight": "blocks.2.norm2.weight",
|
| 218 |
+
"first_stage_model.decoder.norm_out.bias": "conv_norm_out.bias",
|
| 219 |
+
"first_stage_model.decoder.norm_out.weight": "conv_norm_out.weight",
|
| 220 |
+
"first_stage_model.decoder.up.0.block.0.conv1.bias": "blocks.15.conv1.bias",
|
| 221 |
+
"first_stage_model.decoder.up.0.block.0.conv1.weight": "blocks.15.conv1.weight",
|
| 222 |
+
"first_stage_model.decoder.up.0.block.0.conv2.bias": "blocks.15.conv2.bias",
|
| 223 |
+
"first_stage_model.decoder.up.0.block.0.conv2.weight": "blocks.15.conv2.weight",
|
| 224 |
+
"first_stage_model.decoder.up.0.block.0.nin_shortcut.bias": "blocks.15.conv_shortcut.bias",
|
| 225 |
+
"first_stage_model.decoder.up.0.block.0.nin_shortcut.weight": "blocks.15.conv_shortcut.weight",
|
| 226 |
+
"first_stage_model.decoder.up.0.block.0.norm1.bias": "blocks.15.norm1.bias",
|
| 227 |
+
"first_stage_model.decoder.up.0.block.0.norm1.weight": "blocks.15.norm1.weight",
|
| 228 |
+
"first_stage_model.decoder.up.0.block.0.norm2.bias": "blocks.15.norm2.bias",
|
| 229 |
+
"first_stage_model.decoder.up.0.block.0.norm2.weight": "blocks.15.norm2.weight",
|
| 230 |
+
"first_stage_model.decoder.up.0.block.1.conv1.bias": "blocks.16.conv1.bias",
|
| 231 |
+
"first_stage_model.decoder.up.0.block.1.conv1.weight": "blocks.16.conv1.weight",
|
| 232 |
+
"first_stage_model.decoder.up.0.block.1.conv2.bias": "blocks.16.conv2.bias",
|
| 233 |
+
"first_stage_model.decoder.up.0.block.1.conv2.weight": "blocks.16.conv2.weight",
|
| 234 |
+
"first_stage_model.decoder.up.0.block.1.norm1.bias": "blocks.16.norm1.bias",
|
| 235 |
+
"first_stage_model.decoder.up.0.block.1.norm1.weight": "blocks.16.norm1.weight",
|
| 236 |
+
"first_stage_model.decoder.up.0.block.1.norm2.bias": "blocks.16.norm2.bias",
|
| 237 |
+
"first_stage_model.decoder.up.0.block.1.norm2.weight": "blocks.16.norm2.weight",
|
| 238 |
+
"first_stage_model.decoder.up.0.block.2.conv1.bias": "blocks.17.conv1.bias",
|
| 239 |
+
"first_stage_model.decoder.up.0.block.2.conv1.weight": "blocks.17.conv1.weight",
|
| 240 |
+
"first_stage_model.decoder.up.0.block.2.conv2.bias": "blocks.17.conv2.bias",
|
| 241 |
+
"first_stage_model.decoder.up.0.block.2.conv2.weight": "blocks.17.conv2.weight",
|
| 242 |
+
"first_stage_model.decoder.up.0.block.2.norm1.bias": "blocks.17.norm1.bias",
|
| 243 |
+
"first_stage_model.decoder.up.0.block.2.norm1.weight": "blocks.17.norm1.weight",
|
| 244 |
+
"first_stage_model.decoder.up.0.block.2.norm2.bias": "blocks.17.norm2.bias",
|
| 245 |
+
"first_stage_model.decoder.up.0.block.2.norm2.weight": "blocks.17.norm2.weight",
|
| 246 |
+
"first_stage_model.decoder.up.1.block.0.conv1.bias": "blocks.11.conv1.bias",
|
| 247 |
+
"first_stage_model.decoder.up.1.block.0.conv1.weight": "blocks.11.conv1.weight",
|
| 248 |
+
"first_stage_model.decoder.up.1.block.0.conv2.bias": "blocks.11.conv2.bias",
|
| 249 |
+
"first_stage_model.decoder.up.1.block.0.conv2.weight": "blocks.11.conv2.weight",
|
| 250 |
+
"first_stage_model.decoder.up.1.block.0.nin_shortcut.bias": "blocks.11.conv_shortcut.bias",
|
| 251 |
+
"first_stage_model.decoder.up.1.block.0.nin_shortcut.weight": "blocks.11.conv_shortcut.weight",
|
| 252 |
+
"first_stage_model.decoder.up.1.block.0.norm1.bias": "blocks.11.norm1.bias",
|
| 253 |
+
"first_stage_model.decoder.up.1.block.0.norm1.weight": "blocks.11.norm1.weight",
|
| 254 |
+
"first_stage_model.decoder.up.1.block.0.norm2.bias": "blocks.11.norm2.bias",
|
| 255 |
+
"first_stage_model.decoder.up.1.block.0.norm2.weight": "blocks.11.norm2.weight",
|
| 256 |
+
"first_stage_model.decoder.up.1.block.1.conv1.bias": "blocks.12.conv1.bias",
|
| 257 |
+
"first_stage_model.decoder.up.1.block.1.conv1.weight": "blocks.12.conv1.weight",
|
| 258 |
+
"first_stage_model.decoder.up.1.block.1.conv2.bias": "blocks.12.conv2.bias",
|
| 259 |
+
"first_stage_model.decoder.up.1.block.1.conv2.weight": "blocks.12.conv2.weight",
|
| 260 |
+
"first_stage_model.decoder.up.1.block.1.norm1.bias": "blocks.12.norm1.bias",
|
| 261 |
+
"first_stage_model.decoder.up.1.block.1.norm1.weight": "blocks.12.norm1.weight",
|
| 262 |
+
"first_stage_model.decoder.up.1.block.1.norm2.bias": "blocks.12.norm2.bias",
|
| 263 |
+
"first_stage_model.decoder.up.1.block.1.norm2.weight": "blocks.12.norm2.weight",
|
| 264 |
+
"first_stage_model.decoder.up.1.block.2.conv1.bias": "blocks.13.conv1.bias",
|
| 265 |
+
"first_stage_model.decoder.up.1.block.2.conv1.weight": "blocks.13.conv1.weight",
|
| 266 |
+
"first_stage_model.decoder.up.1.block.2.conv2.bias": "blocks.13.conv2.bias",
|
| 267 |
+
"first_stage_model.decoder.up.1.block.2.conv2.weight": "blocks.13.conv2.weight",
|
| 268 |
+
"first_stage_model.decoder.up.1.block.2.norm1.bias": "blocks.13.norm1.bias",
|
| 269 |
+
"first_stage_model.decoder.up.1.block.2.norm1.weight": "blocks.13.norm1.weight",
|
| 270 |
+
"first_stage_model.decoder.up.1.block.2.norm2.bias": "blocks.13.norm2.bias",
|
| 271 |
+
"first_stage_model.decoder.up.1.block.2.norm2.weight": "blocks.13.norm2.weight",
|
| 272 |
+
"first_stage_model.decoder.up.1.upsample.conv.bias": "blocks.14.conv.bias",
|
| 273 |
+
"first_stage_model.decoder.up.1.upsample.conv.weight": "blocks.14.conv.weight",
|
| 274 |
+
"first_stage_model.decoder.up.2.block.0.conv1.bias": "blocks.7.conv1.bias",
|
| 275 |
+
"first_stage_model.decoder.up.2.block.0.conv1.weight": "blocks.7.conv1.weight",
|
| 276 |
+
"first_stage_model.decoder.up.2.block.0.conv2.bias": "blocks.7.conv2.bias",
|
| 277 |
+
"first_stage_model.decoder.up.2.block.0.conv2.weight": "blocks.7.conv2.weight",
|
| 278 |
+
"first_stage_model.decoder.up.2.block.0.norm1.bias": "blocks.7.norm1.bias",
|
| 279 |
+
"first_stage_model.decoder.up.2.block.0.norm1.weight": "blocks.7.norm1.weight",
|
| 280 |
+
"first_stage_model.decoder.up.2.block.0.norm2.bias": "blocks.7.norm2.bias",
|
| 281 |
+
"first_stage_model.decoder.up.2.block.0.norm2.weight": "blocks.7.norm2.weight",
|
| 282 |
+
"first_stage_model.decoder.up.2.block.1.conv1.bias": "blocks.8.conv1.bias",
|
| 283 |
+
"first_stage_model.decoder.up.2.block.1.conv1.weight": "blocks.8.conv1.weight",
|
| 284 |
+
"first_stage_model.decoder.up.2.block.1.conv2.bias": "blocks.8.conv2.bias",
|
| 285 |
+
"first_stage_model.decoder.up.2.block.1.conv2.weight": "blocks.8.conv2.weight",
|
| 286 |
+
"first_stage_model.decoder.up.2.block.1.norm1.bias": "blocks.8.norm1.bias",
|
| 287 |
+
"first_stage_model.decoder.up.2.block.1.norm1.weight": "blocks.8.norm1.weight",
|
| 288 |
+
"first_stage_model.decoder.up.2.block.1.norm2.bias": "blocks.8.norm2.bias",
|
| 289 |
+
"first_stage_model.decoder.up.2.block.1.norm2.weight": "blocks.8.norm2.weight",
|
| 290 |
+
"first_stage_model.decoder.up.2.block.2.conv1.bias": "blocks.9.conv1.bias",
|
| 291 |
+
"first_stage_model.decoder.up.2.block.2.conv1.weight": "blocks.9.conv1.weight",
|
| 292 |
+
"first_stage_model.decoder.up.2.block.2.conv2.bias": "blocks.9.conv2.bias",
|
| 293 |
+
"first_stage_model.decoder.up.2.block.2.conv2.weight": "blocks.9.conv2.weight",
|
| 294 |
+
"first_stage_model.decoder.up.2.block.2.norm1.bias": "blocks.9.norm1.bias",
|
| 295 |
+
"first_stage_model.decoder.up.2.block.2.norm1.weight": "blocks.9.norm1.weight",
|
| 296 |
+
"first_stage_model.decoder.up.2.block.2.norm2.bias": "blocks.9.norm2.bias",
|
| 297 |
+
"first_stage_model.decoder.up.2.block.2.norm2.weight": "blocks.9.norm2.weight",
|
| 298 |
+
"first_stage_model.decoder.up.2.upsample.conv.bias": "blocks.10.conv.bias",
|
| 299 |
+
"first_stage_model.decoder.up.2.upsample.conv.weight": "blocks.10.conv.weight",
|
| 300 |
+
"first_stage_model.decoder.up.3.block.0.conv1.bias": "blocks.3.conv1.bias",
|
| 301 |
+
"first_stage_model.decoder.up.3.block.0.conv1.weight": "blocks.3.conv1.weight",
|
| 302 |
+
"first_stage_model.decoder.up.3.block.0.conv2.bias": "blocks.3.conv2.bias",
|
| 303 |
+
"first_stage_model.decoder.up.3.block.0.conv2.weight": "blocks.3.conv2.weight",
|
| 304 |
+
"first_stage_model.decoder.up.3.block.0.norm1.bias": "blocks.3.norm1.bias",
|
| 305 |
+
"first_stage_model.decoder.up.3.block.0.norm1.weight": "blocks.3.norm1.weight",
|
| 306 |
+
"first_stage_model.decoder.up.3.block.0.norm2.bias": "blocks.3.norm2.bias",
|
| 307 |
+
"first_stage_model.decoder.up.3.block.0.norm2.weight": "blocks.3.norm2.weight",
|
| 308 |
+
"first_stage_model.decoder.up.3.block.1.conv1.bias": "blocks.4.conv1.bias",
|
| 309 |
+
"first_stage_model.decoder.up.3.block.1.conv1.weight": "blocks.4.conv1.weight",
|
| 310 |
+
"first_stage_model.decoder.up.3.block.1.conv2.bias": "blocks.4.conv2.bias",
|
| 311 |
+
"first_stage_model.decoder.up.3.block.1.conv2.weight": "blocks.4.conv2.weight",
|
| 312 |
+
"first_stage_model.decoder.up.3.block.1.norm1.bias": "blocks.4.norm1.bias",
|
| 313 |
+
"first_stage_model.decoder.up.3.block.1.norm1.weight": "blocks.4.norm1.weight",
|
| 314 |
+
"first_stage_model.decoder.up.3.block.1.norm2.bias": "blocks.4.norm2.bias",
|
| 315 |
+
"first_stage_model.decoder.up.3.block.1.norm2.weight": "blocks.4.norm2.weight",
|
| 316 |
+
"first_stage_model.decoder.up.3.block.2.conv1.bias": "blocks.5.conv1.bias",
|
| 317 |
+
"first_stage_model.decoder.up.3.block.2.conv1.weight": "blocks.5.conv1.weight",
|
| 318 |
+
"first_stage_model.decoder.up.3.block.2.conv2.bias": "blocks.5.conv2.bias",
|
| 319 |
+
"first_stage_model.decoder.up.3.block.2.conv2.weight": "blocks.5.conv2.weight",
|
| 320 |
+
"first_stage_model.decoder.up.3.block.2.norm1.bias": "blocks.5.norm1.bias",
|
| 321 |
+
"first_stage_model.decoder.up.3.block.2.norm1.weight": "blocks.5.norm1.weight",
|
| 322 |
+
"first_stage_model.decoder.up.3.block.2.norm2.bias": "blocks.5.norm2.bias",
|
| 323 |
+
"first_stage_model.decoder.up.3.block.2.norm2.weight": "blocks.5.norm2.weight",
|
| 324 |
+
"first_stage_model.decoder.up.3.upsample.conv.bias": "blocks.6.conv.bias",
|
| 325 |
+
"first_stage_model.decoder.up.3.upsample.conv.weight": "blocks.6.conv.weight",
|
| 326 |
+
"first_stage_model.post_quant_conv.bias": "post_quant_conv.bias",
|
| 327 |
+
"first_stage_model.post_quant_conv.weight": "post_quant_conv.weight",
|
| 328 |
+
}
|
| 329 |
+
state_dict_ = {}
|
| 330 |
+
for name in state_dict:
|
| 331 |
+
if name in rename_dict:
|
| 332 |
+
param = state_dict[name]
|
| 333 |
+
if "transformer_blocks" in rename_dict[name]:
|
| 334 |
+
param = param.squeeze()
|
| 335 |
+
state_dict_[rename_dict[name]] = param
|
| 336 |
+
return state_dict_
|
diffsynth/models/sd_vae_encoder.py
ADDED
|
@@ -0,0 +1,282 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_unet import ResnetBlock, DownSampler
|
| 3 |
+
from .sd_vae_decoder import VAEAttentionBlock
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
from einops import rearrange
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SDVAEEncoder(torch.nn.Module):
|
| 9 |
+
def __init__(self):
|
| 10 |
+
super().__init__()
|
| 11 |
+
self.scaling_factor = 0.18215
|
| 12 |
+
self.quant_conv = torch.nn.Conv2d(8, 8, kernel_size=1)
|
| 13 |
+
self.conv_in = torch.nn.Conv2d(3, 128, kernel_size=3, padding=1)
|
| 14 |
+
|
| 15 |
+
self.blocks = torch.nn.ModuleList([
|
| 16 |
+
# DownEncoderBlock2D
|
| 17 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 18 |
+
ResnetBlock(128, 128, eps=1e-6),
|
| 19 |
+
DownSampler(128, padding=0, extra_padding=True),
|
| 20 |
+
# DownEncoderBlock2D
|
| 21 |
+
ResnetBlock(128, 256, eps=1e-6),
|
| 22 |
+
ResnetBlock(256, 256, eps=1e-6),
|
| 23 |
+
DownSampler(256, padding=0, extra_padding=True),
|
| 24 |
+
# DownEncoderBlock2D
|
| 25 |
+
ResnetBlock(256, 512, eps=1e-6),
|
| 26 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 27 |
+
DownSampler(512, padding=0, extra_padding=True),
|
| 28 |
+
# DownEncoderBlock2D
|
| 29 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 30 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 31 |
+
# UNetMidBlock2D
|
| 32 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 33 |
+
VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),
|
| 34 |
+
ResnetBlock(512, 512, eps=1e-6),
|
| 35 |
+
])
|
| 36 |
+
|
| 37 |
+
self.conv_norm_out = torch.nn.GroupNorm(num_channels=512, num_groups=32, eps=1e-6)
|
| 38 |
+
self.conv_act = torch.nn.SiLU()
|
| 39 |
+
self.conv_out = torch.nn.Conv2d(512, 8, kernel_size=3, padding=1)
|
| 40 |
+
|
| 41 |
+
def tiled_forward(self, sample, tile_size=64, tile_stride=32):
|
| 42 |
+
hidden_states = TileWorker().tiled_forward(
|
| 43 |
+
lambda x: self.forward(x),
|
| 44 |
+
sample,
|
| 45 |
+
tile_size,
|
| 46 |
+
tile_stride,
|
| 47 |
+
tile_device=sample.device,
|
| 48 |
+
tile_dtype=sample.dtype
|
| 49 |
+
)
|
| 50 |
+
return hidden_states
|
| 51 |
+
|
| 52 |
+
def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):
|
| 53 |
+
original_dtype = sample.dtype
|
| 54 |
+
sample = sample.to(dtype=next(iter(self.parameters())).dtype)
|
| 55 |
+
# For VAE Decoder, we do not need to apply the tiler on each layer.
|
| 56 |
+
if tiled:
|
| 57 |
+
return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)
|
| 58 |
+
|
| 59 |
+
# 1. pre-process
|
| 60 |
+
hidden_states = self.conv_in(sample)
|
| 61 |
+
time_emb = None
|
| 62 |
+
text_emb = None
|
| 63 |
+
res_stack = None
|
| 64 |
+
|
| 65 |
+
# 2. blocks
|
| 66 |
+
for i, block in enumerate(self.blocks):
|
| 67 |
+
hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)
|
| 68 |
+
|
| 69 |
+
# 3. output
|
| 70 |
+
hidden_states = self.conv_norm_out(hidden_states)
|
| 71 |
+
hidden_states = self.conv_act(hidden_states)
|
| 72 |
+
hidden_states = self.conv_out(hidden_states)
|
| 73 |
+
hidden_states = self.quant_conv(hidden_states)
|
| 74 |
+
hidden_states = hidden_states[:, :4]
|
| 75 |
+
hidden_states *= self.scaling_factor
|
| 76 |
+
hidden_states = hidden_states.to(original_dtype)
|
| 77 |
+
|
| 78 |
+
return hidden_states
|
| 79 |
+
|
| 80 |
+
def encode_video(self, sample, batch_size=8):
|
| 81 |
+
B = sample.shape[0]
|
| 82 |
+
hidden_states = []
|
| 83 |
+
|
| 84 |
+
for i in range(0, sample.shape[2], batch_size):
|
| 85 |
+
|
| 86 |
+
j = min(i + batch_size, sample.shape[2])
|
| 87 |
+
sample_batch = rearrange(sample[:,:,i:j], "B C T H W -> (B T) C H W")
|
| 88 |
+
|
| 89 |
+
hidden_states_batch = self(sample_batch)
|
| 90 |
+
hidden_states_batch = rearrange(hidden_states_batch, "(B T) C H W -> B C T H W", B=B)
|
| 91 |
+
|
| 92 |
+
hidden_states.append(hidden_states_batch)
|
| 93 |
+
|
| 94 |
+
hidden_states = torch.concat(hidden_states, dim=2)
|
| 95 |
+
return hidden_states
|
| 96 |
+
|
| 97 |
+
@staticmethod
|
| 98 |
+
def state_dict_converter():
|
| 99 |
+
return SDVAEEncoderStateDictConverter()
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
class SDVAEEncoderStateDictConverter:
|
| 103 |
+
def __init__(self):
|
| 104 |
+
pass
|
| 105 |
+
|
| 106 |
+
def from_diffusers(self, state_dict):
|
| 107 |
+
# architecture
|
| 108 |
+
block_types = [
|
| 109 |
+
'ResnetBlock', 'ResnetBlock', 'DownSampler',
|
| 110 |
+
'ResnetBlock', 'ResnetBlock', 'DownSampler',
|
| 111 |
+
'ResnetBlock', 'ResnetBlock', 'DownSampler',
|
| 112 |
+
'ResnetBlock', 'ResnetBlock',
|
| 113 |
+
'ResnetBlock', 'VAEAttentionBlock', 'ResnetBlock'
|
| 114 |
+
]
|
| 115 |
+
|
| 116 |
+
# Rename each parameter
|
| 117 |
+
local_rename_dict = {
|
| 118 |
+
"quant_conv": "quant_conv",
|
| 119 |
+
"encoder.conv_in": "conv_in",
|
| 120 |
+
"encoder.mid_block.attentions.0.group_norm": "blocks.12.norm",
|
| 121 |
+
"encoder.mid_block.attentions.0.to_q": "blocks.12.transformer_blocks.0.to_q",
|
| 122 |
+
"encoder.mid_block.attentions.0.to_k": "blocks.12.transformer_blocks.0.to_k",
|
| 123 |
+
"encoder.mid_block.attentions.0.to_v": "blocks.12.transformer_blocks.0.to_v",
|
| 124 |
+
"encoder.mid_block.attentions.0.to_out.0": "blocks.12.transformer_blocks.0.to_out",
|
| 125 |
+
"encoder.mid_block.resnets.0.norm1": "blocks.11.norm1",
|
| 126 |
+
"encoder.mid_block.resnets.0.conv1": "blocks.11.conv1",
|
| 127 |
+
"encoder.mid_block.resnets.0.norm2": "blocks.11.norm2",
|
| 128 |
+
"encoder.mid_block.resnets.0.conv2": "blocks.11.conv2",
|
| 129 |
+
"encoder.mid_block.resnets.1.norm1": "blocks.13.norm1",
|
| 130 |
+
"encoder.mid_block.resnets.1.conv1": "blocks.13.conv1",
|
| 131 |
+
"encoder.mid_block.resnets.1.norm2": "blocks.13.norm2",
|
| 132 |
+
"encoder.mid_block.resnets.1.conv2": "blocks.13.conv2",
|
| 133 |
+
"encoder.conv_norm_out": "conv_norm_out",
|
| 134 |
+
"encoder.conv_out": "conv_out",
|
| 135 |
+
}
|
| 136 |
+
name_list = sorted([name for name in state_dict])
|
| 137 |
+
rename_dict = {}
|
| 138 |
+
block_id = {"ResnetBlock": -1, "DownSampler": -1, "UpSampler": -1}
|
| 139 |
+
last_block_type_with_id = {"ResnetBlock": "", "DownSampler": "", "UpSampler": ""}
|
| 140 |
+
for name in name_list:
|
| 141 |
+
names = name.split(".")
|
| 142 |
+
name_prefix = ".".join(names[:-1])
|
| 143 |
+
if name_prefix in local_rename_dict:
|
| 144 |
+
rename_dict[name] = local_rename_dict[name_prefix] + "." + names[-1]
|
| 145 |
+
elif name.startswith("encoder.down_blocks"):
|
| 146 |
+
block_type = {"resnets": "ResnetBlock", "downsamplers": "DownSampler", "upsamplers": "UpSampler"}[names[3]]
|
| 147 |
+
block_type_with_id = ".".join(names[:5])
|
| 148 |
+
if block_type_with_id != last_block_type_with_id[block_type]:
|
| 149 |
+
block_id[block_type] += 1
|
| 150 |
+
last_block_type_with_id[block_type] = block_type_with_id
|
| 151 |
+
while block_id[block_type] < len(block_types) and block_types[block_id[block_type]] != block_type:
|
| 152 |
+
block_id[block_type] += 1
|
| 153 |
+
block_type_with_id = ".".join(names[:5])
|
| 154 |
+
names = ["blocks", str(block_id[block_type])] + names[5:]
|
| 155 |
+
rename_dict[name] = ".".join(names)
|
| 156 |
+
|
| 157 |
+
# Convert state_dict
|
| 158 |
+
state_dict_ = {}
|
| 159 |
+
for name, param in state_dict.items():
|
| 160 |
+
if name in rename_dict:
|
| 161 |
+
state_dict_[rename_dict[name]] = param
|
| 162 |
+
return state_dict_
|
| 163 |
+
|
| 164 |
+
def from_civitai(self, state_dict):
|
| 165 |
+
rename_dict = {
|
| 166 |
+
"first_stage_model.encoder.conv_in.bias": "conv_in.bias",
|
| 167 |
+
"first_stage_model.encoder.conv_in.weight": "conv_in.weight",
|
| 168 |
+
"first_stage_model.encoder.conv_out.bias": "conv_out.bias",
|
| 169 |
+
"first_stage_model.encoder.conv_out.weight": "conv_out.weight",
|
| 170 |
+
"first_stage_model.encoder.down.0.block.0.conv1.bias": "blocks.0.conv1.bias",
|
| 171 |
+
"first_stage_model.encoder.down.0.block.0.conv1.weight": "blocks.0.conv1.weight",
|
| 172 |
+
"first_stage_model.encoder.down.0.block.0.conv2.bias": "blocks.0.conv2.bias",
|
| 173 |
+
"first_stage_model.encoder.down.0.block.0.conv2.weight": "blocks.0.conv2.weight",
|
| 174 |
+
"first_stage_model.encoder.down.0.block.0.norm1.bias": "blocks.0.norm1.bias",
|
| 175 |
+
"first_stage_model.encoder.down.0.block.0.norm1.weight": "blocks.0.norm1.weight",
|
| 176 |
+
"first_stage_model.encoder.down.0.block.0.norm2.bias": "blocks.0.norm2.bias",
|
| 177 |
+
"first_stage_model.encoder.down.0.block.0.norm2.weight": "blocks.0.norm2.weight",
|
| 178 |
+
"first_stage_model.encoder.down.0.block.1.conv1.bias": "blocks.1.conv1.bias",
|
| 179 |
+
"first_stage_model.encoder.down.0.block.1.conv1.weight": "blocks.1.conv1.weight",
|
| 180 |
+
"first_stage_model.encoder.down.0.block.1.conv2.bias": "blocks.1.conv2.bias",
|
| 181 |
+
"first_stage_model.encoder.down.0.block.1.conv2.weight": "blocks.1.conv2.weight",
|
| 182 |
+
"first_stage_model.encoder.down.0.block.1.norm1.bias": "blocks.1.norm1.bias",
|
| 183 |
+
"first_stage_model.encoder.down.0.block.1.norm1.weight": "blocks.1.norm1.weight",
|
| 184 |
+
"first_stage_model.encoder.down.0.block.1.norm2.bias": "blocks.1.norm2.bias",
|
| 185 |
+
"first_stage_model.encoder.down.0.block.1.norm2.weight": "blocks.1.norm2.weight",
|
| 186 |
+
"first_stage_model.encoder.down.0.downsample.conv.bias": "blocks.2.conv.bias",
|
| 187 |
+
"first_stage_model.encoder.down.0.downsample.conv.weight": "blocks.2.conv.weight",
|
| 188 |
+
"first_stage_model.encoder.down.1.block.0.conv1.bias": "blocks.3.conv1.bias",
|
| 189 |
+
"first_stage_model.encoder.down.1.block.0.conv1.weight": "blocks.3.conv1.weight",
|
| 190 |
+
"first_stage_model.encoder.down.1.block.0.conv2.bias": "blocks.3.conv2.bias",
|
| 191 |
+
"first_stage_model.encoder.down.1.block.0.conv2.weight": "blocks.3.conv2.weight",
|
| 192 |
+
"first_stage_model.encoder.down.1.block.0.nin_shortcut.bias": "blocks.3.conv_shortcut.bias",
|
| 193 |
+
"first_stage_model.encoder.down.1.block.0.nin_shortcut.weight": "blocks.3.conv_shortcut.weight",
|
| 194 |
+
"first_stage_model.encoder.down.1.block.0.norm1.bias": "blocks.3.norm1.bias",
|
| 195 |
+
"first_stage_model.encoder.down.1.block.0.norm1.weight": "blocks.3.norm1.weight",
|
| 196 |
+
"first_stage_model.encoder.down.1.block.0.norm2.bias": "blocks.3.norm2.bias",
|
| 197 |
+
"first_stage_model.encoder.down.1.block.0.norm2.weight": "blocks.3.norm2.weight",
|
| 198 |
+
"first_stage_model.encoder.down.1.block.1.conv1.bias": "blocks.4.conv1.bias",
|
| 199 |
+
"first_stage_model.encoder.down.1.block.1.conv1.weight": "blocks.4.conv1.weight",
|
| 200 |
+
"first_stage_model.encoder.down.1.block.1.conv2.bias": "blocks.4.conv2.bias",
|
| 201 |
+
"first_stage_model.encoder.down.1.block.1.conv2.weight": "blocks.4.conv2.weight",
|
| 202 |
+
"first_stage_model.encoder.down.1.block.1.norm1.bias": "blocks.4.norm1.bias",
|
| 203 |
+
"first_stage_model.encoder.down.1.block.1.norm1.weight": "blocks.4.norm1.weight",
|
| 204 |
+
"first_stage_model.encoder.down.1.block.1.norm2.bias": "blocks.4.norm2.bias",
|
| 205 |
+
"first_stage_model.encoder.down.1.block.1.norm2.weight": "blocks.4.norm2.weight",
|
| 206 |
+
"first_stage_model.encoder.down.1.downsample.conv.bias": "blocks.5.conv.bias",
|
| 207 |
+
"first_stage_model.encoder.down.1.downsample.conv.weight": "blocks.5.conv.weight",
|
| 208 |
+
"first_stage_model.encoder.down.2.block.0.conv1.bias": "blocks.6.conv1.bias",
|
| 209 |
+
"first_stage_model.encoder.down.2.block.0.conv1.weight": "blocks.6.conv1.weight",
|
| 210 |
+
"first_stage_model.encoder.down.2.block.0.conv2.bias": "blocks.6.conv2.bias",
|
| 211 |
+
"first_stage_model.encoder.down.2.block.0.conv2.weight": "blocks.6.conv2.weight",
|
| 212 |
+
"first_stage_model.encoder.down.2.block.0.nin_shortcut.bias": "blocks.6.conv_shortcut.bias",
|
| 213 |
+
"first_stage_model.encoder.down.2.block.0.nin_shortcut.weight": "blocks.6.conv_shortcut.weight",
|
| 214 |
+
"first_stage_model.encoder.down.2.block.0.norm1.bias": "blocks.6.norm1.bias",
|
| 215 |
+
"first_stage_model.encoder.down.2.block.0.norm1.weight": "blocks.6.norm1.weight",
|
| 216 |
+
"first_stage_model.encoder.down.2.block.0.norm2.bias": "blocks.6.norm2.bias",
|
| 217 |
+
"first_stage_model.encoder.down.2.block.0.norm2.weight": "blocks.6.norm2.weight",
|
| 218 |
+
"first_stage_model.encoder.down.2.block.1.conv1.bias": "blocks.7.conv1.bias",
|
| 219 |
+
"first_stage_model.encoder.down.2.block.1.conv1.weight": "blocks.7.conv1.weight",
|
| 220 |
+
"first_stage_model.encoder.down.2.block.1.conv2.bias": "blocks.7.conv2.bias",
|
| 221 |
+
"first_stage_model.encoder.down.2.block.1.conv2.weight": "blocks.7.conv2.weight",
|
| 222 |
+
"first_stage_model.encoder.down.2.block.1.norm1.bias": "blocks.7.norm1.bias",
|
| 223 |
+
"first_stage_model.encoder.down.2.block.1.norm1.weight": "blocks.7.norm1.weight",
|
| 224 |
+
"first_stage_model.encoder.down.2.block.1.norm2.bias": "blocks.7.norm2.bias",
|
| 225 |
+
"first_stage_model.encoder.down.2.block.1.norm2.weight": "blocks.7.norm2.weight",
|
| 226 |
+
"first_stage_model.encoder.down.2.downsample.conv.bias": "blocks.8.conv.bias",
|
| 227 |
+
"first_stage_model.encoder.down.2.downsample.conv.weight": "blocks.8.conv.weight",
|
| 228 |
+
"first_stage_model.encoder.down.3.block.0.conv1.bias": "blocks.9.conv1.bias",
|
| 229 |
+
"first_stage_model.encoder.down.3.block.0.conv1.weight": "blocks.9.conv1.weight",
|
| 230 |
+
"first_stage_model.encoder.down.3.block.0.conv2.bias": "blocks.9.conv2.bias",
|
| 231 |
+
"first_stage_model.encoder.down.3.block.0.conv2.weight": "blocks.9.conv2.weight",
|
| 232 |
+
"first_stage_model.encoder.down.3.block.0.norm1.bias": "blocks.9.norm1.bias",
|
| 233 |
+
"first_stage_model.encoder.down.3.block.0.norm1.weight": "blocks.9.norm1.weight",
|
| 234 |
+
"first_stage_model.encoder.down.3.block.0.norm2.bias": "blocks.9.norm2.bias",
|
| 235 |
+
"first_stage_model.encoder.down.3.block.0.norm2.weight": "blocks.9.norm2.weight",
|
| 236 |
+
"first_stage_model.encoder.down.3.block.1.conv1.bias": "blocks.10.conv1.bias",
|
| 237 |
+
"first_stage_model.encoder.down.3.block.1.conv1.weight": "blocks.10.conv1.weight",
|
| 238 |
+
"first_stage_model.encoder.down.3.block.1.conv2.bias": "blocks.10.conv2.bias",
|
| 239 |
+
"first_stage_model.encoder.down.3.block.1.conv2.weight": "blocks.10.conv2.weight",
|
| 240 |
+
"first_stage_model.encoder.down.3.block.1.norm1.bias": "blocks.10.norm1.bias",
|
| 241 |
+
"first_stage_model.encoder.down.3.block.1.norm1.weight": "blocks.10.norm1.weight",
|
| 242 |
+
"first_stage_model.encoder.down.3.block.1.norm2.bias": "blocks.10.norm2.bias",
|
| 243 |
+
"first_stage_model.encoder.down.3.block.1.norm2.weight": "blocks.10.norm2.weight",
|
| 244 |
+
"first_stage_model.encoder.mid.attn_1.k.bias": "blocks.12.transformer_blocks.0.to_k.bias",
|
| 245 |
+
"first_stage_model.encoder.mid.attn_1.k.weight": "blocks.12.transformer_blocks.0.to_k.weight",
|
| 246 |
+
"first_stage_model.encoder.mid.attn_1.norm.bias": "blocks.12.norm.bias",
|
| 247 |
+
"first_stage_model.encoder.mid.attn_1.norm.weight": "blocks.12.norm.weight",
|
| 248 |
+
"first_stage_model.encoder.mid.attn_1.proj_out.bias": "blocks.12.transformer_blocks.0.to_out.bias",
|
| 249 |
+
"first_stage_model.encoder.mid.attn_1.proj_out.weight": "blocks.12.transformer_blocks.0.to_out.weight",
|
| 250 |
+
"first_stage_model.encoder.mid.attn_1.q.bias": "blocks.12.transformer_blocks.0.to_q.bias",
|
| 251 |
+
"first_stage_model.encoder.mid.attn_1.q.weight": "blocks.12.transformer_blocks.0.to_q.weight",
|
| 252 |
+
"first_stage_model.encoder.mid.attn_1.v.bias": "blocks.12.transformer_blocks.0.to_v.bias",
|
| 253 |
+
"first_stage_model.encoder.mid.attn_1.v.weight": "blocks.12.transformer_blocks.0.to_v.weight",
|
| 254 |
+
"first_stage_model.encoder.mid.block_1.conv1.bias": "blocks.11.conv1.bias",
|
| 255 |
+
"first_stage_model.encoder.mid.block_1.conv1.weight": "blocks.11.conv1.weight",
|
| 256 |
+
"first_stage_model.encoder.mid.block_1.conv2.bias": "blocks.11.conv2.bias",
|
| 257 |
+
"first_stage_model.encoder.mid.block_1.conv2.weight": "blocks.11.conv2.weight",
|
| 258 |
+
"first_stage_model.encoder.mid.block_1.norm1.bias": "blocks.11.norm1.bias",
|
| 259 |
+
"first_stage_model.encoder.mid.block_1.norm1.weight": "blocks.11.norm1.weight",
|
| 260 |
+
"first_stage_model.encoder.mid.block_1.norm2.bias": "blocks.11.norm2.bias",
|
| 261 |
+
"first_stage_model.encoder.mid.block_1.norm2.weight": "blocks.11.norm2.weight",
|
| 262 |
+
"first_stage_model.encoder.mid.block_2.conv1.bias": "blocks.13.conv1.bias",
|
| 263 |
+
"first_stage_model.encoder.mid.block_2.conv1.weight": "blocks.13.conv1.weight",
|
| 264 |
+
"first_stage_model.encoder.mid.block_2.conv2.bias": "blocks.13.conv2.bias",
|
| 265 |
+
"first_stage_model.encoder.mid.block_2.conv2.weight": "blocks.13.conv2.weight",
|
| 266 |
+
"first_stage_model.encoder.mid.block_2.norm1.bias": "blocks.13.norm1.bias",
|
| 267 |
+
"first_stage_model.encoder.mid.block_2.norm1.weight": "blocks.13.norm1.weight",
|
| 268 |
+
"first_stage_model.encoder.mid.block_2.norm2.bias": "blocks.13.norm2.bias",
|
| 269 |
+
"first_stage_model.encoder.mid.block_2.norm2.weight": "blocks.13.norm2.weight",
|
| 270 |
+
"first_stage_model.encoder.norm_out.bias": "conv_norm_out.bias",
|
| 271 |
+
"first_stage_model.encoder.norm_out.weight": "conv_norm_out.weight",
|
| 272 |
+
"first_stage_model.quant_conv.bias": "quant_conv.bias",
|
| 273 |
+
"first_stage_model.quant_conv.weight": "quant_conv.weight",
|
| 274 |
+
}
|
| 275 |
+
state_dict_ = {}
|
| 276 |
+
for name in state_dict:
|
| 277 |
+
if name in rename_dict:
|
| 278 |
+
param = state_dict[name]
|
| 279 |
+
if "transformer_blocks" in rename_dict[name]:
|
| 280 |
+
param = param.squeeze()
|
| 281 |
+
state_dict_[rename_dict[name]] = param
|
| 282 |
+
return state_dict_
|
diffsynth/models/sdxl_controlnet.py
ADDED
|
@@ -0,0 +1,318 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_unet import Timesteps, ResnetBlock, AttentionBlock, PushBlock, DownSampler
|
| 3 |
+
from .sdxl_unet import SDXLUNet
|
| 4 |
+
from .tiler import TileWorker
|
| 5 |
+
from .sd_controlnet import ControlNetConditioningLayer
|
| 6 |
+
from collections import OrderedDict
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class QuickGELU(torch.nn.Module):
|
| 11 |
+
|
| 12 |
+
def forward(self, x: torch.Tensor):
|
| 13 |
+
return x * torch.sigmoid(1.702 * x)
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class ResidualAttentionBlock(torch.nn.Module):
|
| 18 |
+
|
| 19 |
+
def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):
|
| 20 |
+
super().__init__()
|
| 21 |
+
|
| 22 |
+
self.attn = torch.nn.MultiheadAttention(d_model, n_head)
|
| 23 |
+
self.ln_1 = torch.nn.LayerNorm(d_model)
|
| 24 |
+
self.mlp = torch.nn.Sequential(OrderedDict([
|
| 25 |
+
("c_fc", torch.nn.Linear(d_model, d_model * 4)),
|
| 26 |
+
("gelu", QuickGELU()),
|
| 27 |
+
("c_proj", torch.nn.Linear(d_model * 4, d_model))
|
| 28 |
+
]))
|
| 29 |
+
self.ln_2 = torch.nn.LayerNorm(d_model)
|
| 30 |
+
self.attn_mask = attn_mask
|
| 31 |
+
|
| 32 |
+
def attention(self, x: torch.Tensor):
|
| 33 |
+
self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
|
| 34 |
+
return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]
|
| 35 |
+
|
| 36 |
+
def forward(self, x: torch.Tensor):
|
| 37 |
+
x = x + self.attention(self.ln_1(x))
|
| 38 |
+
x = x + self.mlp(self.ln_2(x))
|
| 39 |
+
return x
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
class SDXLControlNetUnion(torch.nn.Module):
|
| 44 |
+
def __init__(self, global_pool=False):
|
| 45 |
+
super().__init__()
|
| 46 |
+
self.time_proj = Timesteps(320)
|
| 47 |
+
self.time_embedding = torch.nn.Sequential(
|
| 48 |
+
torch.nn.Linear(320, 1280),
|
| 49 |
+
torch.nn.SiLU(),
|
| 50 |
+
torch.nn.Linear(1280, 1280)
|
| 51 |
+
)
|
| 52 |
+
self.add_time_proj = Timesteps(256)
|
| 53 |
+
self.add_time_embedding = torch.nn.Sequential(
|
| 54 |
+
torch.nn.Linear(2816, 1280),
|
| 55 |
+
torch.nn.SiLU(),
|
| 56 |
+
torch.nn.Linear(1280, 1280)
|
| 57 |
+
)
|
| 58 |
+
self.control_type_proj = Timesteps(256)
|
| 59 |
+
self.control_type_embedding = torch.nn.Sequential(
|
| 60 |
+
torch.nn.Linear(256 * 8, 1280),
|
| 61 |
+
torch.nn.SiLU(),
|
| 62 |
+
torch.nn.Linear(1280, 1280)
|
| 63 |
+
)
|
| 64 |
+
self.conv_in = torch.nn.Conv2d(4, 320, kernel_size=3, padding=1)
|
| 65 |
+
|
| 66 |
+
self.controlnet_conv_in = ControlNetConditioningLayer(channels=(3, 16, 32, 96, 256, 320))
|
| 67 |
+
self.controlnet_transformer = ResidualAttentionBlock(320, 8)
|
| 68 |
+
self.task_embedding = torch.nn.Parameter(torch.randn(8, 320))
|
| 69 |
+
self.spatial_ch_projs = torch.nn.Linear(320, 320)
|
| 70 |
+
|
| 71 |
+
self.blocks = torch.nn.ModuleList([
|
| 72 |
+
# DownBlock2D
|
| 73 |
+
ResnetBlock(320, 320, 1280),
|
| 74 |
+
PushBlock(),
|
| 75 |
+
ResnetBlock(320, 320, 1280),
|
| 76 |
+
PushBlock(),
|
| 77 |
+
DownSampler(320),
|
| 78 |
+
PushBlock(),
|
| 79 |
+
# CrossAttnDownBlock2D
|
| 80 |
+
ResnetBlock(320, 640, 1280),
|
| 81 |
+
AttentionBlock(10, 64, 640, 2, 2048),
|
| 82 |
+
PushBlock(),
|
| 83 |
+
ResnetBlock(640, 640, 1280),
|
| 84 |
+
AttentionBlock(10, 64, 640, 2, 2048),
|
| 85 |
+
PushBlock(),
|
| 86 |
+
DownSampler(640),
|
| 87 |
+
PushBlock(),
|
| 88 |
+
# CrossAttnDownBlock2D
|
| 89 |
+
ResnetBlock(640, 1280, 1280),
|
| 90 |
+
AttentionBlock(20, 64, 1280, 10, 2048),
|
| 91 |
+
PushBlock(),
|
| 92 |
+
ResnetBlock(1280, 1280, 1280),
|
| 93 |
+
AttentionBlock(20, 64, 1280, 10, 2048),
|
| 94 |
+
PushBlock(),
|
| 95 |
+
# UNetMidBlock2DCrossAttn
|
| 96 |
+
ResnetBlock(1280, 1280, 1280),
|
| 97 |
+
AttentionBlock(20, 64, 1280, 10, 2048),
|
| 98 |
+
ResnetBlock(1280, 1280, 1280),
|
| 99 |
+
PushBlock()
|
| 100 |
+
])
|
| 101 |
+
|
| 102 |
+
self.controlnet_blocks = torch.nn.ModuleList([
|
| 103 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),
|
| 104 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),
|
| 105 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),
|
| 106 |
+
torch.nn.Conv2d(320, 320, kernel_size=(1, 1)),
|
| 107 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),
|
| 108 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),
|
| 109 |
+
torch.nn.Conv2d(640, 640, kernel_size=(1, 1)),
|
| 110 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),
|
| 111 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),
|
| 112 |
+
torch.nn.Conv2d(1280, 1280, kernel_size=(1, 1)),
|
| 113 |
+
])
|
| 114 |
+
|
| 115 |
+
self.global_pool = global_pool
|
| 116 |
+
|
| 117 |
+
# 0 -- openpose
|
| 118 |
+
# 1 -- depth
|
| 119 |
+
# 2 -- hed/pidi/scribble/ted
|
| 120 |
+
# 3 -- canny/lineart/anime_lineart/mlsd
|
| 121 |
+
# 4 -- normal
|
| 122 |
+
# 5 -- segment
|
| 123 |
+
# 6 -- tile
|
| 124 |
+
# 7 -- repaint
|
| 125 |
+
self.task_id = {
|
| 126 |
+
"openpose": 0,
|
| 127 |
+
"depth": 1,
|
| 128 |
+
"softedge": 2,
|
| 129 |
+
"canny": 3,
|
| 130 |
+
"lineart": 3,
|
| 131 |
+
"lineart_anime": 3,
|
| 132 |
+
"tile": 6,
|
| 133 |
+
"inpaint": 7
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
def fuse_condition_to_input(self, hidden_states, task_id, conditioning):
|
| 138 |
+
controlnet_cond = self.controlnet_conv_in(conditioning)
|
| 139 |
+
feat_seq = torch.mean(controlnet_cond, dim=(2, 3))
|
| 140 |
+
feat_seq = feat_seq + self.task_embedding[task_id]
|
| 141 |
+
x = torch.stack([feat_seq, torch.mean(hidden_states, dim=(2, 3))], dim=1)
|
| 142 |
+
x = self.controlnet_transformer(x)
|
| 143 |
+
|
| 144 |
+
alpha = self.spatial_ch_projs(x[:,0]).unsqueeze(-1).unsqueeze(-1)
|
| 145 |
+
controlnet_cond_fuser = controlnet_cond + alpha
|
| 146 |
+
|
| 147 |
+
hidden_states = hidden_states + controlnet_cond_fuser
|
| 148 |
+
return hidden_states
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
def forward(
|
| 152 |
+
self,
|
| 153 |
+
sample, timestep, encoder_hidden_states,
|
| 154 |
+
conditioning, processor_id, add_time_id, add_text_embeds,
|
| 155 |
+
tiled=False, tile_size=64, tile_stride=32,
|
| 156 |
+
unet:SDXLUNet=None,
|
| 157 |
+
**kwargs
|
| 158 |
+
):
|
| 159 |
+
task_id = self.task_id[processor_id]
|
| 160 |
+
|
| 161 |
+
# 1. time
|
| 162 |
+
t_emb = self.time_proj(timestep).to(sample.dtype)
|
| 163 |
+
t_emb = self.time_embedding(t_emb)
|
| 164 |
+
|
| 165 |
+
time_embeds = self.add_time_proj(add_time_id)
|
| 166 |
+
time_embeds = time_embeds.reshape((add_text_embeds.shape[0], -1))
|
| 167 |
+
add_embeds = torch.concat([add_text_embeds, time_embeds], dim=-1)
|
| 168 |
+
add_embeds = add_embeds.to(sample.dtype)
|
| 169 |
+
if unet is not None and unet.is_kolors:
|
| 170 |
+
add_embeds = unet.add_time_embedding(add_embeds)
|
| 171 |
+
else:
|
| 172 |
+
add_embeds = self.add_time_embedding(add_embeds)
|
| 173 |
+
|
| 174 |
+
control_type = torch.zeros((sample.shape[0], 8), dtype=sample.dtype, device=sample.device)
|
| 175 |
+
control_type[:, task_id] = 1
|
| 176 |
+
control_embeds = self.control_type_proj(control_type.flatten())
|
| 177 |
+
control_embeds = control_embeds.reshape((sample.shape[0], -1))
|
| 178 |
+
control_embeds = control_embeds.to(sample.dtype)
|
| 179 |
+
control_embeds = self.control_type_embedding(control_embeds)
|
| 180 |
+
time_emb = t_emb + add_embeds + control_embeds
|
| 181 |
+
|
| 182 |
+
# 2. pre-process
|
| 183 |
+
height, width = sample.shape[2], sample.shape[3]
|
| 184 |
+
hidden_states = self.conv_in(sample)
|
| 185 |
+
hidden_states = self.fuse_condition_to_input(hidden_states, task_id, conditioning)
|
| 186 |
+
text_emb = encoder_hidden_states
|
| 187 |
+
if unet is not None and unet.is_kolors:
|
| 188 |
+
text_emb = unet.text_intermediate_proj(text_emb)
|
| 189 |
+
res_stack = [hidden_states]
|
| 190 |
+
|
| 191 |
+
# 3. blocks
|
| 192 |
+
for i, block in enumerate(self.blocks):
|
| 193 |
+
if tiled and not isinstance(block, PushBlock):
|
| 194 |
+
_, _, inter_height, _ = hidden_states.shape
|
| 195 |
+
resize_scale = inter_height / height
|
| 196 |
+
hidden_states = TileWorker().tiled_forward(
|
| 197 |
+
lambda x: block(x, time_emb, text_emb, res_stack)[0],
|
| 198 |
+
hidden_states,
|
| 199 |
+
int(tile_size * resize_scale),
|
| 200 |
+
int(tile_stride * resize_scale),
|
| 201 |
+
tile_device=hidden_states.device,
|
| 202 |
+
tile_dtype=hidden_states.dtype
|
| 203 |
+
)
|
| 204 |
+
else:
|
| 205 |
+
hidden_states, _, _, _ = block(hidden_states, time_emb, text_emb, res_stack)
|
| 206 |
+
|
| 207 |
+
# 4. ControlNet blocks
|
| 208 |
+
controlnet_res_stack = [block(res) for block, res in zip(self.controlnet_blocks, res_stack)]
|
| 209 |
+
|
| 210 |
+
# pool
|
| 211 |
+
if self.global_pool:
|
| 212 |
+
controlnet_res_stack = [res.mean(dim=(2, 3), keepdim=True) for res in controlnet_res_stack]
|
| 213 |
+
|
| 214 |
+
return controlnet_res_stack
|
| 215 |
+
|
| 216 |
+
@staticmethod
|
| 217 |
+
def state_dict_converter():
|
| 218 |
+
return SDXLControlNetUnionStateDictConverter()
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
class SDXLControlNetUnionStateDictConverter:
|
| 223 |
+
def __init__(self):
|
| 224 |
+
pass
|
| 225 |
+
|
| 226 |
+
def from_diffusers(self, state_dict):
|
| 227 |
+
# architecture
|
| 228 |
+
block_types = [
|
| 229 |
+
"ResnetBlock", "PushBlock", "ResnetBlock", "PushBlock", "DownSampler", "PushBlock",
|
| 230 |
+
"ResnetBlock", "AttentionBlock", "PushBlock", "ResnetBlock", "AttentionBlock", "PushBlock", "DownSampler", "PushBlock",
|
| 231 |
+
"ResnetBlock", "AttentionBlock", "PushBlock", "ResnetBlock", "AttentionBlock", "PushBlock",
|
| 232 |
+
"ResnetBlock", "AttentionBlock", "ResnetBlock", "PushBlock"
|
| 233 |
+
]
|
| 234 |
+
|
| 235 |
+
# controlnet_rename_dict
|
| 236 |
+
controlnet_rename_dict = {
|
| 237 |
+
"controlnet_cond_embedding.conv_in.weight": "controlnet_conv_in.blocks.0.weight",
|
| 238 |
+
"controlnet_cond_embedding.conv_in.bias": "controlnet_conv_in.blocks.0.bias",
|
| 239 |
+
"controlnet_cond_embedding.blocks.0.weight": "controlnet_conv_in.blocks.2.weight",
|
| 240 |
+
"controlnet_cond_embedding.blocks.0.bias": "controlnet_conv_in.blocks.2.bias",
|
| 241 |
+
"controlnet_cond_embedding.blocks.1.weight": "controlnet_conv_in.blocks.4.weight",
|
| 242 |
+
"controlnet_cond_embedding.blocks.1.bias": "controlnet_conv_in.blocks.4.bias",
|
| 243 |
+
"controlnet_cond_embedding.blocks.2.weight": "controlnet_conv_in.blocks.6.weight",
|
| 244 |
+
"controlnet_cond_embedding.blocks.2.bias": "controlnet_conv_in.blocks.6.bias",
|
| 245 |
+
"controlnet_cond_embedding.blocks.3.weight": "controlnet_conv_in.blocks.8.weight",
|
| 246 |
+
"controlnet_cond_embedding.blocks.3.bias": "controlnet_conv_in.blocks.8.bias",
|
| 247 |
+
"controlnet_cond_embedding.blocks.4.weight": "controlnet_conv_in.blocks.10.weight",
|
| 248 |
+
"controlnet_cond_embedding.blocks.4.bias": "controlnet_conv_in.blocks.10.bias",
|
| 249 |
+
"controlnet_cond_embedding.blocks.5.weight": "controlnet_conv_in.blocks.12.weight",
|
| 250 |
+
"controlnet_cond_embedding.blocks.5.bias": "controlnet_conv_in.blocks.12.bias",
|
| 251 |
+
"controlnet_cond_embedding.conv_out.weight": "controlnet_conv_in.blocks.14.weight",
|
| 252 |
+
"controlnet_cond_embedding.conv_out.bias": "controlnet_conv_in.blocks.14.bias",
|
| 253 |
+
"control_add_embedding.linear_1.weight": "control_type_embedding.0.weight",
|
| 254 |
+
"control_add_embedding.linear_1.bias": "control_type_embedding.0.bias",
|
| 255 |
+
"control_add_embedding.linear_2.weight": "control_type_embedding.2.weight",
|
| 256 |
+
"control_add_embedding.linear_2.bias": "control_type_embedding.2.bias",
|
| 257 |
+
}
|
| 258 |
+
|
| 259 |
+
# Rename each parameter
|
| 260 |
+
name_list = sorted([name for name in state_dict])
|
| 261 |
+
rename_dict = {}
|
| 262 |
+
block_id = {"ResnetBlock": -1, "AttentionBlock": -1, "DownSampler": -1, "UpSampler": -1}
|
| 263 |
+
last_block_type_with_id = {"ResnetBlock": "", "AttentionBlock": "", "DownSampler": "", "UpSampler": ""}
|
| 264 |
+
for name in name_list:
|
| 265 |
+
names = name.split(".")
|
| 266 |
+
if names[0] in ["conv_in", "conv_norm_out", "conv_out", "task_embedding", "spatial_ch_projs"]:
|
| 267 |
+
pass
|
| 268 |
+
elif name in controlnet_rename_dict:
|
| 269 |
+
names = controlnet_rename_dict[name].split(".")
|
| 270 |
+
elif names[0] == "controlnet_down_blocks":
|
| 271 |
+
names[0] = "controlnet_blocks"
|
| 272 |
+
elif names[0] == "controlnet_mid_block":
|
| 273 |
+
names = ["controlnet_blocks", "9", names[-1]]
|
| 274 |
+
elif names[0] in ["time_embedding", "add_embedding"]:
|
| 275 |
+
if names[0] == "add_embedding":
|
| 276 |
+
names[0] = "add_time_embedding"
|
| 277 |
+
names[1] = {"linear_1": "0", "linear_2": "2"}[names[1]]
|
| 278 |
+
elif names[0] == "control_add_embedding":
|
| 279 |
+
names[0] = "control_type_embedding"
|
| 280 |
+
elif names[0] == "transformer_layes":
|
| 281 |
+
names[0] = "controlnet_transformer"
|
| 282 |
+
names.pop(1)
|
| 283 |
+
elif names[0] in ["down_blocks", "mid_block", "up_blocks"]:
|
| 284 |
+
if names[0] == "mid_block":
|
| 285 |
+
names.insert(1, "0")
|
| 286 |
+
block_type = {"resnets": "ResnetBlock", "attentions": "AttentionBlock", "downsamplers": "DownSampler", "upsamplers": "UpSampler"}[names[2]]
|
| 287 |
+
block_type_with_id = ".".join(names[:4])
|
| 288 |
+
if block_type_with_id != last_block_type_with_id[block_type]:
|
| 289 |
+
block_id[block_type] += 1
|
| 290 |
+
last_block_type_with_id[block_type] = block_type_with_id
|
| 291 |
+
while block_id[block_type] < len(block_types) and block_types[block_id[block_type]] != block_type:
|
| 292 |
+
block_id[block_type] += 1
|
| 293 |
+
block_type_with_id = ".".join(names[:4])
|
| 294 |
+
names = ["blocks", str(block_id[block_type])] + names[4:]
|
| 295 |
+
if "ff" in names:
|
| 296 |
+
ff_index = names.index("ff")
|
| 297 |
+
component = ".".join(names[ff_index:ff_index+3])
|
| 298 |
+
component = {"ff.net.0": "act_fn", "ff.net.2": "ff"}[component]
|
| 299 |
+
names = names[:ff_index] + [component] + names[ff_index+3:]
|
| 300 |
+
if "to_out" in names:
|
| 301 |
+
names.pop(names.index("to_out") + 1)
|
| 302 |
+
else:
|
| 303 |
+
print(name, state_dict[name].shape)
|
| 304 |
+
# raise ValueError(f"Unknown parameters: {name}")
|
| 305 |
+
rename_dict[name] = ".".join(names)
|
| 306 |
+
|
| 307 |
+
# Convert state_dict
|
| 308 |
+
state_dict_ = {}
|
| 309 |
+
for name, param in state_dict.items():
|
| 310 |
+
if name not in rename_dict:
|
| 311 |
+
continue
|
| 312 |
+
if ".proj_in." in name or ".proj_out." in name:
|
| 313 |
+
param = param.squeeze()
|
| 314 |
+
state_dict_[rename_dict[name]] = param
|
| 315 |
+
return state_dict_
|
| 316 |
+
|
| 317 |
+
def from_civitai(self, state_dict):
|
| 318 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/sdxl_ipadapter.py
ADDED
|
@@ -0,0 +1,122 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .svd_image_encoder import SVDImageEncoder
|
| 2 |
+
from transformers import CLIPImageProcessor
|
| 3 |
+
import torch
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class IpAdapterXLCLIPImageEmbedder(SVDImageEncoder):
|
| 7 |
+
def __init__(self):
|
| 8 |
+
super().__init__(embed_dim=1664, encoder_intermediate_size=8192, projection_dim=1280, num_encoder_layers=48, num_heads=16, head_dim=104)
|
| 9 |
+
self.image_processor = CLIPImageProcessor()
|
| 10 |
+
|
| 11 |
+
def forward(self, image):
|
| 12 |
+
pixel_values = self.image_processor(images=image, return_tensors="pt").pixel_values
|
| 13 |
+
pixel_values = pixel_values.to(device=self.embeddings.class_embedding.device, dtype=self.embeddings.class_embedding.dtype)
|
| 14 |
+
return super().forward(pixel_values)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class IpAdapterImageProjModel(torch.nn.Module):
|
| 18 |
+
def __init__(self, cross_attention_dim=2048, clip_embeddings_dim=1280, clip_extra_context_tokens=4):
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.cross_attention_dim = cross_attention_dim
|
| 21 |
+
self.clip_extra_context_tokens = clip_extra_context_tokens
|
| 22 |
+
self.proj = torch.nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
|
| 23 |
+
self.norm = torch.nn.LayerNorm(cross_attention_dim)
|
| 24 |
+
|
| 25 |
+
def forward(self, image_embeds):
|
| 26 |
+
clip_extra_context_tokens = self.proj(image_embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim)
|
| 27 |
+
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
|
| 28 |
+
return clip_extra_context_tokens
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class IpAdapterModule(torch.nn.Module):
|
| 32 |
+
def __init__(self, input_dim, output_dim):
|
| 33 |
+
super().__init__()
|
| 34 |
+
self.to_k_ip = torch.nn.Linear(input_dim, output_dim, bias=False)
|
| 35 |
+
self.to_v_ip = torch.nn.Linear(input_dim, output_dim, bias=False)
|
| 36 |
+
|
| 37 |
+
def forward(self, hidden_states):
|
| 38 |
+
ip_k = self.to_k_ip(hidden_states)
|
| 39 |
+
ip_v = self.to_v_ip(hidden_states)
|
| 40 |
+
return ip_k, ip_v
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
class SDXLIpAdapter(torch.nn.Module):
|
| 44 |
+
def __init__(self):
|
| 45 |
+
super().__init__()
|
| 46 |
+
shape_list = [(2048, 640)] * 4 + [(2048, 1280)] * 50 + [(2048, 640)] * 6 + [(2048, 1280)] * 10
|
| 47 |
+
self.ipadapter_modules = torch.nn.ModuleList([IpAdapterModule(*shape) for shape in shape_list])
|
| 48 |
+
self.image_proj = IpAdapterImageProjModel()
|
| 49 |
+
self.set_full_adapter()
|
| 50 |
+
|
| 51 |
+
def set_full_adapter(self):
|
| 52 |
+
map_list = sum([
|
| 53 |
+
[(7, i) for i in range(2)],
|
| 54 |
+
[(10, i) for i in range(2)],
|
| 55 |
+
[(15, i) for i in range(10)],
|
| 56 |
+
[(18, i) for i in range(10)],
|
| 57 |
+
[(25, i) for i in range(10)],
|
| 58 |
+
[(28, i) for i in range(10)],
|
| 59 |
+
[(31, i) for i in range(10)],
|
| 60 |
+
[(35, i) for i in range(2)],
|
| 61 |
+
[(38, i) for i in range(2)],
|
| 62 |
+
[(41, i) for i in range(2)],
|
| 63 |
+
[(21, i) for i in range(10)],
|
| 64 |
+
], [])
|
| 65 |
+
self.call_block_id = {i: j for j, i in enumerate(map_list)}
|
| 66 |
+
|
| 67 |
+
def set_less_adapter(self):
|
| 68 |
+
map_list = sum([
|
| 69 |
+
[(7, i) for i in range(2)],
|
| 70 |
+
[(10, i) for i in range(2)],
|
| 71 |
+
[(15, i) for i in range(10)],
|
| 72 |
+
[(18, i) for i in range(10)],
|
| 73 |
+
[(25, i) for i in range(10)],
|
| 74 |
+
[(28, i) for i in range(10)],
|
| 75 |
+
[(31, i) for i in range(10)],
|
| 76 |
+
[(35, i) for i in range(2)],
|
| 77 |
+
[(38, i) for i in range(2)],
|
| 78 |
+
[(41, i) for i in range(2)],
|
| 79 |
+
[(21, i) for i in range(10)],
|
| 80 |
+
], [])
|
| 81 |
+
self.call_block_id = {i: j for j, i in enumerate(map_list) if j>=34 and j<44}
|
| 82 |
+
|
| 83 |
+
def forward(self, hidden_states, scale=1.0):
|
| 84 |
+
hidden_states = self.image_proj(hidden_states)
|
| 85 |
+
hidden_states = hidden_states.view(1, -1, hidden_states.shape[-1])
|
| 86 |
+
ip_kv_dict = {}
|
| 87 |
+
for (block_id, transformer_id) in self.call_block_id:
|
| 88 |
+
ipadapter_id = self.call_block_id[(block_id, transformer_id)]
|
| 89 |
+
ip_k, ip_v = self.ipadapter_modules[ipadapter_id](hidden_states)
|
| 90 |
+
if block_id not in ip_kv_dict:
|
| 91 |
+
ip_kv_dict[block_id] = {}
|
| 92 |
+
ip_kv_dict[block_id][transformer_id] = {
|
| 93 |
+
"ip_k": ip_k,
|
| 94 |
+
"ip_v": ip_v,
|
| 95 |
+
"scale": scale
|
| 96 |
+
}
|
| 97 |
+
return ip_kv_dict
|
| 98 |
+
|
| 99 |
+
@staticmethod
|
| 100 |
+
def state_dict_converter():
|
| 101 |
+
return SDXLIpAdapterStateDictConverter()
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
class SDXLIpAdapterStateDictConverter:
|
| 105 |
+
def __init__(self):
|
| 106 |
+
pass
|
| 107 |
+
|
| 108 |
+
def from_diffusers(self, state_dict):
|
| 109 |
+
state_dict_ = {}
|
| 110 |
+
for name in state_dict["ip_adapter"]:
|
| 111 |
+
names = name.split(".")
|
| 112 |
+
layer_id = str(int(names[0]) // 2)
|
| 113 |
+
name_ = ".".join(["ipadapter_modules"] + [layer_id] + names[1:])
|
| 114 |
+
state_dict_[name_] = state_dict["ip_adapter"][name]
|
| 115 |
+
for name in state_dict["image_proj"]:
|
| 116 |
+
name_ = "image_proj." + name
|
| 117 |
+
state_dict_[name_] = state_dict["image_proj"][name]
|
| 118 |
+
return state_dict_
|
| 119 |
+
|
| 120 |
+
def from_civitai(self, state_dict):
|
| 121 |
+
return self.from_diffusers(state_dict)
|
| 122 |
+
|
diffsynth/models/sdxl_motion.py
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sd_motion import TemporalBlock
|
| 2 |
+
import torch
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class SDXLMotionModel(torch.nn.Module):
|
| 7 |
+
def __init__(self):
|
| 8 |
+
super().__init__()
|
| 9 |
+
self.motion_modules = torch.nn.ModuleList([
|
| 10 |
+
TemporalBlock(8, 320//8, 320, eps=1e-6),
|
| 11 |
+
TemporalBlock(8, 320//8, 320, eps=1e-6),
|
| 12 |
+
|
| 13 |
+
TemporalBlock(8, 640//8, 640, eps=1e-6),
|
| 14 |
+
TemporalBlock(8, 640//8, 640, eps=1e-6),
|
| 15 |
+
|
| 16 |
+
TemporalBlock(8, 1280//8, 1280, eps=1e-6),
|
| 17 |
+
TemporalBlock(8, 1280//8, 1280, eps=1e-6),
|
| 18 |
+
|
| 19 |
+
TemporalBlock(8, 1280//8, 1280, eps=1e-6),
|
| 20 |
+
TemporalBlock(8, 1280//8, 1280, eps=1e-6),
|
| 21 |
+
TemporalBlock(8, 1280//8, 1280, eps=1e-6),
|
| 22 |
+
|
| 23 |
+
TemporalBlock(8, 640//8, 640, eps=1e-6),
|
| 24 |
+
TemporalBlock(8, 640//8, 640, eps=1e-6),
|
| 25 |
+
TemporalBlock(8, 640//8, 640, eps=1e-6),
|
| 26 |
+
|
| 27 |
+
TemporalBlock(8, 320//8, 320, eps=1e-6),
|
| 28 |
+
TemporalBlock(8, 320//8, 320, eps=1e-6),
|
| 29 |
+
TemporalBlock(8, 320//8, 320, eps=1e-6),
|
| 30 |
+
])
|
| 31 |
+
self.call_block_id = {
|
| 32 |
+
0: 0,
|
| 33 |
+
2: 1,
|
| 34 |
+
7: 2,
|
| 35 |
+
10: 3,
|
| 36 |
+
15: 4,
|
| 37 |
+
18: 5,
|
| 38 |
+
25: 6,
|
| 39 |
+
28: 7,
|
| 40 |
+
31: 8,
|
| 41 |
+
35: 9,
|
| 42 |
+
38: 10,
|
| 43 |
+
41: 11,
|
| 44 |
+
44: 12,
|
| 45 |
+
46: 13,
|
| 46 |
+
48: 14,
|
| 47 |
+
}
|
| 48 |
+
|
| 49 |
+
def forward(self):
|
| 50 |
+
pass
|
| 51 |
+
|
| 52 |
+
@staticmethod
|
| 53 |
+
def state_dict_converter():
|
| 54 |
+
return SDMotionModelStateDictConverter()
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
class SDMotionModelStateDictConverter:
|
| 58 |
+
def __init__(self):
|
| 59 |
+
pass
|
| 60 |
+
|
| 61 |
+
def from_diffusers(self, state_dict):
|
| 62 |
+
rename_dict = {
|
| 63 |
+
"norm": "norm",
|
| 64 |
+
"proj_in": "proj_in",
|
| 65 |
+
"transformer_blocks.0.attention_blocks.0.to_q": "transformer_blocks.0.attn1.to_q",
|
| 66 |
+
"transformer_blocks.0.attention_blocks.0.to_k": "transformer_blocks.0.attn1.to_k",
|
| 67 |
+
"transformer_blocks.0.attention_blocks.0.to_v": "transformer_blocks.0.attn1.to_v",
|
| 68 |
+
"transformer_blocks.0.attention_blocks.0.to_out.0": "transformer_blocks.0.attn1.to_out",
|
| 69 |
+
"transformer_blocks.0.attention_blocks.0.pos_encoder": "transformer_blocks.0.pe1",
|
| 70 |
+
"transformer_blocks.0.attention_blocks.1.to_q": "transformer_blocks.0.attn2.to_q",
|
| 71 |
+
"transformer_blocks.0.attention_blocks.1.to_k": "transformer_blocks.0.attn2.to_k",
|
| 72 |
+
"transformer_blocks.0.attention_blocks.1.to_v": "transformer_blocks.0.attn2.to_v",
|
| 73 |
+
"transformer_blocks.0.attention_blocks.1.to_out.0": "transformer_blocks.0.attn2.to_out",
|
| 74 |
+
"transformer_blocks.0.attention_blocks.1.pos_encoder": "transformer_blocks.0.pe2",
|
| 75 |
+
"transformer_blocks.0.norms.0": "transformer_blocks.0.norm1",
|
| 76 |
+
"transformer_blocks.0.norms.1": "transformer_blocks.0.norm2",
|
| 77 |
+
"transformer_blocks.0.ff.net.0.proj": "transformer_blocks.0.act_fn.proj",
|
| 78 |
+
"transformer_blocks.0.ff.net.2": "transformer_blocks.0.ff",
|
| 79 |
+
"transformer_blocks.0.ff_norm": "transformer_blocks.0.norm3",
|
| 80 |
+
"proj_out": "proj_out",
|
| 81 |
+
}
|
| 82 |
+
name_list = sorted([i for i in state_dict if i.startswith("down_blocks.")])
|
| 83 |
+
name_list += sorted([i for i in state_dict if i.startswith("mid_block.")])
|
| 84 |
+
name_list += sorted([i for i in state_dict if i.startswith("up_blocks.")])
|
| 85 |
+
state_dict_ = {}
|
| 86 |
+
last_prefix, module_id = "", -1
|
| 87 |
+
for name in name_list:
|
| 88 |
+
names = name.split(".")
|
| 89 |
+
prefix_index = names.index("temporal_transformer") + 1
|
| 90 |
+
prefix = ".".join(names[:prefix_index])
|
| 91 |
+
if prefix != last_prefix:
|
| 92 |
+
last_prefix = prefix
|
| 93 |
+
module_id += 1
|
| 94 |
+
middle_name = ".".join(names[prefix_index:-1])
|
| 95 |
+
suffix = names[-1]
|
| 96 |
+
if "pos_encoder" in names:
|
| 97 |
+
rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name]])
|
| 98 |
+
else:
|
| 99 |
+
rename = ".".join(["motion_modules", str(module_id), rename_dict[middle_name], suffix])
|
| 100 |
+
state_dict_[rename] = state_dict[name]
|
| 101 |
+
return state_dict_
|
| 102 |
+
|
| 103 |
+
def from_civitai(self, state_dict):
|
| 104 |
+
return self.from_diffusers(state_dict)
|
diffsynth/models/sdxl_text_encoder.py
ADDED
|
@@ -0,0 +1,759 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from .sd_text_encoder import CLIPEncoderLayer
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class SDXLTextEncoder(torch.nn.Module):
|
| 6 |
+
def __init__(self, embed_dim=768, vocab_size=49408, max_position_embeddings=77, num_encoder_layers=11, encoder_intermediate_size=3072):
|
| 7 |
+
super().__init__()
|
| 8 |
+
|
| 9 |
+
# token_embedding
|
| 10 |
+
self.token_embedding = torch.nn.Embedding(vocab_size, embed_dim)
|
| 11 |
+
|
| 12 |
+
# position_embeds (This is a fixed tensor)
|
| 13 |
+
self.position_embeds = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, embed_dim))
|
| 14 |
+
|
| 15 |
+
# encoders
|
| 16 |
+
self.encoders = torch.nn.ModuleList([CLIPEncoderLayer(embed_dim, encoder_intermediate_size) for _ in range(num_encoder_layers)])
|
| 17 |
+
|
| 18 |
+
# attn_mask
|
| 19 |
+
self.attn_mask = self.attention_mask(max_position_embeddings)
|
| 20 |
+
|
| 21 |
+
# The text encoder is different to that in Stable Diffusion 1.x.
|
| 22 |
+
# It does not include final_layer_norm.
|
| 23 |
+
|
| 24 |
+
def attention_mask(self, length):
|
| 25 |
+
mask = torch.empty(length, length)
|
| 26 |
+
mask.fill_(float("-inf"))
|
| 27 |
+
mask.triu_(1)
|
| 28 |
+
return mask
|
| 29 |
+
|
| 30 |
+
def forward(self, input_ids, clip_skip=1):
|
| 31 |
+
embeds = self.token_embedding(input_ids) + self.position_embeds
|
| 32 |
+
attn_mask = self.attn_mask.to(device=embeds.device, dtype=embeds.dtype)
|
| 33 |
+
for encoder_id, encoder in enumerate(self.encoders):
|
| 34 |
+
embeds = encoder(embeds, attn_mask=attn_mask)
|
| 35 |
+
if encoder_id + clip_skip == len(self.encoders):
|
| 36 |
+
break
|
| 37 |
+
return embeds
|
| 38 |
+
|
| 39 |
+
@staticmethod
|
| 40 |
+
def state_dict_converter():
|
| 41 |
+
return SDXLTextEncoderStateDictConverter()
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class SDXLTextEncoder2(torch.nn.Module):
|
| 45 |
+
def __init__(self, embed_dim=1280, vocab_size=49408, max_position_embeddings=77, num_encoder_layers=32, encoder_intermediate_size=5120):
|
| 46 |
+
super().__init__()
|
| 47 |
+
|
| 48 |
+
# token_embedding
|
| 49 |
+
self.token_embedding = torch.nn.Embedding(vocab_size, embed_dim)
|
| 50 |
+
|
| 51 |
+
# position_embeds (This is a fixed tensor)
|
| 52 |
+
self.position_embeds = torch.nn.Parameter(torch.zeros(1, max_position_embeddings, embed_dim))
|
| 53 |
+
|
| 54 |
+
# encoders
|
| 55 |
+
self.encoders = torch.nn.ModuleList([CLIPEncoderLayer(embed_dim, encoder_intermediate_size, num_heads=20, head_dim=64, use_quick_gelu=False) for _ in range(num_encoder_layers)])
|
| 56 |
+
|
| 57 |
+
# attn_mask
|
| 58 |
+
self.attn_mask = self.attention_mask(max_position_embeddings)
|
| 59 |
+
|
| 60 |
+
# final_layer_norm
|
| 61 |
+
self.final_layer_norm = torch.nn.LayerNorm(embed_dim)
|
| 62 |
+
|
| 63 |
+
# text_projection
|
| 64 |
+
self.text_projection = torch.nn.Linear(embed_dim, embed_dim, bias=False)
|
| 65 |
+
|
| 66 |
+
def attention_mask(self, length):
|
| 67 |
+
mask = torch.empty(length, length)
|
| 68 |
+
mask.fill_(float("-inf"))
|
| 69 |
+
mask.triu_(1)
|
| 70 |
+
return mask
|
| 71 |
+
|
| 72 |
+
def forward(self, input_ids, clip_skip=2):
|
| 73 |
+
embeds = self.token_embedding(input_ids) + self.position_embeds
|
| 74 |
+
attn_mask = self.attn_mask.to(device=embeds.device, dtype=embeds.dtype)
|
| 75 |
+
for encoder_id, encoder in enumerate(self.encoders):
|
| 76 |
+
embeds = encoder(embeds, attn_mask=attn_mask)
|
| 77 |
+
if encoder_id + clip_skip == len(self.encoders):
|
| 78 |
+
hidden_states = embeds
|
| 79 |
+
embeds = self.final_layer_norm(embeds)
|
| 80 |
+
pooled_embeds = embeds[torch.arange(embeds.shape[0]), input_ids.to(dtype=torch.int).argmax(dim=-1)]
|
| 81 |
+
pooled_embeds = self.text_projection(pooled_embeds)
|
| 82 |
+
return pooled_embeds, hidden_states
|
| 83 |
+
|
| 84 |
+
@staticmethod
|
| 85 |
+
def state_dict_converter():
|
| 86 |
+
return SDXLTextEncoder2StateDictConverter()
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
class SDXLTextEncoderStateDictConverter:
|
| 90 |
+
def __init__(self):
|
| 91 |
+
pass
|
| 92 |
+
|
| 93 |
+
def from_diffusers(self, state_dict):
|
| 94 |
+
rename_dict = {
|
| 95 |
+
"text_model.embeddings.token_embedding.weight": "token_embedding.weight",
|
| 96 |
+
"text_model.embeddings.position_embedding.weight": "position_embeds",
|
| 97 |
+
"text_model.final_layer_norm.weight": "final_layer_norm.weight",
|
| 98 |
+
"text_model.final_layer_norm.bias": "final_layer_norm.bias"
|
| 99 |
+
}
|
| 100 |
+
attn_rename_dict = {
|
| 101 |
+
"self_attn.q_proj": "attn.to_q",
|
| 102 |
+
"self_attn.k_proj": "attn.to_k",
|
| 103 |
+
"self_attn.v_proj": "attn.to_v",
|
| 104 |
+
"self_attn.out_proj": "attn.to_out",
|
| 105 |
+
"layer_norm1": "layer_norm1",
|
| 106 |
+
"layer_norm2": "layer_norm2",
|
| 107 |
+
"mlp.fc1": "fc1",
|
| 108 |
+
"mlp.fc2": "fc2",
|
| 109 |
+
}
|
| 110 |
+
state_dict_ = {}
|
| 111 |
+
for name in state_dict:
|
| 112 |
+
if name in rename_dict:
|
| 113 |
+
param = state_dict[name]
|
| 114 |
+
if name == "text_model.embeddings.position_embedding.weight":
|
| 115 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 116 |
+
state_dict_[rename_dict[name]] = param
|
| 117 |
+
elif name.startswith("text_model.encoder.layers."):
|
| 118 |
+
param = state_dict[name]
|
| 119 |
+
names = name.split(".")
|
| 120 |
+
layer_id, layer_type, tail = names[3], ".".join(names[4:-1]), names[-1]
|
| 121 |
+
name_ = ".".join(["encoders", layer_id, attn_rename_dict[layer_type], tail])
|
| 122 |
+
state_dict_[name_] = param
|
| 123 |
+
return state_dict_
|
| 124 |
+
|
| 125 |
+
def from_civitai(self, state_dict):
|
| 126 |
+
rename_dict = {
|
| 127 |
+
"conditioner.embedders.0.transformer.text_model.embeddings.position_embedding.weight": "position_embeds",
|
| 128 |
+
"conditioner.embedders.0.transformer.text_model.embeddings.token_embedding.weight": "token_embedding.weight",
|
| 129 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm1.bias": "encoders.0.layer_norm1.bias",
|
| 130 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm1.weight": "encoders.0.layer_norm1.weight",
|
| 131 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm2.bias": "encoders.0.layer_norm2.bias",
|
| 132 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm2.weight": "encoders.0.layer_norm2.weight",
|
| 133 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.mlp.fc1.bias": "encoders.0.fc1.bias",
|
| 134 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.mlp.fc1.weight": "encoders.0.fc1.weight",
|
| 135 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.mlp.fc2.bias": "encoders.0.fc2.bias",
|
| 136 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.mlp.fc2.weight": "encoders.0.fc2.weight",
|
| 137 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.k_proj.bias": "encoders.0.attn.to_k.bias",
|
| 138 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.k_proj.weight": "encoders.0.attn.to_k.weight",
|
| 139 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.out_proj.bias": "encoders.0.attn.to_out.bias",
|
| 140 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.out_proj.weight": "encoders.0.attn.to_out.weight",
|
| 141 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.q_proj.bias": "encoders.0.attn.to_q.bias",
|
| 142 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.q_proj.weight": "encoders.0.attn.to_q.weight",
|
| 143 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.v_proj.bias": "encoders.0.attn.to_v.bias",
|
| 144 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.0.self_attn.v_proj.weight": "encoders.0.attn.to_v.weight",
|
| 145 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.layer_norm1.bias": "encoders.1.layer_norm1.bias",
|
| 146 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.layer_norm1.weight": "encoders.1.layer_norm1.weight",
|
| 147 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.layer_norm2.bias": "encoders.1.layer_norm2.bias",
|
| 148 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.layer_norm2.weight": "encoders.1.layer_norm2.weight",
|
| 149 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.mlp.fc1.bias": "encoders.1.fc1.bias",
|
| 150 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.mlp.fc1.weight": "encoders.1.fc1.weight",
|
| 151 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.mlp.fc2.bias": "encoders.1.fc2.bias",
|
| 152 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.mlp.fc2.weight": "encoders.1.fc2.weight",
|
| 153 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.k_proj.bias": "encoders.1.attn.to_k.bias",
|
| 154 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.k_proj.weight": "encoders.1.attn.to_k.weight",
|
| 155 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.out_proj.bias": "encoders.1.attn.to_out.bias",
|
| 156 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.out_proj.weight": "encoders.1.attn.to_out.weight",
|
| 157 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.q_proj.bias": "encoders.1.attn.to_q.bias",
|
| 158 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.q_proj.weight": "encoders.1.attn.to_q.weight",
|
| 159 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.v_proj.bias": "encoders.1.attn.to_v.bias",
|
| 160 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.1.self_attn.v_proj.weight": "encoders.1.attn.to_v.weight",
|
| 161 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.layer_norm1.bias": "encoders.10.layer_norm1.bias",
|
| 162 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.layer_norm1.weight": "encoders.10.layer_norm1.weight",
|
| 163 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.layer_norm2.bias": "encoders.10.layer_norm2.bias",
|
| 164 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.layer_norm2.weight": "encoders.10.layer_norm2.weight",
|
| 165 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.mlp.fc1.bias": "encoders.10.fc1.bias",
|
| 166 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.mlp.fc1.weight": "encoders.10.fc1.weight",
|
| 167 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.mlp.fc2.bias": "encoders.10.fc2.bias",
|
| 168 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.mlp.fc2.weight": "encoders.10.fc2.weight",
|
| 169 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.k_proj.bias": "encoders.10.attn.to_k.bias",
|
| 170 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.k_proj.weight": "encoders.10.attn.to_k.weight",
|
| 171 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.out_proj.bias": "encoders.10.attn.to_out.bias",
|
| 172 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.out_proj.weight": "encoders.10.attn.to_out.weight",
|
| 173 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.q_proj.bias": "encoders.10.attn.to_q.bias",
|
| 174 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.q_proj.weight": "encoders.10.attn.to_q.weight",
|
| 175 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.v_proj.bias": "encoders.10.attn.to_v.bias",
|
| 176 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.10.self_attn.v_proj.weight": "encoders.10.attn.to_v.weight",
|
| 177 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.layer_norm1.bias": "encoders.2.layer_norm1.bias",
|
| 178 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.layer_norm1.weight": "encoders.2.layer_norm1.weight",
|
| 179 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.layer_norm2.bias": "encoders.2.layer_norm2.bias",
|
| 180 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.layer_norm2.weight": "encoders.2.layer_norm2.weight",
|
| 181 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.mlp.fc1.bias": "encoders.2.fc1.bias",
|
| 182 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.mlp.fc1.weight": "encoders.2.fc1.weight",
|
| 183 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.mlp.fc2.bias": "encoders.2.fc2.bias",
|
| 184 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.mlp.fc2.weight": "encoders.2.fc2.weight",
|
| 185 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.k_proj.bias": "encoders.2.attn.to_k.bias",
|
| 186 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.k_proj.weight": "encoders.2.attn.to_k.weight",
|
| 187 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.out_proj.bias": "encoders.2.attn.to_out.bias",
|
| 188 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.out_proj.weight": "encoders.2.attn.to_out.weight",
|
| 189 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.q_proj.bias": "encoders.2.attn.to_q.bias",
|
| 190 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.q_proj.weight": "encoders.2.attn.to_q.weight",
|
| 191 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.v_proj.bias": "encoders.2.attn.to_v.bias",
|
| 192 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.2.self_attn.v_proj.weight": "encoders.2.attn.to_v.weight",
|
| 193 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.layer_norm1.bias": "encoders.3.layer_norm1.bias",
|
| 194 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.layer_norm1.weight": "encoders.3.layer_norm1.weight",
|
| 195 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.layer_norm2.bias": "encoders.3.layer_norm2.bias",
|
| 196 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.layer_norm2.weight": "encoders.3.layer_norm2.weight",
|
| 197 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.mlp.fc1.bias": "encoders.3.fc1.bias",
|
| 198 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.mlp.fc1.weight": "encoders.3.fc1.weight",
|
| 199 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.mlp.fc2.bias": "encoders.3.fc2.bias",
|
| 200 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.mlp.fc2.weight": "encoders.3.fc2.weight",
|
| 201 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.k_proj.bias": "encoders.3.attn.to_k.bias",
|
| 202 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.k_proj.weight": "encoders.3.attn.to_k.weight",
|
| 203 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.out_proj.bias": "encoders.3.attn.to_out.bias",
|
| 204 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.out_proj.weight": "encoders.3.attn.to_out.weight",
|
| 205 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.q_proj.bias": "encoders.3.attn.to_q.bias",
|
| 206 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.q_proj.weight": "encoders.3.attn.to_q.weight",
|
| 207 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.v_proj.bias": "encoders.3.attn.to_v.bias",
|
| 208 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.3.self_attn.v_proj.weight": "encoders.3.attn.to_v.weight",
|
| 209 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.layer_norm1.bias": "encoders.4.layer_norm1.bias",
|
| 210 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.layer_norm1.weight": "encoders.4.layer_norm1.weight",
|
| 211 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.layer_norm2.bias": "encoders.4.layer_norm2.bias",
|
| 212 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.layer_norm2.weight": "encoders.4.layer_norm2.weight",
|
| 213 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.mlp.fc1.bias": "encoders.4.fc1.bias",
|
| 214 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.mlp.fc1.weight": "encoders.4.fc1.weight",
|
| 215 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.mlp.fc2.bias": "encoders.4.fc2.bias",
|
| 216 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.mlp.fc2.weight": "encoders.4.fc2.weight",
|
| 217 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.k_proj.bias": "encoders.4.attn.to_k.bias",
|
| 218 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.k_proj.weight": "encoders.4.attn.to_k.weight",
|
| 219 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.out_proj.bias": "encoders.4.attn.to_out.bias",
|
| 220 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.out_proj.weight": "encoders.4.attn.to_out.weight",
|
| 221 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.q_proj.bias": "encoders.4.attn.to_q.bias",
|
| 222 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.q_proj.weight": "encoders.4.attn.to_q.weight",
|
| 223 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.v_proj.bias": "encoders.4.attn.to_v.bias",
|
| 224 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.4.self_attn.v_proj.weight": "encoders.4.attn.to_v.weight",
|
| 225 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.layer_norm1.bias": "encoders.5.layer_norm1.bias",
|
| 226 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.layer_norm1.weight": "encoders.5.layer_norm1.weight",
|
| 227 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.layer_norm2.bias": "encoders.5.layer_norm2.bias",
|
| 228 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.layer_norm2.weight": "encoders.5.layer_norm2.weight",
|
| 229 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.mlp.fc1.bias": "encoders.5.fc1.bias",
|
| 230 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.mlp.fc1.weight": "encoders.5.fc1.weight",
|
| 231 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.mlp.fc2.bias": "encoders.5.fc2.bias",
|
| 232 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.mlp.fc2.weight": "encoders.5.fc2.weight",
|
| 233 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.k_proj.bias": "encoders.5.attn.to_k.bias",
|
| 234 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.k_proj.weight": "encoders.5.attn.to_k.weight",
|
| 235 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.out_proj.bias": "encoders.5.attn.to_out.bias",
|
| 236 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.out_proj.weight": "encoders.5.attn.to_out.weight",
|
| 237 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.q_proj.bias": "encoders.5.attn.to_q.bias",
|
| 238 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.q_proj.weight": "encoders.5.attn.to_q.weight",
|
| 239 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.v_proj.bias": "encoders.5.attn.to_v.bias",
|
| 240 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.5.self_attn.v_proj.weight": "encoders.5.attn.to_v.weight",
|
| 241 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.layer_norm1.bias": "encoders.6.layer_norm1.bias",
|
| 242 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.layer_norm1.weight": "encoders.6.layer_norm1.weight",
|
| 243 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.layer_norm2.bias": "encoders.6.layer_norm2.bias",
|
| 244 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.layer_norm2.weight": "encoders.6.layer_norm2.weight",
|
| 245 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.mlp.fc1.bias": "encoders.6.fc1.bias",
|
| 246 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.mlp.fc1.weight": "encoders.6.fc1.weight",
|
| 247 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.mlp.fc2.bias": "encoders.6.fc2.bias",
|
| 248 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.mlp.fc2.weight": "encoders.6.fc2.weight",
|
| 249 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.k_proj.bias": "encoders.6.attn.to_k.bias",
|
| 250 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.k_proj.weight": "encoders.6.attn.to_k.weight",
|
| 251 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.out_proj.bias": "encoders.6.attn.to_out.bias",
|
| 252 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.out_proj.weight": "encoders.6.attn.to_out.weight",
|
| 253 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.q_proj.bias": "encoders.6.attn.to_q.bias",
|
| 254 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.q_proj.weight": "encoders.6.attn.to_q.weight",
|
| 255 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.v_proj.bias": "encoders.6.attn.to_v.bias",
|
| 256 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.6.self_attn.v_proj.weight": "encoders.6.attn.to_v.weight",
|
| 257 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.layer_norm1.bias": "encoders.7.layer_norm1.bias",
|
| 258 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.layer_norm1.weight": "encoders.7.layer_norm1.weight",
|
| 259 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.layer_norm2.bias": "encoders.7.layer_norm2.bias",
|
| 260 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.layer_norm2.weight": "encoders.7.layer_norm2.weight",
|
| 261 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.mlp.fc1.bias": "encoders.7.fc1.bias",
|
| 262 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.mlp.fc1.weight": "encoders.7.fc1.weight",
|
| 263 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.mlp.fc2.bias": "encoders.7.fc2.bias",
|
| 264 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.mlp.fc2.weight": "encoders.7.fc2.weight",
|
| 265 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.k_proj.bias": "encoders.7.attn.to_k.bias",
|
| 266 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.k_proj.weight": "encoders.7.attn.to_k.weight",
|
| 267 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.out_proj.bias": "encoders.7.attn.to_out.bias",
|
| 268 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.out_proj.weight": "encoders.7.attn.to_out.weight",
|
| 269 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.q_proj.bias": "encoders.7.attn.to_q.bias",
|
| 270 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.q_proj.weight": "encoders.7.attn.to_q.weight",
|
| 271 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.v_proj.bias": "encoders.7.attn.to_v.bias",
|
| 272 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.7.self_attn.v_proj.weight": "encoders.7.attn.to_v.weight",
|
| 273 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.layer_norm1.bias": "encoders.8.layer_norm1.bias",
|
| 274 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.layer_norm1.weight": "encoders.8.layer_norm1.weight",
|
| 275 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.layer_norm2.bias": "encoders.8.layer_norm2.bias",
|
| 276 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.layer_norm2.weight": "encoders.8.layer_norm2.weight",
|
| 277 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.mlp.fc1.bias": "encoders.8.fc1.bias",
|
| 278 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.mlp.fc1.weight": "encoders.8.fc1.weight",
|
| 279 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.mlp.fc2.bias": "encoders.8.fc2.bias",
|
| 280 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.mlp.fc2.weight": "encoders.8.fc2.weight",
|
| 281 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.k_proj.bias": "encoders.8.attn.to_k.bias",
|
| 282 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.k_proj.weight": "encoders.8.attn.to_k.weight",
|
| 283 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.out_proj.bias": "encoders.8.attn.to_out.bias",
|
| 284 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.out_proj.weight": "encoders.8.attn.to_out.weight",
|
| 285 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.q_proj.bias": "encoders.8.attn.to_q.bias",
|
| 286 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.q_proj.weight": "encoders.8.attn.to_q.weight",
|
| 287 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.v_proj.bias": "encoders.8.attn.to_v.bias",
|
| 288 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.8.self_attn.v_proj.weight": "encoders.8.attn.to_v.weight",
|
| 289 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.layer_norm1.bias": "encoders.9.layer_norm1.bias",
|
| 290 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.layer_norm1.weight": "encoders.9.layer_norm1.weight",
|
| 291 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.layer_norm2.bias": "encoders.9.layer_norm2.bias",
|
| 292 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.layer_norm2.weight": "encoders.9.layer_norm2.weight",
|
| 293 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.mlp.fc1.bias": "encoders.9.fc1.bias",
|
| 294 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.mlp.fc1.weight": "encoders.9.fc1.weight",
|
| 295 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.mlp.fc2.bias": "encoders.9.fc2.bias",
|
| 296 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.mlp.fc2.weight": "encoders.9.fc2.weight",
|
| 297 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.k_proj.bias": "encoders.9.attn.to_k.bias",
|
| 298 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.k_proj.weight": "encoders.9.attn.to_k.weight",
|
| 299 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.out_proj.bias": "encoders.9.attn.to_out.bias",
|
| 300 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.out_proj.weight": "encoders.9.attn.to_out.weight",
|
| 301 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.q_proj.bias": "encoders.9.attn.to_q.bias",
|
| 302 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.q_proj.weight": "encoders.9.attn.to_q.weight",
|
| 303 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.v_proj.bias": "encoders.9.attn.to_v.bias",
|
| 304 |
+
"conditioner.embedders.0.transformer.text_model.encoder.layers.9.self_attn.v_proj.weight": "encoders.9.attn.to_v.weight",
|
| 305 |
+
}
|
| 306 |
+
state_dict_ = {}
|
| 307 |
+
for name in state_dict:
|
| 308 |
+
if name in rename_dict:
|
| 309 |
+
param = state_dict[name]
|
| 310 |
+
if name == "conditioner.embedders.0.transformer.text_model.embeddings.position_embedding.weight":
|
| 311 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 312 |
+
state_dict_[rename_dict[name]] = param
|
| 313 |
+
return state_dict_
|
| 314 |
+
|
| 315 |
+
|
| 316 |
+
class SDXLTextEncoder2StateDictConverter:
|
| 317 |
+
def __init__(self):
|
| 318 |
+
pass
|
| 319 |
+
|
| 320 |
+
def from_diffusers(self, state_dict):
|
| 321 |
+
rename_dict = {
|
| 322 |
+
"text_model.embeddings.token_embedding.weight": "token_embedding.weight",
|
| 323 |
+
"text_model.embeddings.position_embedding.weight": "position_embeds",
|
| 324 |
+
"text_model.final_layer_norm.weight": "final_layer_norm.weight",
|
| 325 |
+
"text_model.final_layer_norm.bias": "final_layer_norm.bias",
|
| 326 |
+
"text_projection.weight": "text_projection.weight"
|
| 327 |
+
}
|
| 328 |
+
attn_rename_dict = {
|
| 329 |
+
"self_attn.q_proj": "attn.to_q",
|
| 330 |
+
"self_attn.k_proj": "attn.to_k",
|
| 331 |
+
"self_attn.v_proj": "attn.to_v",
|
| 332 |
+
"self_attn.out_proj": "attn.to_out",
|
| 333 |
+
"layer_norm1": "layer_norm1",
|
| 334 |
+
"layer_norm2": "layer_norm2",
|
| 335 |
+
"mlp.fc1": "fc1",
|
| 336 |
+
"mlp.fc2": "fc2",
|
| 337 |
+
}
|
| 338 |
+
state_dict_ = {}
|
| 339 |
+
for name in state_dict:
|
| 340 |
+
if name in rename_dict:
|
| 341 |
+
param = state_dict[name]
|
| 342 |
+
if name == "text_model.embeddings.position_embedding.weight":
|
| 343 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 344 |
+
state_dict_[rename_dict[name]] = param
|
| 345 |
+
elif name.startswith("text_model.encoder.layers."):
|
| 346 |
+
param = state_dict[name]
|
| 347 |
+
names = name.split(".")
|
| 348 |
+
layer_id, layer_type, tail = names[3], ".".join(names[4:-1]), names[-1]
|
| 349 |
+
name_ = ".".join(["encoders", layer_id, attn_rename_dict[layer_type], tail])
|
| 350 |
+
state_dict_[name_] = param
|
| 351 |
+
return state_dict_
|
| 352 |
+
|
| 353 |
+
def from_civitai(self, state_dict):
|
| 354 |
+
rename_dict = {
|
| 355 |
+
"conditioner.embedders.1.model.ln_final.bias": "final_layer_norm.bias",
|
| 356 |
+
"conditioner.embedders.1.model.ln_final.weight": "final_layer_norm.weight",
|
| 357 |
+
"conditioner.embedders.1.model.positional_embedding": "position_embeds",
|
| 358 |
+
"conditioner.embedders.1.model.token_embedding.weight": "token_embedding.weight",
|
| 359 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_bias": ['encoders.0.attn.to_q.bias', 'encoders.0.attn.to_k.bias', 'encoders.0.attn.to_v.bias'],
|
| 360 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight": ['encoders.0.attn.to_q.weight', 'encoders.0.attn.to_k.weight', 'encoders.0.attn.to_v.weight'],
|
| 361 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.attn.out_proj.bias": "encoders.0.attn.to_out.bias",
|
| 362 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.attn.out_proj.weight": "encoders.0.attn.to_out.weight",
|
| 363 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.ln_1.bias": "encoders.0.layer_norm1.bias",
|
| 364 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.ln_1.weight": "encoders.0.layer_norm1.weight",
|
| 365 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.ln_2.bias": "encoders.0.layer_norm2.bias",
|
| 366 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.ln_2.weight": "encoders.0.layer_norm2.weight",
|
| 367 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_fc.bias": "encoders.0.fc1.bias",
|
| 368 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_fc.weight": "encoders.0.fc1.weight",
|
| 369 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_proj.bias": "encoders.0.fc2.bias",
|
| 370 |
+
"conditioner.embedders.1.model.transformer.resblocks.0.mlp.c_proj.weight": "encoders.0.fc2.weight",
|
| 371 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.attn.in_proj_bias": ['encoders.1.attn.to_q.bias', 'encoders.1.attn.to_k.bias', 'encoders.1.attn.to_v.bias'],
|
| 372 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.attn.in_proj_weight": ['encoders.1.attn.to_q.weight', 'encoders.1.attn.to_k.weight', 'encoders.1.attn.to_v.weight'],
|
| 373 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.attn.out_proj.bias": "encoders.1.attn.to_out.bias",
|
| 374 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.attn.out_proj.weight": "encoders.1.attn.to_out.weight",
|
| 375 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.ln_1.bias": "encoders.1.layer_norm1.bias",
|
| 376 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.ln_1.weight": "encoders.1.layer_norm1.weight",
|
| 377 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.ln_2.bias": "encoders.1.layer_norm2.bias",
|
| 378 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.ln_2.weight": "encoders.1.layer_norm2.weight",
|
| 379 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_fc.bias": "encoders.1.fc1.bias",
|
| 380 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_fc.weight": "encoders.1.fc1.weight",
|
| 381 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_proj.bias": "encoders.1.fc2.bias",
|
| 382 |
+
"conditioner.embedders.1.model.transformer.resblocks.1.mlp.c_proj.weight": "encoders.1.fc2.weight",
|
| 383 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.attn.in_proj_bias": ['encoders.10.attn.to_q.bias', 'encoders.10.attn.to_k.bias', 'encoders.10.attn.to_v.bias'],
|
| 384 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.attn.in_proj_weight": ['encoders.10.attn.to_q.weight', 'encoders.10.attn.to_k.weight', 'encoders.10.attn.to_v.weight'],
|
| 385 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.attn.out_proj.bias": "encoders.10.attn.to_out.bias",
|
| 386 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.attn.out_proj.weight": "encoders.10.attn.to_out.weight",
|
| 387 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.ln_1.bias": "encoders.10.layer_norm1.bias",
|
| 388 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.ln_1.weight": "encoders.10.layer_norm1.weight",
|
| 389 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.ln_2.bias": "encoders.10.layer_norm2.bias",
|
| 390 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.ln_2.weight": "encoders.10.layer_norm2.weight",
|
| 391 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_fc.bias": "encoders.10.fc1.bias",
|
| 392 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_fc.weight": "encoders.10.fc1.weight",
|
| 393 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_proj.bias": "encoders.10.fc2.bias",
|
| 394 |
+
"conditioner.embedders.1.model.transformer.resblocks.10.mlp.c_proj.weight": "encoders.10.fc2.weight",
|
| 395 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.attn.in_proj_bias": ['encoders.11.attn.to_q.bias', 'encoders.11.attn.to_k.bias', 'encoders.11.attn.to_v.bias'],
|
| 396 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.attn.in_proj_weight": ['encoders.11.attn.to_q.weight', 'encoders.11.attn.to_k.weight', 'encoders.11.attn.to_v.weight'],
|
| 397 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.attn.out_proj.bias": "encoders.11.attn.to_out.bias",
|
| 398 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.attn.out_proj.weight": "encoders.11.attn.to_out.weight",
|
| 399 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.ln_1.bias": "encoders.11.layer_norm1.bias",
|
| 400 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.ln_1.weight": "encoders.11.layer_norm1.weight",
|
| 401 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.ln_2.bias": "encoders.11.layer_norm2.bias",
|
| 402 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.ln_2.weight": "encoders.11.layer_norm2.weight",
|
| 403 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_fc.bias": "encoders.11.fc1.bias",
|
| 404 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_fc.weight": "encoders.11.fc1.weight",
|
| 405 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_proj.bias": "encoders.11.fc2.bias",
|
| 406 |
+
"conditioner.embedders.1.model.transformer.resblocks.11.mlp.c_proj.weight": "encoders.11.fc2.weight",
|
| 407 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.attn.in_proj_bias": ['encoders.12.attn.to_q.bias', 'encoders.12.attn.to_k.bias', 'encoders.12.attn.to_v.bias'],
|
| 408 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.attn.in_proj_weight": ['encoders.12.attn.to_q.weight', 'encoders.12.attn.to_k.weight', 'encoders.12.attn.to_v.weight'],
|
| 409 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.attn.out_proj.bias": "encoders.12.attn.to_out.bias",
|
| 410 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.attn.out_proj.weight": "encoders.12.attn.to_out.weight",
|
| 411 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.ln_1.bias": "encoders.12.layer_norm1.bias",
|
| 412 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.ln_1.weight": "encoders.12.layer_norm1.weight",
|
| 413 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.ln_2.bias": "encoders.12.layer_norm2.bias",
|
| 414 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.ln_2.weight": "encoders.12.layer_norm2.weight",
|
| 415 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_fc.bias": "encoders.12.fc1.bias",
|
| 416 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_fc.weight": "encoders.12.fc1.weight",
|
| 417 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_proj.bias": "encoders.12.fc2.bias",
|
| 418 |
+
"conditioner.embedders.1.model.transformer.resblocks.12.mlp.c_proj.weight": "encoders.12.fc2.weight",
|
| 419 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.attn.in_proj_bias": ['encoders.13.attn.to_q.bias', 'encoders.13.attn.to_k.bias', 'encoders.13.attn.to_v.bias'],
|
| 420 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.attn.in_proj_weight": ['encoders.13.attn.to_q.weight', 'encoders.13.attn.to_k.weight', 'encoders.13.attn.to_v.weight'],
|
| 421 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.attn.out_proj.bias": "encoders.13.attn.to_out.bias",
|
| 422 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.attn.out_proj.weight": "encoders.13.attn.to_out.weight",
|
| 423 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.ln_1.bias": "encoders.13.layer_norm1.bias",
|
| 424 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.ln_1.weight": "encoders.13.layer_norm1.weight",
|
| 425 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.ln_2.bias": "encoders.13.layer_norm2.bias",
|
| 426 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.ln_2.weight": "encoders.13.layer_norm2.weight",
|
| 427 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_fc.bias": "encoders.13.fc1.bias",
|
| 428 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_fc.weight": "encoders.13.fc1.weight",
|
| 429 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_proj.bias": "encoders.13.fc2.bias",
|
| 430 |
+
"conditioner.embedders.1.model.transformer.resblocks.13.mlp.c_proj.weight": "encoders.13.fc2.weight",
|
| 431 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.attn.in_proj_bias": ['encoders.14.attn.to_q.bias', 'encoders.14.attn.to_k.bias', 'encoders.14.attn.to_v.bias'],
|
| 432 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.attn.in_proj_weight": ['encoders.14.attn.to_q.weight', 'encoders.14.attn.to_k.weight', 'encoders.14.attn.to_v.weight'],
|
| 433 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.attn.out_proj.bias": "encoders.14.attn.to_out.bias",
|
| 434 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.attn.out_proj.weight": "encoders.14.attn.to_out.weight",
|
| 435 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.ln_1.bias": "encoders.14.layer_norm1.bias",
|
| 436 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.ln_1.weight": "encoders.14.layer_norm1.weight",
|
| 437 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.ln_2.bias": "encoders.14.layer_norm2.bias",
|
| 438 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.ln_2.weight": "encoders.14.layer_norm2.weight",
|
| 439 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_fc.bias": "encoders.14.fc1.bias",
|
| 440 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_fc.weight": "encoders.14.fc1.weight",
|
| 441 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_proj.bias": "encoders.14.fc2.bias",
|
| 442 |
+
"conditioner.embedders.1.model.transformer.resblocks.14.mlp.c_proj.weight": "encoders.14.fc2.weight",
|
| 443 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.attn.in_proj_bias": ['encoders.15.attn.to_q.bias', 'encoders.15.attn.to_k.bias', 'encoders.15.attn.to_v.bias'],
|
| 444 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.attn.in_proj_weight": ['encoders.15.attn.to_q.weight', 'encoders.15.attn.to_k.weight', 'encoders.15.attn.to_v.weight'],
|
| 445 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.attn.out_proj.bias": "encoders.15.attn.to_out.bias",
|
| 446 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.attn.out_proj.weight": "encoders.15.attn.to_out.weight",
|
| 447 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.ln_1.bias": "encoders.15.layer_norm1.bias",
|
| 448 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.ln_1.weight": "encoders.15.layer_norm1.weight",
|
| 449 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.ln_2.bias": "encoders.15.layer_norm2.bias",
|
| 450 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.ln_2.weight": "encoders.15.layer_norm2.weight",
|
| 451 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_fc.bias": "encoders.15.fc1.bias",
|
| 452 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_fc.weight": "encoders.15.fc1.weight",
|
| 453 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_proj.bias": "encoders.15.fc2.bias",
|
| 454 |
+
"conditioner.embedders.1.model.transformer.resblocks.15.mlp.c_proj.weight": "encoders.15.fc2.weight",
|
| 455 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.attn.in_proj_bias": ['encoders.16.attn.to_q.bias', 'encoders.16.attn.to_k.bias', 'encoders.16.attn.to_v.bias'],
|
| 456 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.attn.in_proj_weight": ['encoders.16.attn.to_q.weight', 'encoders.16.attn.to_k.weight', 'encoders.16.attn.to_v.weight'],
|
| 457 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.attn.out_proj.bias": "encoders.16.attn.to_out.bias",
|
| 458 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.attn.out_proj.weight": "encoders.16.attn.to_out.weight",
|
| 459 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.ln_1.bias": "encoders.16.layer_norm1.bias",
|
| 460 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.ln_1.weight": "encoders.16.layer_norm1.weight",
|
| 461 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.ln_2.bias": "encoders.16.layer_norm2.bias",
|
| 462 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.ln_2.weight": "encoders.16.layer_norm2.weight",
|
| 463 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_fc.bias": "encoders.16.fc1.bias",
|
| 464 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_fc.weight": "encoders.16.fc1.weight",
|
| 465 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_proj.bias": "encoders.16.fc2.bias",
|
| 466 |
+
"conditioner.embedders.1.model.transformer.resblocks.16.mlp.c_proj.weight": "encoders.16.fc2.weight",
|
| 467 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.attn.in_proj_bias": ['encoders.17.attn.to_q.bias', 'encoders.17.attn.to_k.bias', 'encoders.17.attn.to_v.bias'],
|
| 468 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.attn.in_proj_weight": ['encoders.17.attn.to_q.weight', 'encoders.17.attn.to_k.weight', 'encoders.17.attn.to_v.weight'],
|
| 469 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.attn.out_proj.bias": "encoders.17.attn.to_out.bias",
|
| 470 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.attn.out_proj.weight": "encoders.17.attn.to_out.weight",
|
| 471 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.ln_1.bias": "encoders.17.layer_norm1.bias",
|
| 472 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.ln_1.weight": "encoders.17.layer_norm1.weight",
|
| 473 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.ln_2.bias": "encoders.17.layer_norm2.bias",
|
| 474 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.ln_2.weight": "encoders.17.layer_norm2.weight",
|
| 475 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_fc.bias": "encoders.17.fc1.bias",
|
| 476 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_fc.weight": "encoders.17.fc1.weight",
|
| 477 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_proj.bias": "encoders.17.fc2.bias",
|
| 478 |
+
"conditioner.embedders.1.model.transformer.resblocks.17.mlp.c_proj.weight": "encoders.17.fc2.weight",
|
| 479 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.attn.in_proj_bias": ['encoders.18.attn.to_q.bias', 'encoders.18.attn.to_k.bias', 'encoders.18.attn.to_v.bias'],
|
| 480 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.attn.in_proj_weight": ['encoders.18.attn.to_q.weight', 'encoders.18.attn.to_k.weight', 'encoders.18.attn.to_v.weight'],
|
| 481 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.attn.out_proj.bias": "encoders.18.attn.to_out.bias",
|
| 482 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.attn.out_proj.weight": "encoders.18.attn.to_out.weight",
|
| 483 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.ln_1.bias": "encoders.18.layer_norm1.bias",
|
| 484 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.ln_1.weight": "encoders.18.layer_norm1.weight",
|
| 485 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.ln_2.bias": "encoders.18.layer_norm2.bias",
|
| 486 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.ln_2.weight": "encoders.18.layer_norm2.weight",
|
| 487 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_fc.bias": "encoders.18.fc1.bias",
|
| 488 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_fc.weight": "encoders.18.fc1.weight",
|
| 489 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_proj.bias": "encoders.18.fc2.bias",
|
| 490 |
+
"conditioner.embedders.1.model.transformer.resblocks.18.mlp.c_proj.weight": "encoders.18.fc2.weight",
|
| 491 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.attn.in_proj_bias": ['encoders.19.attn.to_q.bias', 'encoders.19.attn.to_k.bias', 'encoders.19.attn.to_v.bias'],
|
| 492 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.attn.in_proj_weight": ['encoders.19.attn.to_q.weight', 'encoders.19.attn.to_k.weight', 'encoders.19.attn.to_v.weight'],
|
| 493 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.attn.out_proj.bias": "encoders.19.attn.to_out.bias",
|
| 494 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.attn.out_proj.weight": "encoders.19.attn.to_out.weight",
|
| 495 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.ln_1.bias": "encoders.19.layer_norm1.bias",
|
| 496 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.ln_1.weight": "encoders.19.layer_norm1.weight",
|
| 497 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.ln_2.bias": "encoders.19.layer_norm2.bias",
|
| 498 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.ln_2.weight": "encoders.19.layer_norm2.weight",
|
| 499 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_fc.bias": "encoders.19.fc1.bias",
|
| 500 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_fc.weight": "encoders.19.fc1.weight",
|
| 501 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_proj.bias": "encoders.19.fc2.bias",
|
| 502 |
+
"conditioner.embedders.1.model.transformer.resblocks.19.mlp.c_proj.weight": "encoders.19.fc2.weight",
|
| 503 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.attn.in_proj_bias": ['encoders.2.attn.to_q.bias', 'encoders.2.attn.to_k.bias', 'encoders.2.attn.to_v.bias'],
|
| 504 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.attn.in_proj_weight": ['encoders.2.attn.to_q.weight', 'encoders.2.attn.to_k.weight', 'encoders.2.attn.to_v.weight'],
|
| 505 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.attn.out_proj.bias": "encoders.2.attn.to_out.bias",
|
| 506 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.attn.out_proj.weight": "encoders.2.attn.to_out.weight",
|
| 507 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.ln_1.bias": "encoders.2.layer_norm1.bias",
|
| 508 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.ln_1.weight": "encoders.2.layer_norm1.weight",
|
| 509 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.ln_2.bias": "encoders.2.layer_norm2.bias",
|
| 510 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.ln_2.weight": "encoders.2.layer_norm2.weight",
|
| 511 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_fc.bias": "encoders.2.fc1.bias",
|
| 512 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_fc.weight": "encoders.2.fc1.weight",
|
| 513 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_proj.bias": "encoders.2.fc2.bias",
|
| 514 |
+
"conditioner.embedders.1.model.transformer.resblocks.2.mlp.c_proj.weight": "encoders.2.fc2.weight",
|
| 515 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.attn.in_proj_bias": ['encoders.20.attn.to_q.bias', 'encoders.20.attn.to_k.bias', 'encoders.20.attn.to_v.bias'],
|
| 516 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.attn.in_proj_weight": ['encoders.20.attn.to_q.weight', 'encoders.20.attn.to_k.weight', 'encoders.20.attn.to_v.weight'],
|
| 517 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.attn.out_proj.bias": "encoders.20.attn.to_out.bias",
|
| 518 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.attn.out_proj.weight": "encoders.20.attn.to_out.weight",
|
| 519 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.ln_1.bias": "encoders.20.layer_norm1.bias",
|
| 520 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.ln_1.weight": "encoders.20.layer_norm1.weight",
|
| 521 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.ln_2.bias": "encoders.20.layer_norm2.bias",
|
| 522 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.ln_2.weight": "encoders.20.layer_norm2.weight",
|
| 523 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_fc.bias": "encoders.20.fc1.bias",
|
| 524 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_fc.weight": "encoders.20.fc1.weight",
|
| 525 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_proj.bias": "encoders.20.fc2.bias",
|
| 526 |
+
"conditioner.embedders.1.model.transformer.resblocks.20.mlp.c_proj.weight": "encoders.20.fc2.weight",
|
| 527 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.attn.in_proj_bias": ['encoders.21.attn.to_q.bias', 'encoders.21.attn.to_k.bias', 'encoders.21.attn.to_v.bias'],
|
| 528 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.attn.in_proj_weight": ['encoders.21.attn.to_q.weight', 'encoders.21.attn.to_k.weight', 'encoders.21.attn.to_v.weight'],
|
| 529 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.attn.out_proj.bias": "encoders.21.attn.to_out.bias",
|
| 530 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.attn.out_proj.weight": "encoders.21.attn.to_out.weight",
|
| 531 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.ln_1.bias": "encoders.21.layer_norm1.bias",
|
| 532 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.ln_1.weight": "encoders.21.layer_norm1.weight",
|
| 533 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.ln_2.bias": "encoders.21.layer_norm2.bias",
|
| 534 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.ln_2.weight": "encoders.21.layer_norm2.weight",
|
| 535 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_fc.bias": "encoders.21.fc1.bias",
|
| 536 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_fc.weight": "encoders.21.fc1.weight",
|
| 537 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_proj.bias": "encoders.21.fc2.bias",
|
| 538 |
+
"conditioner.embedders.1.model.transformer.resblocks.21.mlp.c_proj.weight": "encoders.21.fc2.weight",
|
| 539 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.attn.in_proj_bias": ['encoders.22.attn.to_q.bias', 'encoders.22.attn.to_k.bias', 'encoders.22.attn.to_v.bias'],
|
| 540 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.attn.in_proj_weight": ['encoders.22.attn.to_q.weight', 'encoders.22.attn.to_k.weight', 'encoders.22.attn.to_v.weight'],
|
| 541 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.attn.out_proj.bias": "encoders.22.attn.to_out.bias",
|
| 542 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.attn.out_proj.weight": "encoders.22.attn.to_out.weight",
|
| 543 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.ln_1.bias": "encoders.22.layer_norm1.bias",
|
| 544 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.ln_1.weight": "encoders.22.layer_norm1.weight",
|
| 545 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.ln_2.bias": "encoders.22.layer_norm2.bias",
|
| 546 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.ln_2.weight": "encoders.22.layer_norm2.weight",
|
| 547 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_fc.bias": "encoders.22.fc1.bias",
|
| 548 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_fc.weight": "encoders.22.fc1.weight",
|
| 549 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_proj.bias": "encoders.22.fc2.bias",
|
| 550 |
+
"conditioner.embedders.1.model.transformer.resblocks.22.mlp.c_proj.weight": "encoders.22.fc2.weight",
|
| 551 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.attn.in_proj_bias": ['encoders.23.attn.to_q.bias', 'encoders.23.attn.to_k.bias', 'encoders.23.attn.to_v.bias'],
|
| 552 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.attn.in_proj_weight": ['encoders.23.attn.to_q.weight', 'encoders.23.attn.to_k.weight', 'encoders.23.attn.to_v.weight'],
|
| 553 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.attn.out_proj.bias": "encoders.23.attn.to_out.bias",
|
| 554 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.attn.out_proj.weight": "encoders.23.attn.to_out.weight",
|
| 555 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.ln_1.bias": "encoders.23.layer_norm1.bias",
|
| 556 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.ln_1.weight": "encoders.23.layer_norm1.weight",
|
| 557 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.ln_2.bias": "encoders.23.layer_norm2.bias",
|
| 558 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.ln_2.weight": "encoders.23.layer_norm2.weight",
|
| 559 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_fc.bias": "encoders.23.fc1.bias",
|
| 560 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_fc.weight": "encoders.23.fc1.weight",
|
| 561 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_proj.bias": "encoders.23.fc2.bias",
|
| 562 |
+
"conditioner.embedders.1.model.transformer.resblocks.23.mlp.c_proj.weight": "encoders.23.fc2.weight",
|
| 563 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.attn.in_proj_bias": ['encoders.24.attn.to_q.bias', 'encoders.24.attn.to_k.bias', 'encoders.24.attn.to_v.bias'],
|
| 564 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.attn.in_proj_weight": ['encoders.24.attn.to_q.weight', 'encoders.24.attn.to_k.weight', 'encoders.24.attn.to_v.weight'],
|
| 565 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.attn.out_proj.bias": "encoders.24.attn.to_out.bias",
|
| 566 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.attn.out_proj.weight": "encoders.24.attn.to_out.weight",
|
| 567 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.ln_1.bias": "encoders.24.layer_norm1.bias",
|
| 568 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.ln_1.weight": "encoders.24.layer_norm1.weight",
|
| 569 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.ln_2.bias": "encoders.24.layer_norm2.bias",
|
| 570 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.ln_2.weight": "encoders.24.layer_norm2.weight",
|
| 571 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_fc.bias": "encoders.24.fc1.bias",
|
| 572 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_fc.weight": "encoders.24.fc1.weight",
|
| 573 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_proj.bias": "encoders.24.fc2.bias",
|
| 574 |
+
"conditioner.embedders.1.model.transformer.resblocks.24.mlp.c_proj.weight": "encoders.24.fc2.weight",
|
| 575 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.attn.in_proj_bias": ['encoders.25.attn.to_q.bias', 'encoders.25.attn.to_k.bias', 'encoders.25.attn.to_v.bias'],
|
| 576 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.attn.in_proj_weight": ['encoders.25.attn.to_q.weight', 'encoders.25.attn.to_k.weight', 'encoders.25.attn.to_v.weight'],
|
| 577 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.attn.out_proj.bias": "encoders.25.attn.to_out.bias",
|
| 578 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.attn.out_proj.weight": "encoders.25.attn.to_out.weight",
|
| 579 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.ln_1.bias": "encoders.25.layer_norm1.bias",
|
| 580 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.ln_1.weight": "encoders.25.layer_norm1.weight",
|
| 581 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.ln_2.bias": "encoders.25.layer_norm2.bias",
|
| 582 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.ln_2.weight": "encoders.25.layer_norm2.weight",
|
| 583 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_fc.bias": "encoders.25.fc1.bias",
|
| 584 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_fc.weight": "encoders.25.fc1.weight",
|
| 585 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_proj.bias": "encoders.25.fc2.bias",
|
| 586 |
+
"conditioner.embedders.1.model.transformer.resblocks.25.mlp.c_proj.weight": "encoders.25.fc2.weight",
|
| 587 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.attn.in_proj_bias": ['encoders.26.attn.to_q.bias', 'encoders.26.attn.to_k.bias', 'encoders.26.attn.to_v.bias'],
|
| 588 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.attn.in_proj_weight": ['encoders.26.attn.to_q.weight', 'encoders.26.attn.to_k.weight', 'encoders.26.attn.to_v.weight'],
|
| 589 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.attn.out_proj.bias": "encoders.26.attn.to_out.bias",
|
| 590 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.attn.out_proj.weight": "encoders.26.attn.to_out.weight",
|
| 591 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.ln_1.bias": "encoders.26.layer_norm1.bias",
|
| 592 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.ln_1.weight": "encoders.26.layer_norm1.weight",
|
| 593 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.ln_2.bias": "encoders.26.layer_norm2.bias",
|
| 594 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.ln_2.weight": "encoders.26.layer_norm2.weight",
|
| 595 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_fc.bias": "encoders.26.fc1.bias",
|
| 596 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_fc.weight": "encoders.26.fc1.weight",
|
| 597 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_proj.bias": "encoders.26.fc2.bias",
|
| 598 |
+
"conditioner.embedders.1.model.transformer.resblocks.26.mlp.c_proj.weight": "encoders.26.fc2.weight",
|
| 599 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.attn.in_proj_bias": ['encoders.27.attn.to_q.bias', 'encoders.27.attn.to_k.bias', 'encoders.27.attn.to_v.bias'],
|
| 600 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.attn.in_proj_weight": ['encoders.27.attn.to_q.weight', 'encoders.27.attn.to_k.weight', 'encoders.27.attn.to_v.weight'],
|
| 601 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.attn.out_proj.bias": "encoders.27.attn.to_out.bias",
|
| 602 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.attn.out_proj.weight": "encoders.27.attn.to_out.weight",
|
| 603 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.ln_1.bias": "encoders.27.layer_norm1.bias",
|
| 604 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.ln_1.weight": "encoders.27.layer_norm1.weight",
|
| 605 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.ln_2.bias": "encoders.27.layer_norm2.bias",
|
| 606 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.ln_2.weight": "encoders.27.layer_norm2.weight",
|
| 607 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_fc.bias": "encoders.27.fc1.bias",
|
| 608 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_fc.weight": "encoders.27.fc1.weight",
|
| 609 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_proj.bias": "encoders.27.fc2.bias",
|
| 610 |
+
"conditioner.embedders.1.model.transformer.resblocks.27.mlp.c_proj.weight": "encoders.27.fc2.weight",
|
| 611 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.attn.in_proj_bias": ['encoders.28.attn.to_q.bias', 'encoders.28.attn.to_k.bias', 'encoders.28.attn.to_v.bias'],
|
| 612 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.attn.in_proj_weight": ['encoders.28.attn.to_q.weight', 'encoders.28.attn.to_k.weight', 'encoders.28.attn.to_v.weight'],
|
| 613 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.attn.out_proj.bias": "encoders.28.attn.to_out.bias",
|
| 614 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.attn.out_proj.weight": "encoders.28.attn.to_out.weight",
|
| 615 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.ln_1.bias": "encoders.28.layer_norm1.bias",
|
| 616 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.ln_1.weight": "encoders.28.layer_norm1.weight",
|
| 617 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.ln_2.bias": "encoders.28.layer_norm2.bias",
|
| 618 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.ln_2.weight": "encoders.28.layer_norm2.weight",
|
| 619 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_fc.bias": "encoders.28.fc1.bias",
|
| 620 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_fc.weight": "encoders.28.fc1.weight",
|
| 621 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_proj.bias": "encoders.28.fc2.bias",
|
| 622 |
+
"conditioner.embedders.1.model.transformer.resblocks.28.mlp.c_proj.weight": "encoders.28.fc2.weight",
|
| 623 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.attn.in_proj_bias": ['encoders.29.attn.to_q.bias', 'encoders.29.attn.to_k.bias', 'encoders.29.attn.to_v.bias'],
|
| 624 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.attn.in_proj_weight": ['encoders.29.attn.to_q.weight', 'encoders.29.attn.to_k.weight', 'encoders.29.attn.to_v.weight'],
|
| 625 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.attn.out_proj.bias": "encoders.29.attn.to_out.bias",
|
| 626 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.attn.out_proj.weight": "encoders.29.attn.to_out.weight",
|
| 627 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.ln_1.bias": "encoders.29.layer_norm1.bias",
|
| 628 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.ln_1.weight": "encoders.29.layer_norm1.weight",
|
| 629 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.ln_2.bias": "encoders.29.layer_norm2.bias",
|
| 630 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.ln_2.weight": "encoders.29.layer_norm2.weight",
|
| 631 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_fc.bias": "encoders.29.fc1.bias",
|
| 632 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_fc.weight": "encoders.29.fc1.weight",
|
| 633 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_proj.bias": "encoders.29.fc2.bias",
|
| 634 |
+
"conditioner.embedders.1.model.transformer.resblocks.29.mlp.c_proj.weight": "encoders.29.fc2.weight",
|
| 635 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.attn.in_proj_bias": ['encoders.3.attn.to_q.bias', 'encoders.3.attn.to_k.bias', 'encoders.3.attn.to_v.bias'],
|
| 636 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.attn.in_proj_weight": ['encoders.3.attn.to_q.weight', 'encoders.3.attn.to_k.weight', 'encoders.3.attn.to_v.weight'],
|
| 637 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.attn.out_proj.bias": "encoders.3.attn.to_out.bias",
|
| 638 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.attn.out_proj.weight": "encoders.3.attn.to_out.weight",
|
| 639 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.ln_1.bias": "encoders.3.layer_norm1.bias",
|
| 640 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.ln_1.weight": "encoders.3.layer_norm1.weight",
|
| 641 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.ln_2.bias": "encoders.3.layer_norm2.bias",
|
| 642 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.ln_2.weight": "encoders.3.layer_norm2.weight",
|
| 643 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_fc.bias": "encoders.3.fc1.bias",
|
| 644 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_fc.weight": "encoders.3.fc1.weight",
|
| 645 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_proj.bias": "encoders.3.fc2.bias",
|
| 646 |
+
"conditioner.embedders.1.model.transformer.resblocks.3.mlp.c_proj.weight": "encoders.3.fc2.weight",
|
| 647 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.attn.in_proj_bias": ['encoders.30.attn.to_q.bias', 'encoders.30.attn.to_k.bias', 'encoders.30.attn.to_v.bias'],
|
| 648 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.attn.in_proj_weight": ['encoders.30.attn.to_q.weight', 'encoders.30.attn.to_k.weight', 'encoders.30.attn.to_v.weight'],
|
| 649 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.attn.out_proj.bias": "encoders.30.attn.to_out.bias",
|
| 650 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.attn.out_proj.weight": "encoders.30.attn.to_out.weight",
|
| 651 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.ln_1.bias": "encoders.30.layer_norm1.bias",
|
| 652 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.ln_1.weight": "encoders.30.layer_norm1.weight",
|
| 653 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.ln_2.bias": "encoders.30.layer_norm2.bias",
|
| 654 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.ln_2.weight": "encoders.30.layer_norm2.weight",
|
| 655 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_fc.bias": "encoders.30.fc1.bias",
|
| 656 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_fc.weight": "encoders.30.fc1.weight",
|
| 657 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_proj.bias": "encoders.30.fc2.bias",
|
| 658 |
+
"conditioner.embedders.1.model.transformer.resblocks.30.mlp.c_proj.weight": "encoders.30.fc2.weight",
|
| 659 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.attn.in_proj_bias": ['encoders.31.attn.to_q.bias', 'encoders.31.attn.to_k.bias', 'encoders.31.attn.to_v.bias'],
|
| 660 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.attn.in_proj_weight": ['encoders.31.attn.to_q.weight', 'encoders.31.attn.to_k.weight', 'encoders.31.attn.to_v.weight'],
|
| 661 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.attn.out_proj.bias": "encoders.31.attn.to_out.bias",
|
| 662 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.attn.out_proj.weight": "encoders.31.attn.to_out.weight",
|
| 663 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.ln_1.bias": "encoders.31.layer_norm1.bias",
|
| 664 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.ln_1.weight": "encoders.31.layer_norm1.weight",
|
| 665 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.ln_2.bias": "encoders.31.layer_norm2.bias",
|
| 666 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.ln_2.weight": "encoders.31.layer_norm2.weight",
|
| 667 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_fc.bias": "encoders.31.fc1.bias",
|
| 668 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_fc.weight": "encoders.31.fc1.weight",
|
| 669 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_proj.bias": "encoders.31.fc2.bias",
|
| 670 |
+
"conditioner.embedders.1.model.transformer.resblocks.31.mlp.c_proj.weight": "encoders.31.fc2.weight",
|
| 671 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.attn.in_proj_bias": ['encoders.4.attn.to_q.bias', 'encoders.4.attn.to_k.bias', 'encoders.4.attn.to_v.bias'],
|
| 672 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.attn.in_proj_weight": ['encoders.4.attn.to_q.weight', 'encoders.4.attn.to_k.weight', 'encoders.4.attn.to_v.weight'],
|
| 673 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.attn.out_proj.bias": "encoders.4.attn.to_out.bias",
|
| 674 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.attn.out_proj.weight": "encoders.4.attn.to_out.weight",
|
| 675 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.ln_1.bias": "encoders.4.layer_norm1.bias",
|
| 676 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.ln_1.weight": "encoders.4.layer_norm1.weight",
|
| 677 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.ln_2.bias": "encoders.4.layer_norm2.bias",
|
| 678 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.ln_2.weight": "encoders.4.layer_norm2.weight",
|
| 679 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_fc.bias": "encoders.4.fc1.bias",
|
| 680 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_fc.weight": "encoders.4.fc1.weight",
|
| 681 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_proj.bias": "encoders.4.fc2.bias",
|
| 682 |
+
"conditioner.embedders.1.model.transformer.resblocks.4.mlp.c_proj.weight": "encoders.4.fc2.weight",
|
| 683 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.attn.in_proj_bias": ['encoders.5.attn.to_q.bias', 'encoders.5.attn.to_k.bias', 'encoders.5.attn.to_v.bias'],
|
| 684 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.attn.in_proj_weight": ['encoders.5.attn.to_q.weight', 'encoders.5.attn.to_k.weight', 'encoders.5.attn.to_v.weight'],
|
| 685 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.attn.out_proj.bias": "encoders.5.attn.to_out.bias",
|
| 686 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.attn.out_proj.weight": "encoders.5.attn.to_out.weight",
|
| 687 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.ln_1.bias": "encoders.5.layer_norm1.bias",
|
| 688 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.ln_1.weight": "encoders.5.layer_norm1.weight",
|
| 689 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.ln_2.bias": "encoders.5.layer_norm2.bias",
|
| 690 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.ln_2.weight": "encoders.5.layer_norm2.weight",
|
| 691 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_fc.bias": "encoders.5.fc1.bias",
|
| 692 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_fc.weight": "encoders.5.fc1.weight",
|
| 693 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_proj.bias": "encoders.5.fc2.bias",
|
| 694 |
+
"conditioner.embedders.1.model.transformer.resblocks.5.mlp.c_proj.weight": "encoders.5.fc2.weight",
|
| 695 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.attn.in_proj_bias": ['encoders.6.attn.to_q.bias', 'encoders.6.attn.to_k.bias', 'encoders.6.attn.to_v.bias'],
|
| 696 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.attn.in_proj_weight": ['encoders.6.attn.to_q.weight', 'encoders.6.attn.to_k.weight', 'encoders.6.attn.to_v.weight'],
|
| 697 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.attn.out_proj.bias": "encoders.6.attn.to_out.bias",
|
| 698 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.attn.out_proj.weight": "encoders.6.attn.to_out.weight",
|
| 699 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.ln_1.bias": "encoders.6.layer_norm1.bias",
|
| 700 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.ln_1.weight": "encoders.6.layer_norm1.weight",
|
| 701 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.ln_2.bias": "encoders.6.layer_norm2.bias",
|
| 702 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.ln_2.weight": "encoders.6.layer_norm2.weight",
|
| 703 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_fc.bias": "encoders.6.fc1.bias",
|
| 704 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_fc.weight": "encoders.6.fc1.weight",
|
| 705 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_proj.bias": "encoders.6.fc2.bias",
|
| 706 |
+
"conditioner.embedders.1.model.transformer.resblocks.6.mlp.c_proj.weight": "encoders.6.fc2.weight",
|
| 707 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.attn.in_proj_bias": ['encoders.7.attn.to_q.bias', 'encoders.7.attn.to_k.bias', 'encoders.7.attn.to_v.bias'],
|
| 708 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.attn.in_proj_weight": ['encoders.7.attn.to_q.weight', 'encoders.7.attn.to_k.weight', 'encoders.7.attn.to_v.weight'],
|
| 709 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.attn.out_proj.bias": "encoders.7.attn.to_out.bias",
|
| 710 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.attn.out_proj.weight": "encoders.7.attn.to_out.weight",
|
| 711 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.ln_1.bias": "encoders.7.layer_norm1.bias",
|
| 712 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.ln_1.weight": "encoders.7.layer_norm1.weight",
|
| 713 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.ln_2.bias": "encoders.7.layer_norm2.bias",
|
| 714 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.ln_2.weight": "encoders.7.layer_norm2.weight",
|
| 715 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_fc.bias": "encoders.7.fc1.bias",
|
| 716 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_fc.weight": "encoders.7.fc1.weight",
|
| 717 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_proj.bias": "encoders.7.fc2.bias",
|
| 718 |
+
"conditioner.embedders.1.model.transformer.resblocks.7.mlp.c_proj.weight": "encoders.7.fc2.weight",
|
| 719 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.attn.in_proj_bias": ['encoders.8.attn.to_q.bias', 'encoders.8.attn.to_k.bias', 'encoders.8.attn.to_v.bias'],
|
| 720 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.attn.in_proj_weight": ['encoders.8.attn.to_q.weight', 'encoders.8.attn.to_k.weight', 'encoders.8.attn.to_v.weight'],
|
| 721 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.attn.out_proj.bias": "encoders.8.attn.to_out.bias",
|
| 722 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.attn.out_proj.weight": "encoders.8.attn.to_out.weight",
|
| 723 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.ln_1.bias": "encoders.8.layer_norm1.bias",
|
| 724 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.ln_1.weight": "encoders.8.layer_norm1.weight",
|
| 725 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.ln_2.bias": "encoders.8.layer_norm2.bias",
|
| 726 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.ln_2.weight": "encoders.8.layer_norm2.weight",
|
| 727 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_fc.bias": "encoders.8.fc1.bias",
|
| 728 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_fc.weight": "encoders.8.fc1.weight",
|
| 729 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_proj.bias": "encoders.8.fc2.bias",
|
| 730 |
+
"conditioner.embedders.1.model.transformer.resblocks.8.mlp.c_proj.weight": "encoders.8.fc2.weight",
|
| 731 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.attn.in_proj_bias": ['encoders.9.attn.to_q.bias', 'encoders.9.attn.to_k.bias', 'encoders.9.attn.to_v.bias'],
|
| 732 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.attn.in_proj_weight": ['encoders.9.attn.to_q.weight', 'encoders.9.attn.to_k.weight', 'encoders.9.attn.to_v.weight'],
|
| 733 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.attn.out_proj.bias": "encoders.9.attn.to_out.bias",
|
| 734 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.attn.out_proj.weight": "encoders.9.attn.to_out.weight",
|
| 735 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.ln_1.bias": "encoders.9.layer_norm1.bias",
|
| 736 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.ln_1.weight": "encoders.9.layer_norm1.weight",
|
| 737 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.ln_2.bias": "encoders.9.layer_norm2.bias",
|
| 738 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.ln_2.weight": "encoders.9.layer_norm2.weight",
|
| 739 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_fc.bias": "encoders.9.fc1.bias",
|
| 740 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_fc.weight": "encoders.9.fc1.weight",
|
| 741 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_proj.bias": "encoders.9.fc2.bias",
|
| 742 |
+
"conditioner.embedders.1.model.transformer.resblocks.9.mlp.c_proj.weight": "encoders.9.fc2.weight",
|
| 743 |
+
"conditioner.embedders.1.model.text_projection": "text_projection.weight",
|
| 744 |
+
}
|
| 745 |
+
state_dict_ = {}
|
| 746 |
+
for name in state_dict:
|
| 747 |
+
if name in rename_dict:
|
| 748 |
+
param = state_dict[name]
|
| 749 |
+
if name == "conditioner.embedders.1.model.positional_embedding":
|
| 750 |
+
param = param.reshape((1, param.shape[0], param.shape[1]))
|
| 751 |
+
elif name == "conditioner.embedders.1.model.text_projection":
|
| 752 |
+
param = param.T
|
| 753 |
+
if isinstance(rename_dict[name], str):
|
| 754 |
+
state_dict_[rename_dict[name]] = param
|
| 755 |
+
else:
|
| 756 |
+
length = param.shape[0] // 3
|
| 757 |
+
for i, rename in enumerate(rename_dict[name]):
|
| 758 |
+
state_dict_[rename] = param[i*length: i*length+length]
|
| 759 |
+
return state_dict_
|
diffsynth/models/sdxl_unet.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
diffsynth/models/sdxl_vae_decoder.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sd_vae_decoder import SDVAEDecoder, SDVAEDecoderStateDictConverter
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
class SDXLVAEDecoder(SDVAEDecoder):
|
| 5 |
+
def __init__(self, upcast_to_float32=True):
|
| 6 |
+
super().__init__()
|
| 7 |
+
self.scaling_factor = 0.13025
|
| 8 |
+
|
| 9 |
+
@staticmethod
|
| 10 |
+
def state_dict_converter():
|
| 11 |
+
return SDXLVAEDecoderStateDictConverter()
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class SDXLVAEDecoderStateDictConverter(SDVAEDecoderStateDictConverter):
|
| 15 |
+
def __init__(self):
|
| 16 |
+
super().__init__()
|
| 17 |
+
|
| 18 |
+
def from_diffusers(self, state_dict):
|
| 19 |
+
state_dict = super().from_diffusers(state_dict)
|
| 20 |
+
return state_dict, {"upcast_to_float32": True}
|
| 21 |
+
|
| 22 |
+
def from_civitai(self, state_dict):
|
| 23 |
+
state_dict = super().from_civitai(state_dict)
|
| 24 |
+
return state_dict, {"upcast_to_float32": True}
|
diffsynth/models/sdxl_vae_encoder.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .sd_vae_encoder import SDVAEEncoderStateDictConverter, SDVAEEncoder
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
class SDXLVAEEncoder(SDVAEEncoder):
|
| 5 |
+
def __init__(self, upcast_to_float32=True):
|
| 6 |
+
super().__init__()
|
| 7 |
+
self.scaling_factor = 0.13025
|
| 8 |
+
|
| 9 |
+
@staticmethod
|
| 10 |
+
def state_dict_converter():
|
| 11 |
+
return SDXLVAEEncoderStateDictConverter()
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class SDXLVAEEncoderStateDictConverter(SDVAEEncoderStateDictConverter):
|
| 15 |
+
def __init__(self):
|
| 16 |
+
super().__init__()
|
| 17 |
+
|
| 18 |
+
def from_diffusers(self, state_dict):
|
| 19 |
+
state_dict = super().from_diffusers(state_dict)
|
| 20 |
+
return state_dict, {"upcast_to_float32": True}
|
| 21 |
+
|
| 22 |
+
def from_civitai(self, state_dict):
|
| 23 |
+
state_dict = super().from_civitai(state_dict)
|
| 24 |
+
return state_dict, {"upcast_to_float32": True}
|
diffsynth/models/step1x_connector.py
ADDED
|
@@ -0,0 +1,683 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
|
| 3 |
+
import torch, math
|
| 4 |
+
import torch.nn
|
| 5 |
+
from einops import rearrange
|
| 6 |
+
from torch import nn
|
| 7 |
+
from functools import partial
|
| 8 |
+
from einops import rearrange
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def attention(q, k, v, attn_mask, mode="torch"):
|
| 13 |
+
q = q.transpose(1, 2)
|
| 14 |
+
k = k.transpose(1, 2)
|
| 15 |
+
v = v.transpose(1, 2)
|
| 16 |
+
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 17 |
+
x = rearrange(x, "b n s d -> b s (n d)")
|
| 18 |
+
return x
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class MLP(nn.Module):
|
| 23 |
+
"""MLP as used in Vision Transformer, MLP-Mixer and related networks"""
|
| 24 |
+
|
| 25 |
+
def __init__(
|
| 26 |
+
self,
|
| 27 |
+
in_channels,
|
| 28 |
+
hidden_channels=None,
|
| 29 |
+
out_features=None,
|
| 30 |
+
act_layer=nn.GELU,
|
| 31 |
+
norm_layer=None,
|
| 32 |
+
bias=True,
|
| 33 |
+
drop=0.0,
|
| 34 |
+
use_conv=False,
|
| 35 |
+
device=None,
|
| 36 |
+
dtype=None,
|
| 37 |
+
):
|
| 38 |
+
super().__init__()
|
| 39 |
+
out_features = out_features or in_channels
|
| 40 |
+
hidden_channels = hidden_channels or in_channels
|
| 41 |
+
bias = (bias, bias)
|
| 42 |
+
drop_probs = (drop, drop)
|
| 43 |
+
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
|
| 44 |
+
|
| 45 |
+
self.fc1 = linear_layer(
|
| 46 |
+
in_channels, hidden_channels, bias=bias[0], device=device, dtype=dtype
|
| 47 |
+
)
|
| 48 |
+
self.act = act_layer()
|
| 49 |
+
self.drop1 = nn.Dropout(drop_probs[0])
|
| 50 |
+
self.norm = (
|
| 51 |
+
norm_layer(hidden_channels, device=device, dtype=dtype)
|
| 52 |
+
if norm_layer is not None
|
| 53 |
+
else nn.Identity()
|
| 54 |
+
)
|
| 55 |
+
self.fc2 = linear_layer(
|
| 56 |
+
hidden_channels, out_features, bias=bias[1], device=device, dtype=dtype
|
| 57 |
+
)
|
| 58 |
+
self.drop2 = nn.Dropout(drop_probs[1])
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
x = self.fc1(x)
|
| 62 |
+
x = self.act(x)
|
| 63 |
+
x = self.drop1(x)
|
| 64 |
+
x = self.norm(x)
|
| 65 |
+
x = self.fc2(x)
|
| 66 |
+
x = self.drop2(x)
|
| 67 |
+
return x
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
class TextProjection(nn.Module):
|
| 71 |
+
"""
|
| 72 |
+
Projects text embeddings. Also handles dropout for classifier-free guidance.
|
| 73 |
+
|
| 74 |
+
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
|
| 75 |
+
"""
|
| 76 |
+
|
| 77 |
+
def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):
|
| 78 |
+
factory_kwargs = {"dtype": dtype, "device": device}
|
| 79 |
+
super().__init__()
|
| 80 |
+
self.linear_1 = nn.Linear(
|
| 81 |
+
in_features=in_channels,
|
| 82 |
+
out_features=hidden_size,
|
| 83 |
+
bias=True,
|
| 84 |
+
**factory_kwargs,
|
| 85 |
+
)
|
| 86 |
+
self.act_1 = act_layer()
|
| 87 |
+
self.linear_2 = nn.Linear(
|
| 88 |
+
in_features=hidden_size,
|
| 89 |
+
out_features=hidden_size,
|
| 90 |
+
bias=True,
|
| 91 |
+
**factory_kwargs,
|
| 92 |
+
)
|
| 93 |
+
|
| 94 |
+
def forward(self, caption):
|
| 95 |
+
hidden_states = self.linear_1(caption)
|
| 96 |
+
hidden_states = self.act_1(hidden_states)
|
| 97 |
+
hidden_states = self.linear_2(hidden_states)
|
| 98 |
+
return hidden_states
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
class TimestepEmbedder(nn.Module):
|
| 102 |
+
"""
|
| 103 |
+
Embeds scalar timesteps into vector representations.
|
| 104 |
+
"""
|
| 105 |
+
|
| 106 |
+
def __init__(
|
| 107 |
+
self,
|
| 108 |
+
hidden_size,
|
| 109 |
+
act_layer,
|
| 110 |
+
frequency_embedding_size=256,
|
| 111 |
+
max_period=10000,
|
| 112 |
+
out_size=None,
|
| 113 |
+
dtype=None,
|
| 114 |
+
device=None,
|
| 115 |
+
):
|
| 116 |
+
factory_kwargs = {"dtype": dtype, "device": device}
|
| 117 |
+
super().__init__()
|
| 118 |
+
self.frequency_embedding_size = frequency_embedding_size
|
| 119 |
+
self.max_period = max_period
|
| 120 |
+
if out_size is None:
|
| 121 |
+
out_size = hidden_size
|
| 122 |
+
|
| 123 |
+
self.mlp = nn.Sequential(
|
| 124 |
+
nn.Linear(
|
| 125 |
+
frequency_embedding_size, hidden_size, bias=True, **factory_kwargs
|
| 126 |
+
),
|
| 127 |
+
act_layer(),
|
| 128 |
+
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
|
| 129 |
+
)
|
| 130 |
+
nn.init.normal_(self.mlp[0].weight, std=0.02) # type: ignore
|
| 131 |
+
nn.init.normal_(self.mlp[2].weight, std=0.02) # type: ignore
|
| 132 |
+
|
| 133 |
+
@staticmethod
|
| 134 |
+
def timestep_embedding(t, dim, max_period=10000):
|
| 135 |
+
"""
|
| 136 |
+
Create sinusoidal timestep embeddings.
|
| 137 |
+
|
| 138 |
+
Args:
|
| 139 |
+
t (torch.Tensor): a 1-D Tensor of N indices, one per batch element. These may be fractional.
|
| 140 |
+
dim (int): the dimension of the output.
|
| 141 |
+
max_period (int): controls the minimum frequency of the embeddings.
|
| 142 |
+
|
| 143 |
+
Returns:
|
| 144 |
+
embedding (torch.Tensor): An (N, D) Tensor of positional embeddings.
|
| 145 |
+
|
| 146 |
+
.. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
| 147 |
+
"""
|
| 148 |
+
half = dim // 2
|
| 149 |
+
freqs = torch.exp(
|
| 150 |
+
-math.log(max_period)
|
| 151 |
+
* torch.arange(start=0, end=half, dtype=torch.float32)
|
| 152 |
+
/ half
|
| 153 |
+
).to(device=t.device)
|
| 154 |
+
args = t[:, None].float() * freqs[None]
|
| 155 |
+
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
| 156 |
+
if dim % 2:
|
| 157 |
+
embedding = torch.cat(
|
| 158 |
+
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
|
| 159 |
+
)
|
| 160 |
+
return embedding
|
| 161 |
+
|
| 162 |
+
def forward(self, t):
|
| 163 |
+
t_freq = self.timestep_embedding(
|
| 164 |
+
t, self.frequency_embedding_size, self.max_period
|
| 165 |
+
).type(t.dtype) # type: ignore
|
| 166 |
+
t_emb = self.mlp(t_freq)
|
| 167 |
+
return t_emb
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
def apply_gate(x, gate=None, tanh=False):
|
| 171 |
+
"""AI is creating summary for apply_gate
|
| 172 |
+
|
| 173 |
+
Args:
|
| 174 |
+
x (torch.Tensor): input tensor.
|
| 175 |
+
gate (torch.Tensor, optional): gate tensor. Defaults to None.
|
| 176 |
+
tanh (bool, optional): whether to use tanh function. Defaults to False.
|
| 177 |
+
|
| 178 |
+
Returns:
|
| 179 |
+
torch.Tensor: the output tensor after apply gate.
|
| 180 |
+
"""
|
| 181 |
+
if gate is None:
|
| 182 |
+
return x
|
| 183 |
+
if tanh:
|
| 184 |
+
return x * gate.unsqueeze(1).tanh()
|
| 185 |
+
else:
|
| 186 |
+
return x * gate.unsqueeze(1)
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
class RMSNorm(nn.Module):
|
| 190 |
+
def __init__(
|
| 191 |
+
self,
|
| 192 |
+
dim: int,
|
| 193 |
+
elementwise_affine=True,
|
| 194 |
+
eps: float = 1e-6,
|
| 195 |
+
device=None,
|
| 196 |
+
dtype=None,
|
| 197 |
+
):
|
| 198 |
+
"""
|
| 199 |
+
Initialize the RMSNorm normalization layer.
|
| 200 |
+
|
| 201 |
+
Args:
|
| 202 |
+
dim (int): The dimension of the input tensor.
|
| 203 |
+
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
|
| 204 |
+
|
| 205 |
+
Attributes:
|
| 206 |
+
eps (float): A small value added to the denominator for numerical stability.
|
| 207 |
+
weight (nn.Parameter): Learnable scaling parameter.
|
| 208 |
+
|
| 209 |
+
"""
|
| 210 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 211 |
+
super().__init__()
|
| 212 |
+
self.eps = eps
|
| 213 |
+
if elementwise_affine:
|
| 214 |
+
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
|
| 215 |
+
|
| 216 |
+
def _norm(self, x):
|
| 217 |
+
"""
|
| 218 |
+
Apply the RMSNorm normalization to the input tensor.
|
| 219 |
+
|
| 220 |
+
Args:
|
| 221 |
+
x (torch.Tensor): The input tensor.
|
| 222 |
+
|
| 223 |
+
Returns:
|
| 224 |
+
torch.Tensor: The normalized tensor.
|
| 225 |
+
|
| 226 |
+
"""
|
| 227 |
+
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
| 228 |
+
|
| 229 |
+
def forward(self, x):
|
| 230 |
+
"""
|
| 231 |
+
Forward pass through the RMSNorm layer.
|
| 232 |
+
|
| 233 |
+
Args:
|
| 234 |
+
x (torch.Tensor): The input tensor.
|
| 235 |
+
|
| 236 |
+
Returns:
|
| 237 |
+
torch.Tensor: The output tensor after applying RMSNorm.
|
| 238 |
+
|
| 239 |
+
"""
|
| 240 |
+
output = self._norm(x.float()).type_as(x)
|
| 241 |
+
if hasattr(self, "weight"):
|
| 242 |
+
output = output * self.weight
|
| 243 |
+
return output
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def get_norm_layer(norm_layer):
|
| 247 |
+
"""
|
| 248 |
+
Get the normalization layer.
|
| 249 |
+
|
| 250 |
+
Args:
|
| 251 |
+
norm_layer (str): The type of normalization layer.
|
| 252 |
+
|
| 253 |
+
Returns:
|
| 254 |
+
norm_layer (nn.Module): The normalization layer.
|
| 255 |
+
"""
|
| 256 |
+
if norm_layer == "layer":
|
| 257 |
+
return nn.LayerNorm
|
| 258 |
+
elif norm_layer == "rms":
|
| 259 |
+
return RMSNorm
|
| 260 |
+
else:
|
| 261 |
+
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
def get_activation_layer(act_type):
|
| 265 |
+
"""get activation layer
|
| 266 |
+
|
| 267 |
+
Args:
|
| 268 |
+
act_type (str): the activation type
|
| 269 |
+
|
| 270 |
+
Returns:
|
| 271 |
+
torch.nn.functional: the activation layer
|
| 272 |
+
"""
|
| 273 |
+
if act_type == "gelu":
|
| 274 |
+
return lambda: nn.GELU()
|
| 275 |
+
elif act_type == "gelu_tanh":
|
| 276 |
+
return lambda: nn.GELU(approximate="tanh")
|
| 277 |
+
elif act_type == "relu":
|
| 278 |
+
return nn.ReLU
|
| 279 |
+
elif act_type == "silu":
|
| 280 |
+
return nn.SiLU
|
| 281 |
+
else:
|
| 282 |
+
raise ValueError(f"Unknown activation type: {act_type}")
|
| 283 |
+
|
| 284 |
+
class IndividualTokenRefinerBlock(torch.nn.Module):
|
| 285 |
+
def __init__(
|
| 286 |
+
self,
|
| 287 |
+
hidden_size,
|
| 288 |
+
heads_num,
|
| 289 |
+
mlp_width_ratio: str = 4.0,
|
| 290 |
+
mlp_drop_rate: float = 0.0,
|
| 291 |
+
act_type: str = "silu",
|
| 292 |
+
qk_norm: bool = False,
|
| 293 |
+
qk_norm_type: str = "layer",
|
| 294 |
+
qkv_bias: bool = True,
|
| 295 |
+
need_CA: bool = False,
|
| 296 |
+
dtype: Optional[torch.dtype] = None,
|
| 297 |
+
device: Optional[torch.device] = None,
|
| 298 |
+
):
|
| 299 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 300 |
+
super().__init__()
|
| 301 |
+
self.need_CA = need_CA
|
| 302 |
+
self.heads_num = heads_num
|
| 303 |
+
head_dim = hidden_size // heads_num
|
| 304 |
+
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
|
| 305 |
+
|
| 306 |
+
self.norm1 = nn.LayerNorm(
|
| 307 |
+
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
| 308 |
+
)
|
| 309 |
+
self.self_attn_qkv = nn.Linear(
|
| 310 |
+
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
|
| 311 |
+
)
|
| 312 |
+
qk_norm_layer = get_norm_layer(qk_norm_type)
|
| 313 |
+
self.self_attn_q_norm = (
|
| 314 |
+
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
| 315 |
+
if qk_norm
|
| 316 |
+
else nn.Identity()
|
| 317 |
+
)
|
| 318 |
+
self.self_attn_k_norm = (
|
| 319 |
+
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
| 320 |
+
if qk_norm
|
| 321 |
+
else nn.Identity()
|
| 322 |
+
)
|
| 323 |
+
self.self_attn_proj = nn.Linear(
|
| 324 |
+
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
| 325 |
+
)
|
| 326 |
+
|
| 327 |
+
self.norm2 = nn.LayerNorm(
|
| 328 |
+
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
| 329 |
+
)
|
| 330 |
+
act_layer = get_activation_layer(act_type)
|
| 331 |
+
self.mlp = MLP(
|
| 332 |
+
in_channels=hidden_size,
|
| 333 |
+
hidden_channels=mlp_hidden_dim,
|
| 334 |
+
act_layer=act_layer,
|
| 335 |
+
drop=mlp_drop_rate,
|
| 336 |
+
**factory_kwargs,
|
| 337 |
+
)
|
| 338 |
+
|
| 339 |
+
self.adaLN_modulation = nn.Sequential(
|
| 340 |
+
act_layer(),
|
| 341 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
|
| 342 |
+
)
|
| 343 |
+
|
| 344 |
+
if self.need_CA:
|
| 345 |
+
self.cross_attnblock=CrossAttnBlock(hidden_size=hidden_size,
|
| 346 |
+
heads_num=heads_num,
|
| 347 |
+
mlp_width_ratio=mlp_width_ratio,
|
| 348 |
+
mlp_drop_rate=mlp_drop_rate,
|
| 349 |
+
act_type=act_type,
|
| 350 |
+
qk_norm=qk_norm,
|
| 351 |
+
qk_norm_type=qk_norm_type,
|
| 352 |
+
qkv_bias=qkv_bias,
|
| 353 |
+
**factory_kwargs,)
|
| 354 |
+
# Zero-initialize the modulation
|
| 355 |
+
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
| 356 |
+
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
| 357 |
+
|
| 358 |
+
def forward(
|
| 359 |
+
self,
|
| 360 |
+
x: torch.Tensor,
|
| 361 |
+
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
|
| 362 |
+
attn_mask: torch.Tensor = None,
|
| 363 |
+
y: torch.Tensor = None,
|
| 364 |
+
):
|
| 365 |
+
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 366 |
+
|
| 367 |
+
norm_x = self.norm1(x)
|
| 368 |
+
qkv = self.self_attn_qkv(norm_x)
|
| 369 |
+
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
| 370 |
+
# Apply QK-Norm if needed
|
| 371 |
+
q = self.self_attn_q_norm(q).to(v)
|
| 372 |
+
k = self.self_attn_k_norm(k).to(v)
|
| 373 |
+
|
| 374 |
+
# Self-Attention
|
| 375 |
+
attn = attention(q, k, v, mode="torch", attn_mask=attn_mask)
|
| 376 |
+
|
| 377 |
+
x = x + apply_gate(self.self_attn_proj(attn), gate_msa)
|
| 378 |
+
|
| 379 |
+
if self.need_CA:
|
| 380 |
+
x = self.cross_attnblock(x, c, attn_mask, y)
|
| 381 |
+
|
| 382 |
+
# FFN Layer
|
| 383 |
+
x = x + apply_gate(self.mlp(self.norm2(x)), gate_mlp)
|
| 384 |
+
|
| 385 |
+
return x
|
| 386 |
+
|
| 387 |
+
|
| 388 |
+
|
| 389 |
+
|
| 390 |
+
class CrossAttnBlock(torch.nn.Module):
|
| 391 |
+
def __init__(
|
| 392 |
+
self,
|
| 393 |
+
hidden_size,
|
| 394 |
+
heads_num,
|
| 395 |
+
mlp_width_ratio: str = 4.0,
|
| 396 |
+
mlp_drop_rate: float = 0.0,
|
| 397 |
+
act_type: str = "silu",
|
| 398 |
+
qk_norm: bool = False,
|
| 399 |
+
qk_norm_type: str = "layer",
|
| 400 |
+
qkv_bias: bool = True,
|
| 401 |
+
dtype: Optional[torch.dtype] = None,
|
| 402 |
+
device: Optional[torch.device] = None,
|
| 403 |
+
):
|
| 404 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 405 |
+
super().__init__()
|
| 406 |
+
self.heads_num = heads_num
|
| 407 |
+
head_dim = hidden_size // heads_num
|
| 408 |
+
|
| 409 |
+
self.norm1 = nn.LayerNorm(
|
| 410 |
+
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
| 411 |
+
)
|
| 412 |
+
self.norm1_2 = nn.LayerNorm(
|
| 413 |
+
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
| 414 |
+
)
|
| 415 |
+
self.self_attn_q = nn.Linear(
|
| 416 |
+
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
| 417 |
+
)
|
| 418 |
+
self.self_attn_kv = nn.Linear(
|
| 419 |
+
hidden_size, hidden_size*2, bias=qkv_bias, **factory_kwargs
|
| 420 |
+
)
|
| 421 |
+
qk_norm_layer = get_norm_layer(qk_norm_type)
|
| 422 |
+
self.self_attn_q_norm = (
|
| 423 |
+
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
| 424 |
+
if qk_norm
|
| 425 |
+
else nn.Identity()
|
| 426 |
+
)
|
| 427 |
+
self.self_attn_k_norm = (
|
| 428 |
+
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
| 429 |
+
if qk_norm
|
| 430 |
+
else nn.Identity()
|
| 431 |
+
)
|
| 432 |
+
self.self_attn_proj = nn.Linear(
|
| 433 |
+
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
| 434 |
+
)
|
| 435 |
+
|
| 436 |
+
self.norm2 = nn.LayerNorm(
|
| 437 |
+
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
| 438 |
+
)
|
| 439 |
+
act_layer = get_activation_layer(act_type)
|
| 440 |
+
|
| 441 |
+
self.adaLN_modulation = nn.Sequential(
|
| 442 |
+
act_layer(),
|
| 443 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
|
| 444 |
+
)
|
| 445 |
+
# Zero-initialize the modulation
|
| 446 |
+
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
| 447 |
+
nn.init.zeros_(self.adaLN_modulation[1].bias)
|
| 448 |
+
|
| 449 |
+
def forward(
|
| 450 |
+
self,
|
| 451 |
+
x: torch.Tensor,
|
| 452 |
+
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
|
| 453 |
+
attn_mask: torch.Tensor = None,
|
| 454 |
+
y: torch.Tensor=None,
|
| 455 |
+
|
| 456 |
+
):
|
| 457 |
+
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 458 |
+
|
| 459 |
+
norm_x = self.norm1(x)
|
| 460 |
+
norm_y = self.norm1_2(y)
|
| 461 |
+
q = self.self_attn_q(norm_x)
|
| 462 |
+
q = rearrange(q, "B L (H D) -> B L H D", H=self.heads_num)
|
| 463 |
+
kv = self.self_attn_kv(norm_y)
|
| 464 |
+
k, v = rearrange(kv, "B L (K H D) -> K B L H D", K=2, H=self.heads_num)
|
| 465 |
+
# Apply QK-Norm if needed
|
| 466 |
+
q = self.self_attn_q_norm(q).to(v)
|
| 467 |
+
k = self.self_attn_k_norm(k).to(v)
|
| 468 |
+
|
| 469 |
+
# Self-Attention
|
| 470 |
+
attn = attention(q, k, v, mode="torch", attn_mask=attn_mask)
|
| 471 |
+
|
| 472 |
+
x = x + apply_gate(self.self_attn_proj(attn), gate_msa)
|
| 473 |
+
|
| 474 |
+
return x
|
| 475 |
+
|
| 476 |
+
|
| 477 |
+
|
| 478 |
+
class IndividualTokenRefiner(torch.nn.Module):
|
| 479 |
+
def __init__(
|
| 480 |
+
self,
|
| 481 |
+
hidden_size,
|
| 482 |
+
heads_num,
|
| 483 |
+
depth,
|
| 484 |
+
mlp_width_ratio: float = 4.0,
|
| 485 |
+
mlp_drop_rate: float = 0.0,
|
| 486 |
+
act_type: str = "silu",
|
| 487 |
+
qk_norm: bool = False,
|
| 488 |
+
qk_norm_type: str = "layer",
|
| 489 |
+
qkv_bias: bool = True,
|
| 490 |
+
need_CA:bool=False,
|
| 491 |
+
dtype: Optional[torch.dtype] = None,
|
| 492 |
+
device: Optional[torch.device] = None,
|
| 493 |
+
):
|
| 494 |
+
|
| 495 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 496 |
+
super().__init__()
|
| 497 |
+
self.need_CA = need_CA
|
| 498 |
+
self.blocks = nn.ModuleList(
|
| 499 |
+
[
|
| 500 |
+
IndividualTokenRefinerBlock(
|
| 501 |
+
hidden_size=hidden_size,
|
| 502 |
+
heads_num=heads_num,
|
| 503 |
+
mlp_width_ratio=mlp_width_ratio,
|
| 504 |
+
mlp_drop_rate=mlp_drop_rate,
|
| 505 |
+
act_type=act_type,
|
| 506 |
+
qk_norm=qk_norm,
|
| 507 |
+
qk_norm_type=qk_norm_type,
|
| 508 |
+
qkv_bias=qkv_bias,
|
| 509 |
+
need_CA=self.need_CA,
|
| 510 |
+
**factory_kwargs,
|
| 511 |
+
)
|
| 512 |
+
for _ in range(depth)
|
| 513 |
+
]
|
| 514 |
+
)
|
| 515 |
+
|
| 516 |
+
|
| 517 |
+
def forward(
|
| 518 |
+
self,
|
| 519 |
+
x: torch.Tensor,
|
| 520 |
+
c: torch.LongTensor,
|
| 521 |
+
mask: Optional[torch.Tensor] = None,
|
| 522 |
+
y:torch.Tensor=None,
|
| 523 |
+
):
|
| 524 |
+
self_attn_mask = None
|
| 525 |
+
if mask is not None:
|
| 526 |
+
batch_size = mask.shape[0]
|
| 527 |
+
seq_len = mask.shape[1]
|
| 528 |
+
mask = mask.to(x.device)
|
| 529 |
+
# batch_size x 1 x seq_len x seq_len
|
| 530 |
+
self_attn_mask_1 = mask.view(batch_size, 1, 1, seq_len).repeat(
|
| 531 |
+
1, 1, seq_len, 1
|
| 532 |
+
)
|
| 533 |
+
# batch_size x 1 x seq_len x seq_len
|
| 534 |
+
self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
|
| 535 |
+
# batch_size x 1 x seq_len x seq_len, 1 for broadcasting of heads_num
|
| 536 |
+
self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
|
| 537 |
+
# avoids self-attention weight being NaN for padding tokens
|
| 538 |
+
self_attn_mask[:, :, :, 0] = True
|
| 539 |
+
|
| 540 |
+
|
| 541 |
+
for block in self.blocks:
|
| 542 |
+
x = block(x, c, self_attn_mask,y)
|
| 543 |
+
|
| 544 |
+
return x
|
| 545 |
+
|
| 546 |
+
|
| 547 |
+
class SingleTokenRefiner(torch.nn.Module):
|
| 548 |
+
"""
|
| 549 |
+
A single token refiner block for llm text embedding refine.
|
| 550 |
+
"""
|
| 551 |
+
def __init__(
|
| 552 |
+
self,
|
| 553 |
+
in_channels,
|
| 554 |
+
hidden_size,
|
| 555 |
+
heads_num,
|
| 556 |
+
depth,
|
| 557 |
+
mlp_width_ratio: float = 4.0,
|
| 558 |
+
mlp_drop_rate: float = 0.0,
|
| 559 |
+
act_type: str = "silu",
|
| 560 |
+
qk_norm: bool = False,
|
| 561 |
+
qk_norm_type: str = "layer",
|
| 562 |
+
qkv_bias: bool = True,
|
| 563 |
+
need_CA:bool=False,
|
| 564 |
+
attn_mode: str = "torch",
|
| 565 |
+
dtype: Optional[torch.dtype] = None,
|
| 566 |
+
device: Optional[torch.device] = None,
|
| 567 |
+
):
|
| 568 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 569 |
+
super().__init__()
|
| 570 |
+
self.attn_mode = attn_mode
|
| 571 |
+
self.need_CA = need_CA
|
| 572 |
+
assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."
|
| 573 |
+
|
| 574 |
+
self.input_embedder = nn.Linear(
|
| 575 |
+
in_channels, hidden_size, bias=True, **factory_kwargs
|
| 576 |
+
)
|
| 577 |
+
if self.need_CA:
|
| 578 |
+
self.input_embedder_CA = nn.Linear(
|
| 579 |
+
in_channels, hidden_size, bias=True, **factory_kwargs
|
| 580 |
+
)
|
| 581 |
+
|
| 582 |
+
act_layer = get_activation_layer(act_type)
|
| 583 |
+
# Build timestep embedding layer
|
| 584 |
+
self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)
|
| 585 |
+
# Build context embedding layer
|
| 586 |
+
self.c_embedder = TextProjection(
|
| 587 |
+
in_channels, hidden_size, act_layer, **factory_kwargs
|
| 588 |
+
)
|
| 589 |
+
|
| 590 |
+
self.individual_token_refiner = IndividualTokenRefiner(
|
| 591 |
+
hidden_size=hidden_size,
|
| 592 |
+
heads_num=heads_num,
|
| 593 |
+
depth=depth,
|
| 594 |
+
mlp_width_ratio=mlp_width_ratio,
|
| 595 |
+
mlp_drop_rate=mlp_drop_rate,
|
| 596 |
+
act_type=act_type,
|
| 597 |
+
qk_norm=qk_norm,
|
| 598 |
+
qk_norm_type=qk_norm_type,
|
| 599 |
+
qkv_bias=qkv_bias,
|
| 600 |
+
need_CA=need_CA,
|
| 601 |
+
**factory_kwargs,
|
| 602 |
+
)
|
| 603 |
+
|
| 604 |
+
def forward(
|
| 605 |
+
self,
|
| 606 |
+
x: torch.Tensor,
|
| 607 |
+
t: torch.LongTensor,
|
| 608 |
+
mask: Optional[torch.LongTensor] = None,
|
| 609 |
+
y: torch.LongTensor=None,
|
| 610 |
+
):
|
| 611 |
+
timestep_aware_representations = self.t_embedder(t)
|
| 612 |
+
|
| 613 |
+
if mask is None:
|
| 614 |
+
context_aware_representations = x.mean(dim=1)
|
| 615 |
+
else:
|
| 616 |
+
mask_float = mask.unsqueeze(-1) # [b, s1, 1]
|
| 617 |
+
context_aware_representations = (x * mask_float).sum(
|
| 618 |
+
dim=1
|
| 619 |
+
) / mask_float.sum(dim=1)
|
| 620 |
+
context_aware_representations = self.c_embedder(context_aware_representations)
|
| 621 |
+
c = timestep_aware_representations + context_aware_representations
|
| 622 |
+
|
| 623 |
+
x = self.input_embedder(x)
|
| 624 |
+
if self.need_CA:
|
| 625 |
+
y = self.input_embedder_CA(y)
|
| 626 |
+
x = self.individual_token_refiner(x, c, mask, y)
|
| 627 |
+
else:
|
| 628 |
+
x = self.individual_token_refiner(x, c, mask)
|
| 629 |
+
|
| 630 |
+
return x
|
| 631 |
+
|
| 632 |
+
|
| 633 |
+
class Qwen2Connector(torch.nn.Module):
|
| 634 |
+
def __init__(
|
| 635 |
+
self,
|
| 636 |
+
# biclip_dim=1024,
|
| 637 |
+
in_channels=3584,
|
| 638 |
+
hidden_size=4096,
|
| 639 |
+
heads_num=32,
|
| 640 |
+
depth=2,
|
| 641 |
+
need_CA=False,
|
| 642 |
+
device=None,
|
| 643 |
+
dtype=torch.bfloat16,
|
| 644 |
+
):
|
| 645 |
+
super().__init__()
|
| 646 |
+
factory_kwargs = {"device": device, "dtype":dtype}
|
| 647 |
+
|
| 648 |
+
self.S =SingleTokenRefiner(in_channels=in_channels,hidden_size=hidden_size,heads_num=heads_num,depth=depth,need_CA=need_CA,**factory_kwargs)
|
| 649 |
+
self.global_proj_out=nn.Linear(in_channels,768)
|
| 650 |
+
|
| 651 |
+
self.scale_factor = nn.Parameter(torch.zeros(1))
|
| 652 |
+
with torch.no_grad():
|
| 653 |
+
self.scale_factor.data += -(1 - 0.09)
|
| 654 |
+
|
| 655 |
+
def forward(self, x,t,mask):
|
| 656 |
+
mask_float = mask.unsqueeze(-1) # [b, s1, 1]
|
| 657 |
+
x_mean = (x * mask_float).sum(
|
| 658 |
+
dim=1
|
| 659 |
+
) / mask_float.sum(dim=1) * (1 + self.scale_factor.to(dtype=x.dtype, device=x.device))
|
| 660 |
+
|
| 661 |
+
global_out=self.global_proj_out(x_mean)
|
| 662 |
+
encoder_hidden_states = self.S(x,t,mask)
|
| 663 |
+
return encoder_hidden_states,global_out
|
| 664 |
+
|
| 665 |
+
@staticmethod
|
| 666 |
+
def state_dict_converter():
|
| 667 |
+
return Qwen2ConnectorStateDictConverter()
|
| 668 |
+
|
| 669 |
+
|
| 670 |
+
class Qwen2ConnectorStateDictConverter:
|
| 671 |
+
def __init__(self):
|
| 672 |
+
pass
|
| 673 |
+
|
| 674 |
+
def from_diffusers(self, state_dict):
|
| 675 |
+
return state_dict
|
| 676 |
+
|
| 677 |
+
def from_civitai(self, state_dict):
|
| 678 |
+
state_dict_ = {}
|
| 679 |
+
for name, param in state_dict.items():
|
| 680 |
+
if name.startswith("connector."):
|
| 681 |
+
name_ = name[len("connector."):]
|
| 682 |
+
state_dict_[name_] = param
|
| 683 |
+
return state_dict_
|
diffsynth/models/stepvideo_dit.py
ADDED
|
@@ -0,0 +1,940 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2025 StepFun Inc. All Rights Reserved.
|
| 2 |
+
#
|
| 3 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 4 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 5 |
+
# in the Software without restriction, including without limitation the rights
|
| 6 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 7 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 8 |
+
# furnished to do so, subject to the following conditions:
|
| 9 |
+
#
|
| 10 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 11 |
+
# copies or substantial portions of the Software.
|
| 12 |
+
# ==============================================================================
|
| 13 |
+
from typing import Dict, Optional, Tuple, Union, List
|
| 14 |
+
import torch, math
|
| 15 |
+
from torch import nn
|
| 16 |
+
from einops import rearrange, repeat
|
| 17 |
+
from tqdm import tqdm
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class RMSNorm(nn.Module):
|
| 21 |
+
def __init__(
|
| 22 |
+
self,
|
| 23 |
+
dim: int,
|
| 24 |
+
elementwise_affine=True,
|
| 25 |
+
eps: float = 1e-6,
|
| 26 |
+
device=None,
|
| 27 |
+
dtype=None,
|
| 28 |
+
):
|
| 29 |
+
"""
|
| 30 |
+
Initialize the RMSNorm normalization layer.
|
| 31 |
+
|
| 32 |
+
Args:
|
| 33 |
+
dim (int): The dimension of the input tensor.
|
| 34 |
+
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
|
| 35 |
+
|
| 36 |
+
Attributes:
|
| 37 |
+
eps (float): A small value added to the denominator for numerical stability.
|
| 38 |
+
weight (nn.Parameter): Learnable scaling parameter.
|
| 39 |
+
|
| 40 |
+
"""
|
| 41 |
+
factory_kwargs = {"device": device, "dtype": dtype}
|
| 42 |
+
super().__init__()
|
| 43 |
+
self.eps = eps
|
| 44 |
+
if elementwise_affine:
|
| 45 |
+
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
|
| 46 |
+
|
| 47 |
+
def _norm(self, x):
|
| 48 |
+
"""
|
| 49 |
+
Apply the RMSNorm normalization to the input tensor.
|
| 50 |
+
|
| 51 |
+
Args:
|
| 52 |
+
x (torch.Tensor): The input tensor.
|
| 53 |
+
|
| 54 |
+
Returns:
|
| 55 |
+
torch.Tensor: The normalized tensor.
|
| 56 |
+
|
| 57 |
+
"""
|
| 58 |
+
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
"""
|
| 62 |
+
Forward pass through the RMSNorm layer.
|
| 63 |
+
|
| 64 |
+
Args:
|
| 65 |
+
x (torch.Tensor): The input tensor.
|
| 66 |
+
|
| 67 |
+
Returns:
|
| 68 |
+
torch.Tensor: The output tensor after applying RMSNorm.
|
| 69 |
+
|
| 70 |
+
"""
|
| 71 |
+
output = self._norm(x.float()).type_as(x)
|
| 72 |
+
if hasattr(self, "weight"):
|
| 73 |
+
output = output * self.weight
|
| 74 |
+
return output
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
ACTIVATION_FUNCTIONS = {
|
| 78 |
+
"swish": nn.SiLU(),
|
| 79 |
+
"silu": nn.SiLU(),
|
| 80 |
+
"mish": nn.Mish(),
|
| 81 |
+
"gelu": nn.GELU(),
|
| 82 |
+
"relu": nn.ReLU(),
|
| 83 |
+
}
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def get_activation(act_fn: str) -> nn.Module:
|
| 87 |
+
"""Helper function to get activation function from string.
|
| 88 |
+
|
| 89 |
+
Args:
|
| 90 |
+
act_fn (str): Name of activation function.
|
| 91 |
+
|
| 92 |
+
Returns:
|
| 93 |
+
nn.Module: Activation function.
|
| 94 |
+
"""
|
| 95 |
+
|
| 96 |
+
act_fn = act_fn.lower()
|
| 97 |
+
if act_fn in ACTIVATION_FUNCTIONS:
|
| 98 |
+
return ACTIVATION_FUNCTIONS[act_fn]
|
| 99 |
+
else:
|
| 100 |
+
raise ValueError(f"Unsupported activation function: {act_fn}")
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def get_timestep_embedding(
|
| 104 |
+
timesteps: torch.Tensor,
|
| 105 |
+
embedding_dim: int,
|
| 106 |
+
flip_sin_to_cos: bool = False,
|
| 107 |
+
downscale_freq_shift: float = 1,
|
| 108 |
+
scale: float = 1,
|
| 109 |
+
max_period: int = 10000,
|
| 110 |
+
):
|
| 111 |
+
"""
|
| 112 |
+
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
|
| 113 |
+
|
| 114 |
+
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
| 115 |
+
These may be fractional.
|
| 116 |
+
:param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the
|
| 117 |
+
embeddings. :return: an [N x dim] Tensor of positional embeddings.
|
| 118 |
+
"""
|
| 119 |
+
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
|
| 120 |
+
|
| 121 |
+
half_dim = embedding_dim // 2
|
| 122 |
+
exponent = -math.log(max_period) * torch.arange(
|
| 123 |
+
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
|
| 124 |
+
)
|
| 125 |
+
exponent = exponent / (half_dim - downscale_freq_shift)
|
| 126 |
+
|
| 127 |
+
emb = torch.exp(exponent)
|
| 128 |
+
emb = timesteps[:, None].float() * emb[None, :]
|
| 129 |
+
|
| 130 |
+
# scale embeddings
|
| 131 |
+
emb = scale * emb
|
| 132 |
+
|
| 133 |
+
# concat sine and cosine embeddings
|
| 134 |
+
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
|
| 135 |
+
|
| 136 |
+
# flip sine and cosine embeddings
|
| 137 |
+
if flip_sin_to_cos:
|
| 138 |
+
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
| 139 |
+
|
| 140 |
+
# zero pad
|
| 141 |
+
if embedding_dim % 2 == 1:
|
| 142 |
+
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
|
| 143 |
+
return emb
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
class Timesteps(nn.Module):
|
| 147 |
+
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float):
|
| 148 |
+
super().__init__()
|
| 149 |
+
self.num_channels = num_channels
|
| 150 |
+
self.flip_sin_to_cos = flip_sin_to_cos
|
| 151 |
+
self.downscale_freq_shift = downscale_freq_shift
|
| 152 |
+
|
| 153 |
+
def forward(self, timesteps):
|
| 154 |
+
t_emb = get_timestep_embedding(
|
| 155 |
+
timesteps,
|
| 156 |
+
self.num_channels,
|
| 157 |
+
flip_sin_to_cos=self.flip_sin_to_cos,
|
| 158 |
+
downscale_freq_shift=self.downscale_freq_shift,
|
| 159 |
+
)
|
| 160 |
+
return t_emb
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
class TimestepEmbedding(nn.Module):
|
| 164 |
+
def __init__(
|
| 165 |
+
self,
|
| 166 |
+
in_channels: int,
|
| 167 |
+
time_embed_dim: int,
|
| 168 |
+
act_fn: str = "silu",
|
| 169 |
+
out_dim: int = None,
|
| 170 |
+
post_act_fn: Optional[str] = None,
|
| 171 |
+
cond_proj_dim=None,
|
| 172 |
+
sample_proj_bias=True
|
| 173 |
+
):
|
| 174 |
+
super().__init__()
|
| 175 |
+
linear_cls = nn.Linear
|
| 176 |
+
|
| 177 |
+
self.linear_1 = linear_cls(
|
| 178 |
+
in_channels,
|
| 179 |
+
time_embed_dim,
|
| 180 |
+
bias=sample_proj_bias,
|
| 181 |
+
)
|
| 182 |
+
|
| 183 |
+
if cond_proj_dim is not None:
|
| 184 |
+
self.cond_proj = linear_cls(
|
| 185 |
+
cond_proj_dim,
|
| 186 |
+
in_channels,
|
| 187 |
+
bias=False,
|
| 188 |
+
)
|
| 189 |
+
else:
|
| 190 |
+
self.cond_proj = None
|
| 191 |
+
|
| 192 |
+
self.act = get_activation(act_fn)
|
| 193 |
+
|
| 194 |
+
if out_dim is not None:
|
| 195 |
+
time_embed_dim_out = out_dim
|
| 196 |
+
else:
|
| 197 |
+
time_embed_dim_out = time_embed_dim
|
| 198 |
+
|
| 199 |
+
self.linear_2 = linear_cls(
|
| 200 |
+
time_embed_dim,
|
| 201 |
+
time_embed_dim_out,
|
| 202 |
+
bias=sample_proj_bias,
|
| 203 |
+
)
|
| 204 |
+
|
| 205 |
+
if post_act_fn is None:
|
| 206 |
+
self.post_act = None
|
| 207 |
+
else:
|
| 208 |
+
self.post_act = get_activation(post_act_fn)
|
| 209 |
+
|
| 210 |
+
def forward(self, sample, condition=None):
|
| 211 |
+
if condition is not None:
|
| 212 |
+
sample = sample + self.cond_proj(condition)
|
| 213 |
+
sample = self.linear_1(sample)
|
| 214 |
+
|
| 215 |
+
if self.act is not None:
|
| 216 |
+
sample = self.act(sample)
|
| 217 |
+
|
| 218 |
+
sample = self.linear_2(sample)
|
| 219 |
+
|
| 220 |
+
if self.post_act is not None:
|
| 221 |
+
sample = self.post_act(sample)
|
| 222 |
+
return sample
|
| 223 |
+
|
| 224 |
+
|
| 225 |
+
class PixArtAlphaCombinedTimestepSizeEmbeddings(nn.Module):
|
| 226 |
+
def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool = False):
|
| 227 |
+
super().__init__()
|
| 228 |
+
|
| 229 |
+
self.outdim = size_emb_dim
|
| 230 |
+
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
|
| 231 |
+
self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
|
| 232 |
+
|
| 233 |
+
self.use_additional_conditions = use_additional_conditions
|
| 234 |
+
if self.use_additional_conditions:
|
| 235 |
+
self.additional_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
|
| 236 |
+
self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim)
|
| 237 |
+
self.nframe_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
|
| 238 |
+
self.fps_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)
|
| 239 |
+
|
| 240 |
+
def forward(self, timestep, resolution=None, nframe=None, fps=None):
|
| 241 |
+
hidden_dtype = timestep.dtype
|
| 242 |
+
|
| 243 |
+
timesteps_proj = self.time_proj(timestep)
|
| 244 |
+
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype)) # (N, D)
|
| 245 |
+
|
| 246 |
+
if self.use_additional_conditions:
|
| 247 |
+
batch_size = timestep.shape[0]
|
| 248 |
+
resolution_emb = self.additional_condition_proj(resolution.flatten()).to(hidden_dtype)
|
| 249 |
+
resolution_emb = self.resolution_embedder(resolution_emb).reshape(batch_size, -1)
|
| 250 |
+
nframe_emb = self.additional_condition_proj(nframe.flatten()).to(hidden_dtype)
|
| 251 |
+
nframe_emb = self.nframe_embedder(nframe_emb).reshape(batch_size, -1)
|
| 252 |
+
conditioning = timesteps_emb + resolution_emb + nframe_emb
|
| 253 |
+
|
| 254 |
+
if fps is not None:
|
| 255 |
+
fps_emb = self.additional_condition_proj(fps.flatten()).to(hidden_dtype)
|
| 256 |
+
fps_emb = self.fps_embedder(fps_emb).reshape(batch_size, -1)
|
| 257 |
+
conditioning = conditioning + fps_emb
|
| 258 |
+
else:
|
| 259 |
+
conditioning = timesteps_emb
|
| 260 |
+
|
| 261 |
+
return conditioning
|
| 262 |
+
|
| 263 |
+
|
| 264 |
+
class AdaLayerNormSingle(nn.Module):
|
| 265 |
+
r"""
|
| 266 |
+
Norm layer adaptive layer norm single (adaLN-single).
|
| 267 |
+
|
| 268 |
+
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
|
| 269 |
+
|
| 270 |
+
Parameters:
|
| 271 |
+
embedding_dim (`int`): The size of each embedding vector.
|
| 272 |
+
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
|
| 273 |
+
"""
|
| 274 |
+
def __init__(self, embedding_dim: int, use_additional_conditions: bool = False, time_step_rescale=1000):
|
| 275 |
+
super().__init__()
|
| 276 |
+
|
| 277 |
+
self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings(
|
| 278 |
+
embedding_dim, size_emb_dim=embedding_dim // 2, use_additional_conditions=use_additional_conditions
|
| 279 |
+
)
|
| 280 |
+
|
| 281 |
+
self.silu = nn.SiLU()
|
| 282 |
+
self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True)
|
| 283 |
+
|
| 284 |
+
self.time_step_rescale = time_step_rescale ## timestep usually in [0, 1], we rescale it to [0,1000] for stability
|
| 285 |
+
|
| 286 |
+
def forward(
|
| 287 |
+
self,
|
| 288 |
+
timestep: torch.Tensor,
|
| 289 |
+
added_cond_kwargs: Dict[str, torch.Tensor] = None,
|
| 290 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 291 |
+
embedded_timestep = self.emb(timestep*self.time_step_rescale, **added_cond_kwargs)
|
| 292 |
+
|
| 293 |
+
out = self.linear(self.silu(embedded_timestep))
|
| 294 |
+
|
| 295 |
+
return out, embedded_timestep
|
| 296 |
+
|
| 297 |
+
|
| 298 |
+
class PixArtAlphaTextProjection(nn.Module):
|
| 299 |
+
"""
|
| 300 |
+
Projects caption embeddings. Also handles dropout for classifier-free guidance.
|
| 301 |
+
|
| 302 |
+
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
|
| 303 |
+
"""
|
| 304 |
+
|
| 305 |
+
def __init__(self, in_features, hidden_size):
|
| 306 |
+
super().__init__()
|
| 307 |
+
self.linear_1 = nn.Linear(
|
| 308 |
+
in_features,
|
| 309 |
+
hidden_size,
|
| 310 |
+
bias=True,
|
| 311 |
+
)
|
| 312 |
+
self.act_1 = nn.GELU(approximate="tanh")
|
| 313 |
+
self.linear_2 = nn.Linear(
|
| 314 |
+
hidden_size,
|
| 315 |
+
hidden_size,
|
| 316 |
+
bias=True,
|
| 317 |
+
)
|
| 318 |
+
|
| 319 |
+
def forward(self, caption):
|
| 320 |
+
hidden_states = self.linear_1(caption)
|
| 321 |
+
hidden_states = self.act_1(hidden_states)
|
| 322 |
+
hidden_states = self.linear_2(hidden_states)
|
| 323 |
+
return hidden_states
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
class Attention(nn.Module):
|
| 327 |
+
def __init__(self):
|
| 328 |
+
super().__init__()
|
| 329 |
+
|
| 330 |
+
def attn_processor(self, attn_type):
|
| 331 |
+
if attn_type == 'torch':
|
| 332 |
+
return self.torch_attn_func
|
| 333 |
+
elif attn_type == 'parallel':
|
| 334 |
+
return self.parallel_attn_func
|
| 335 |
+
else:
|
| 336 |
+
raise Exception('Not supported attention type...')
|
| 337 |
+
|
| 338 |
+
def torch_attn_func(
|
| 339 |
+
self,
|
| 340 |
+
q,
|
| 341 |
+
k,
|
| 342 |
+
v,
|
| 343 |
+
attn_mask=None,
|
| 344 |
+
causal=False,
|
| 345 |
+
drop_rate=0.0,
|
| 346 |
+
**kwargs
|
| 347 |
+
):
|
| 348 |
+
|
| 349 |
+
if attn_mask is not None and attn_mask.dtype != torch.bool:
|
| 350 |
+
attn_mask = attn_mask.to(q.dtype)
|
| 351 |
+
|
| 352 |
+
if attn_mask is not None and attn_mask.ndim == 3: ## no head
|
| 353 |
+
n_heads = q.shape[2]
|
| 354 |
+
attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1)
|
| 355 |
+
|
| 356 |
+
q, k, v = map(lambda x: rearrange(x, 'b s h d -> b h s d'), (q, k, v))
|
| 357 |
+
if attn_mask is not None:
|
| 358 |
+
attn_mask = attn_mask.to(q.device)
|
| 359 |
+
x = torch.nn.functional.scaled_dot_product_attention(
|
| 360 |
+
q, k, v, attn_mask=attn_mask, dropout_p=drop_rate, is_causal=causal
|
| 361 |
+
)
|
| 362 |
+
x = rearrange(x, 'b h s d -> b s h d')
|
| 363 |
+
return x
|
| 364 |
+
|
| 365 |
+
|
| 366 |
+
class RoPE1D:
|
| 367 |
+
def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):
|
| 368 |
+
self.base = freq
|
| 369 |
+
self.F0 = F0
|
| 370 |
+
self.scaling_factor = scaling_factor
|
| 371 |
+
self.cache = {}
|
| 372 |
+
|
| 373 |
+
def get_cos_sin(self, D, seq_len, device, dtype):
|
| 374 |
+
if (D, seq_len, device, dtype) not in self.cache:
|
| 375 |
+
inv_freq = 1.0 / (self.base ** (torch.arange(0, D, 2).float().to(device) / D))
|
| 376 |
+
t = torch.arange(seq_len, device=device, dtype=inv_freq.dtype)
|
| 377 |
+
freqs = torch.einsum("i,j->ij", t, inv_freq).to(dtype)
|
| 378 |
+
freqs = torch.cat((freqs, freqs), dim=-1)
|
| 379 |
+
cos = freqs.cos() # (Seq, Dim)
|
| 380 |
+
sin = freqs.sin()
|
| 381 |
+
self.cache[D, seq_len, device, dtype] = (cos, sin)
|
| 382 |
+
return self.cache[D, seq_len, device, dtype]
|
| 383 |
+
|
| 384 |
+
@staticmethod
|
| 385 |
+
def rotate_half(x):
|
| 386 |
+
x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2:]
|
| 387 |
+
return torch.cat((-x2, x1), dim=-1)
|
| 388 |
+
|
| 389 |
+
def apply_rope1d(self, tokens, pos1d, cos, sin):
|
| 390 |
+
assert pos1d.ndim == 2
|
| 391 |
+
cos = torch.nn.functional.embedding(pos1d, cos)[:, :, None, :]
|
| 392 |
+
sin = torch.nn.functional.embedding(pos1d, sin)[:, :, None, :]
|
| 393 |
+
return (tokens * cos) + (self.rotate_half(tokens) * sin)
|
| 394 |
+
|
| 395 |
+
def __call__(self, tokens, positions):
|
| 396 |
+
"""
|
| 397 |
+
input:
|
| 398 |
+
* tokens: batch_size x ntokens x nheads x dim
|
| 399 |
+
* positions: batch_size x ntokens (t position of each token)
|
| 400 |
+
output:
|
| 401 |
+
* tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)
|
| 402 |
+
"""
|
| 403 |
+
D = tokens.size(3)
|
| 404 |
+
assert positions.ndim == 2 # Batch, Seq
|
| 405 |
+
cos, sin = self.get_cos_sin(D, int(positions.max()) + 1, tokens.device, tokens.dtype)
|
| 406 |
+
tokens = self.apply_rope1d(tokens, positions, cos, sin)
|
| 407 |
+
return tokens
|
| 408 |
+
|
| 409 |
+
|
| 410 |
+
class RoPE3D(RoPE1D):
|
| 411 |
+
def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):
|
| 412 |
+
super(RoPE3D, self).__init__(freq, F0, scaling_factor)
|
| 413 |
+
self.position_cache = {}
|
| 414 |
+
|
| 415 |
+
def get_mesh_3d(self, rope_positions, bsz):
|
| 416 |
+
f, h, w = rope_positions
|
| 417 |
+
|
| 418 |
+
if f"{f}-{h}-{w}" not in self.position_cache:
|
| 419 |
+
x = torch.arange(f, device='cpu')
|
| 420 |
+
y = torch.arange(h, device='cpu')
|
| 421 |
+
z = torch.arange(w, device='cpu')
|
| 422 |
+
self.position_cache[f"{f}-{h}-{w}"] = torch.cartesian_prod(x, y, z).view(1, f*h*w, 3).expand(bsz, -1, 3)
|
| 423 |
+
return self.position_cache[f"{f}-{h}-{w}"]
|
| 424 |
+
|
| 425 |
+
def __call__(self, tokens, rope_positions, ch_split, parallel=False):
|
| 426 |
+
"""
|
| 427 |
+
input:
|
| 428 |
+
* tokens: batch_size x ntokens x nheads x dim
|
| 429 |
+
* rope_positions: list of (f, h, w)
|
| 430 |
+
output:
|
| 431 |
+
* tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)
|
| 432 |
+
"""
|
| 433 |
+
assert sum(ch_split) == tokens.size(-1);
|
| 434 |
+
|
| 435 |
+
mesh_grid = self.get_mesh_3d(rope_positions, bsz=tokens.shape[0])
|
| 436 |
+
out = []
|
| 437 |
+
for i, (D, x) in enumerate(zip(ch_split, torch.split(tokens, ch_split, dim=-1))):
|
| 438 |
+
cos, sin = self.get_cos_sin(D, int(mesh_grid.max()) + 1, tokens.device, tokens.dtype)
|
| 439 |
+
|
| 440 |
+
if parallel:
|
| 441 |
+
pass
|
| 442 |
+
else:
|
| 443 |
+
mesh = mesh_grid[:, :, i].clone()
|
| 444 |
+
x = self.apply_rope1d(x, mesh.to(tokens.device), cos, sin)
|
| 445 |
+
out.append(x)
|
| 446 |
+
|
| 447 |
+
tokens = torch.cat(out, dim=-1)
|
| 448 |
+
return tokens
|
| 449 |
+
|
| 450 |
+
|
| 451 |
+
class SelfAttention(Attention):
|
| 452 |
+
def __init__(self, hidden_dim, head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type='torch'):
|
| 453 |
+
super().__init__()
|
| 454 |
+
self.head_dim = head_dim
|
| 455 |
+
self.n_heads = hidden_dim // head_dim
|
| 456 |
+
|
| 457 |
+
self.wqkv = nn.Linear(hidden_dim, hidden_dim*3, bias=bias)
|
| 458 |
+
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
| 459 |
+
|
| 460 |
+
self.with_rope = with_rope
|
| 461 |
+
self.with_qk_norm = with_qk_norm
|
| 462 |
+
if self.with_qk_norm:
|
| 463 |
+
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
|
| 464 |
+
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
|
| 465 |
+
|
| 466 |
+
if self.with_rope:
|
| 467 |
+
self.rope_3d = RoPE3D(freq=1e4, F0=1.0, scaling_factor=1.0)
|
| 468 |
+
self.rope_ch_split = [64, 32, 32]
|
| 469 |
+
|
| 470 |
+
self.core_attention = self.attn_processor(attn_type=attn_type)
|
| 471 |
+
self.parallel = attn_type=='parallel'
|
| 472 |
+
|
| 473 |
+
def apply_rope3d(self, x, fhw_positions, rope_ch_split, parallel=True):
|
| 474 |
+
x = self.rope_3d(x, fhw_positions, rope_ch_split, parallel)
|
| 475 |
+
return x
|
| 476 |
+
|
| 477 |
+
def forward(
|
| 478 |
+
self,
|
| 479 |
+
x,
|
| 480 |
+
cu_seqlens=None,
|
| 481 |
+
max_seqlen=None,
|
| 482 |
+
rope_positions=None,
|
| 483 |
+
attn_mask=None
|
| 484 |
+
):
|
| 485 |
+
xqkv = self.wqkv(x)
|
| 486 |
+
xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3*self.head_dim)
|
| 487 |
+
|
| 488 |
+
xq, xk, xv = torch.split(xqkv, [self.head_dim]*3, dim=-1) ## seq_len, n, dim
|
| 489 |
+
|
| 490 |
+
if self.with_qk_norm:
|
| 491 |
+
xq = self.q_norm(xq)
|
| 492 |
+
xk = self.k_norm(xk)
|
| 493 |
+
|
| 494 |
+
if self.with_rope:
|
| 495 |
+
xq = self.apply_rope3d(xq, rope_positions, self.rope_ch_split, parallel=self.parallel)
|
| 496 |
+
xk = self.apply_rope3d(xk, rope_positions, self.rope_ch_split, parallel=self.parallel)
|
| 497 |
+
|
| 498 |
+
output = self.core_attention(
|
| 499 |
+
xq,
|
| 500 |
+
xk,
|
| 501 |
+
xv,
|
| 502 |
+
cu_seqlens=cu_seqlens,
|
| 503 |
+
max_seqlen=max_seqlen,
|
| 504 |
+
attn_mask=attn_mask
|
| 505 |
+
)
|
| 506 |
+
output = rearrange(output, 'b s h d -> b s (h d)')
|
| 507 |
+
output = self.wo(output)
|
| 508 |
+
|
| 509 |
+
return output
|
| 510 |
+
|
| 511 |
+
|
| 512 |
+
class CrossAttention(Attention):
|
| 513 |
+
def __init__(self, hidden_dim, head_dim, bias=False, with_qk_norm=True, attn_type='torch'):
|
| 514 |
+
super().__init__()
|
| 515 |
+
self.head_dim = head_dim
|
| 516 |
+
self.n_heads = hidden_dim // head_dim
|
| 517 |
+
|
| 518 |
+
self.wq = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
| 519 |
+
self.wkv = nn.Linear(hidden_dim, hidden_dim*2, bias=bias)
|
| 520 |
+
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
| 521 |
+
|
| 522 |
+
self.with_qk_norm = with_qk_norm
|
| 523 |
+
if self.with_qk_norm:
|
| 524 |
+
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
|
| 525 |
+
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
|
| 526 |
+
|
| 527 |
+
self.core_attention = self.attn_processor(attn_type=attn_type)
|
| 528 |
+
|
| 529 |
+
def forward(
|
| 530 |
+
self,
|
| 531 |
+
x: torch.Tensor,
|
| 532 |
+
encoder_hidden_states: torch.Tensor,
|
| 533 |
+
attn_mask=None
|
| 534 |
+
):
|
| 535 |
+
xq = self.wq(x)
|
| 536 |
+
xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)
|
| 537 |
+
|
| 538 |
+
xkv = self.wkv(encoder_hidden_states)
|
| 539 |
+
xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2*self.head_dim)
|
| 540 |
+
|
| 541 |
+
xk, xv = torch.split(xkv, [self.head_dim]*2, dim=-1) ## seq_len, n, dim
|
| 542 |
+
|
| 543 |
+
if self.with_qk_norm:
|
| 544 |
+
xq = self.q_norm(xq)
|
| 545 |
+
xk = self.k_norm(xk)
|
| 546 |
+
|
| 547 |
+
output = self.core_attention(
|
| 548 |
+
xq,
|
| 549 |
+
xk,
|
| 550 |
+
xv,
|
| 551 |
+
attn_mask=attn_mask
|
| 552 |
+
)
|
| 553 |
+
|
| 554 |
+
output = rearrange(output, 'b s h d -> b s (h d)')
|
| 555 |
+
output = self.wo(output)
|
| 556 |
+
|
| 557 |
+
return output
|
| 558 |
+
|
| 559 |
+
|
| 560 |
+
class GELU(nn.Module):
|
| 561 |
+
r"""
|
| 562 |
+
GELU activation function with tanh approximation support with `approximate="tanh"`.
|
| 563 |
+
|
| 564 |
+
Parameters:
|
| 565 |
+
dim_in (`int`): The number of channels in the input.
|
| 566 |
+
dim_out (`int`): The number of channels in the output.
|
| 567 |
+
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
|
| 568 |
+
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
| 569 |
+
"""
|
| 570 |
+
|
| 571 |
+
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
|
| 572 |
+
super().__init__()
|
| 573 |
+
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
| 574 |
+
self.approximate = approximate
|
| 575 |
+
|
| 576 |
+
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
| 577 |
+
return torch.nn.functional.gelu(gate, approximate=self.approximate)
|
| 578 |
+
|
| 579 |
+
def forward(self, hidden_states):
|
| 580 |
+
hidden_states = self.proj(hidden_states)
|
| 581 |
+
hidden_states = self.gelu(hidden_states)
|
| 582 |
+
return hidden_states
|
| 583 |
+
|
| 584 |
+
|
| 585 |
+
class FeedForward(nn.Module):
|
| 586 |
+
def __init__(
|
| 587 |
+
self,
|
| 588 |
+
dim: int,
|
| 589 |
+
inner_dim: Optional[int] = None,
|
| 590 |
+
dim_out: Optional[int] = None,
|
| 591 |
+
mult: int = 4,
|
| 592 |
+
bias: bool = False,
|
| 593 |
+
):
|
| 594 |
+
super().__init__()
|
| 595 |
+
inner_dim = dim*mult if inner_dim is None else inner_dim
|
| 596 |
+
dim_out = dim if dim_out is None else dim_out
|
| 597 |
+
self.net = nn.ModuleList([
|
| 598 |
+
GELU(dim, inner_dim, approximate="tanh", bias=bias),
|
| 599 |
+
nn.Identity(),
|
| 600 |
+
nn.Linear(inner_dim, dim_out, bias=bias)
|
| 601 |
+
])
|
| 602 |
+
|
| 603 |
+
|
| 604 |
+
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
| 605 |
+
for module in self.net:
|
| 606 |
+
hidden_states = module(hidden_states)
|
| 607 |
+
return hidden_states
|
| 608 |
+
|
| 609 |
+
|
| 610 |
+
def modulate(x, scale, shift):
|
| 611 |
+
x = x * (1 + scale) + shift
|
| 612 |
+
return x
|
| 613 |
+
|
| 614 |
+
|
| 615 |
+
def gate(x, gate):
|
| 616 |
+
x = gate * x
|
| 617 |
+
return x
|
| 618 |
+
|
| 619 |
+
|
| 620 |
+
class StepVideoTransformerBlock(nn.Module):
|
| 621 |
+
r"""
|
| 622 |
+
A basic Transformer block.
|
| 623 |
+
|
| 624 |
+
Parameters:
|
| 625 |
+
dim (`int`): The number of channels in the input and output.
|
| 626 |
+
num_attention_heads (`int`): The number of heads to use for multi-head attention.
|
| 627 |
+
attention_head_dim (`int`): The number of channels in each head.
|
| 628 |
+
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
| 629 |
+
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
|
| 630 |
+
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
|
| 631 |
+
num_embeds_ada_norm (:
|
| 632 |
+
obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
|
| 633 |
+
attention_bias (:
|
| 634 |
+
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
|
| 635 |
+
only_cross_attention (`bool`, *optional*):
|
| 636 |
+
Whether to use only cross-attention layers. In this case two cross attention layers are used.
|
| 637 |
+
double_self_attention (`bool`, *optional*):
|
| 638 |
+
Whether to use two self-attention layers. In this case no cross attention layers are used.
|
| 639 |
+
upcast_attention (`bool`, *optional*):
|
| 640 |
+
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
|
| 641 |
+
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
|
| 642 |
+
Whether to use learnable elementwise affine parameters for normalization.
|
| 643 |
+
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
|
| 644 |
+
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
|
| 645 |
+
final_dropout (`bool` *optional*, defaults to False):
|
| 646 |
+
Whether to apply a final dropout after the last feed-forward layer.
|
| 647 |
+
attention_type (`str`, *optional*, defaults to `"default"`):
|
| 648 |
+
The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
|
| 649 |
+
positional_embeddings (`str`, *optional*, defaults to `None`):
|
| 650 |
+
The type of positional embeddings to apply to.
|
| 651 |
+
num_positional_embeddings (`int`, *optional*, defaults to `None`):
|
| 652 |
+
The maximum number of positional embeddings to apply.
|
| 653 |
+
"""
|
| 654 |
+
|
| 655 |
+
def __init__(
|
| 656 |
+
self,
|
| 657 |
+
dim: int,
|
| 658 |
+
attention_head_dim: int,
|
| 659 |
+
norm_eps: float = 1e-5,
|
| 660 |
+
ff_inner_dim: Optional[int] = None,
|
| 661 |
+
ff_bias: bool = False,
|
| 662 |
+
attention_type: str = 'parallel'
|
| 663 |
+
):
|
| 664 |
+
super().__init__()
|
| 665 |
+
self.dim = dim
|
| 666 |
+
self.norm1 = nn.LayerNorm(dim, eps=norm_eps)
|
| 667 |
+
self.attn1 = SelfAttention(dim, attention_head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type=attention_type)
|
| 668 |
+
|
| 669 |
+
self.norm2 = nn.LayerNorm(dim, eps=norm_eps)
|
| 670 |
+
self.attn2 = CrossAttention(dim, attention_head_dim, bias=False, with_qk_norm=True, attn_type='torch')
|
| 671 |
+
|
| 672 |
+
self.ff = FeedForward(dim=dim, inner_dim=ff_inner_dim, dim_out=dim, bias=ff_bias)
|
| 673 |
+
|
| 674 |
+
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) /dim**0.5)
|
| 675 |
+
|
| 676 |
+
@torch.no_grad()
|
| 677 |
+
def forward(
|
| 678 |
+
self,
|
| 679 |
+
q: torch.Tensor,
|
| 680 |
+
kv: Optional[torch.Tensor] = None,
|
| 681 |
+
timestep: Optional[torch.LongTensor] = None,
|
| 682 |
+
attn_mask = None,
|
| 683 |
+
rope_positions: list = None,
|
| 684 |
+
) -> torch.Tensor:
|
| 685 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
| 686 |
+
torch.clone(chunk) for chunk in (self.scale_shift_table[None].to(dtype=q.dtype, device=q.device) + timestep.reshape(-1, 6, self.dim)).chunk(6, dim=1)
|
| 687 |
+
)
|
| 688 |
+
|
| 689 |
+
scale_shift_q = modulate(self.norm1(q), scale_msa, shift_msa)
|
| 690 |
+
|
| 691 |
+
attn_q = self.attn1(
|
| 692 |
+
scale_shift_q,
|
| 693 |
+
rope_positions=rope_positions
|
| 694 |
+
)
|
| 695 |
+
|
| 696 |
+
q = gate(attn_q, gate_msa) + q
|
| 697 |
+
|
| 698 |
+
attn_q = self.attn2(
|
| 699 |
+
q,
|
| 700 |
+
kv,
|
| 701 |
+
attn_mask
|
| 702 |
+
)
|
| 703 |
+
|
| 704 |
+
q = attn_q + q
|
| 705 |
+
|
| 706 |
+
scale_shift_q = modulate(self.norm2(q), scale_mlp, shift_mlp)
|
| 707 |
+
|
| 708 |
+
ff_output = self.ff(scale_shift_q)
|
| 709 |
+
|
| 710 |
+
q = gate(ff_output, gate_mlp) + q
|
| 711 |
+
|
| 712 |
+
return q
|
| 713 |
+
|
| 714 |
+
|
| 715 |
+
class PatchEmbed(nn.Module):
|
| 716 |
+
"""2D Image to Patch Embedding"""
|
| 717 |
+
|
| 718 |
+
def __init__(
|
| 719 |
+
self,
|
| 720 |
+
patch_size=64,
|
| 721 |
+
in_channels=3,
|
| 722 |
+
embed_dim=768,
|
| 723 |
+
layer_norm=False,
|
| 724 |
+
flatten=True,
|
| 725 |
+
bias=True,
|
| 726 |
+
):
|
| 727 |
+
super().__init__()
|
| 728 |
+
|
| 729 |
+
self.flatten = flatten
|
| 730 |
+
self.layer_norm = layer_norm
|
| 731 |
+
|
| 732 |
+
self.proj = nn.Conv2d(
|
| 733 |
+
in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
|
| 734 |
+
)
|
| 735 |
+
|
| 736 |
+
def forward(self, latent):
|
| 737 |
+
latent = self.proj(latent).to(latent.dtype)
|
| 738 |
+
if self.flatten:
|
| 739 |
+
latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC
|
| 740 |
+
if self.layer_norm:
|
| 741 |
+
latent = self.norm(latent)
|
| 742 |
+
|
| 743 |
+
return latent
|
| 744 |
+
|
| 745 |
+
|
| 746 |
+
class StepVideoModel(torch.nn.Module):
|
| 747 |
+
def __init__(
|
| 748 |
+
self,
|
| 749 |
+
num_attention_heads: int = 48,
|
| 750 |
+
attention_head_dim: int = 128,
|
| 751 |
+
in_channels: int = 64,
|
| 752 |
+
out_channels: Optional[int] = 64,
|
| 753 |
+
num_layers: int = 48,
|
| 754 |
+
dropout: float = 0.0,
|
| 755 |
+
patch_size: int = 1,
|
| 756 |
+
norm_type: str = "ada_norm_single",
|
| 757 |
+
norm_elementwise_affine: bool = False,
|
| 758 |
+
norm_eps: float = 1e-6,
|
| 759 |
+
use_additional_conditions: Optional[bool] = False,
|
| 760 |
+
caption_channels: Optional[Union[int, List, Tuple]] = [6144, 1024],
|
| 761 |
+
attention_type: Optional[str] = "torch",
|
| 762 |
+
):
|
| 763 |
+
super().__init__()
|
| 764 |
+
|
| 765 |
+
# Set some common variables used across the board.
|
| 766 |
+
self.inner_dim = num_attention_heads * attention_head_dim
|
| 767 |
+
self.out_channels = in_channels if out_channels is None else out_channels
|
| 768 |
+
|
| 769 |
+
self.use_additional_conditions = use_additional_conditions
|
| 770 |
+
|
| 771 |
+
self.pos_embed = PatchEmbed(
|
| 772 |
+
patch_size=patch_size,
|
| 773 |
+
in_channels=in_channels,
|
| 774 |
+
embed_dim=self.inner_dim,
|
| 775 |
+
)
|
| 776 |
+
|
| 777 |
+
self.transformer_blocks = nn.ModuleList(
|
| 778 |
+
[
|
| 779 |
+
StepVideoTransformerBlock(
|
| 780 |
+
dim=self.inner_dim,
|
| 781 |
+
attention_head_dim=attention_head_dim,
|
| 782 |
+
attention_type=attention_type
|
| 783 |
+
)
|
| 784 |
+
for _ in range(num_layers)
|
| 785 |
+
]
|
| 786 |
+
)
|
| 787 |
+
|
| 788 |
+
# 3. Output blocks.
|
| 789 |
+
self.norm_out = nn.LayerNorm(self.inner_dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine)
|
| 790 |
+
self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5)
|
| 791 |
+
self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels)
|
| 792 |
+
self.patch_size = patch_size
|
| 793 |
+
|
| 794 |
+
self.adaln_single = AdaLayerNormSingle(
|
| 795 |
+
self.inner_dim, use_additional_conditions=self.use_additional_conditions
|
| 796 |
+
)
|
| 797 |
+
|
| 798 |
+
if isinstance(caption_channels, int):
|
| 799 |
+
caption_channel = caption_channels
|
| 800 |
+
else:
|
| 801 |
+
caption_channel, clip_channel = caption_channels
|
| 802 |
+
self.clip_projection = nn.Linear(clip_channel, self.inner_dim)
|
| 803 |
+
|
| 804 |
+
self.caption_norm = nn.LayerNorm(caption_channel, eps=norm_eps, elementwise_affine=norm_elementwise_affine)
|
| 805 |
+
|
| 806 |
+
self.caption_projection = PixArtAlphaTextProjection(
|
| 807 |
+
in_features=caption_channel, hidden_size=self.inner_dim
|
| 808 |
+
)
|
| 809 |
+
|
| 810 |
+
self.parallel = attention_type=='parallel'
|
| 811 |
+
|
| 812 |
+
def patchfy(self, hidden_states):
|
| 813 |
+
hidden_states = rearrange(hidden_states, 'b f c h w -> (b f) c h w')
|
| 814 |
+
hidden_states = self.pos_embed(hidden_states)
|
| 815 |
+
return hidden_states
|
| 816 |
+
|
| 817 |
+
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states, q_seqlen):
|
| 818 |
+
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
|
| 819 |
+
mask = torch.zeros([len(kv_seqlens), q_seqlen, max(kv_seqlens)], dtype=torch.bool, device=encoder_attention_mask.device)
|
| 820 |
+
encoder_hidden_states = encoder_hidden_states[:,: max(kv_seqlens)]
|
| 821 |
+
for i, kv_len in enumerate(kv_seqlens):
|
| 822 |
+
mask[i, :, :kv_len] = 1
|
| 823 |
+
return encoder_hidden_states, mask
|
| 824 |
+
|
| 825 |
+
|
| 826 |
+
def block_forward(
|
| 827 |
+
self,
|
| 828 |
+
hidden_states,
|
| 829 |
+
encoder_hidden_states=None,
|
| 830 |
+
timestep=None,
|
| 831 |
+
rope_positions=None,
|
| 832 |
+
attn_mask=None,
|
| 833 |
+
parallel=True
|
| 834 |
+
):
|
| 835 |
+
for block in tqdm(self.transformer_blocks, desc="Transformer blocks"):
|
| 836 |
+
hidden_states = block(
|
| 837 |
+
hidden_states,
|
| 838 |
+
encoder_hidden_states,
|
| 839 |
+
timestep=timestep,
|
| 840 |
+
attn_mask=attn_mask,
|
| 841 |
+
rope_positions=rope_positions
|
| 842 |
+
)
|
| 843 |
+
|
| 844 |
+
return hidden_states
|
| 845 |
+
|
| 846 |
+
|
| 847 |
+
@torch.inference_mode()
|
| 848 |
+
def forward(
|
| 849 |
+
self,
|
| 850 |
+
hidden_states: torch.Tensor,
|
| 851 |
+
encoder_hidden_states: Optional[torch.Tensor] = None,
|
| 852 |
+
encoder_hidden_states_2: Optional[torch.Tensor] = None,
|
| 853 |
+
timestep: Optional[torch.LongTensor] = None,
|
| 854 |
+
added_cond_kwargs: Dict[str, torch.Tensor] = None,
|
| 855 |
+
encoder_attention_mask: Optional[torch.Tensor] = None,
|
| 856 |
+
fps: torch.Tensor=None,
|
| 857 |
+
return_dict: bool = False,
|
| 858 |
+
):
|
| 859 |
+
assert hidden_states.ndim==5; "hidden_states's shape should be (bsz, f, ch, h ,w)"
|
| 860 |
+
|
| 861 |
+
bsz, frame, _, height, width = hidden_states.shape
|
| 862 |
+
height, width = height // self.patch_size, width // self.patch_size
|
| 863 |
+
|
| 864 |
+
hidden_states = self.patchfy(hidden_states)
|
| 865 |
+
len_frame = hidden_states.shape[1]
|
| 866 |
+
|
| 867 |
+
if self.use_additional_conditions:
|
| 868 |
+
added_cond_kwargs = {
|
| 869 |
+
"resolution": torch.tensor([(height, width)]*bsz, device=hidden_states.device, dtype=hidden_states.dtype),
|
| 870 |
+
"nframe": torch.tensor([frame]*bsz, device=hidden_states.device, dtype=hidden_states.dtype),
|
| 871 |
+
"fps": fps
|
| 872 |
+
}
|
| 873 |
+
else:
|
| 874 |
+
added_cond_kwargs = {}
|
| 875 |
+
|
| 876 |
+
timestep, embedded_timestep = self.adaln_single(
|
| 877 |
+
timestep, added_cond_kwargs=added_cond_kwargs
|
| 878 |
+
)
|
| 879 |
+
|
| 880 |
+
encoder_hidden_states = self.caption_projection(self.caption_norm(encoder_hidden_states))
|
| 881 |
+
|
| 882 |
+
if encoder_hidden_states_2 is not None and hasattr(self, 'clip_projection'):
|
| 883 |
+
clip_embedding = self.clip_projection(encoder_hidden_states_2)
|
| 884 |
+
encoder_hidden_states = torch.cat([clip_embedding, encoder_hidden_states], dim=1)
|
| 885 |
+
|
| 886 |
+
hidden_states = rearrange(hidden_states, '(b f) l d-> b (f l) d', b=bsz, f=frame, l=len_frame).contiguous()
|
| 887 |
+
encoder_hidden_states, attn_mask = self.prepare_attn_mask(encoder_attention_mask, encoder_hidden_states, q_seqlen=frame*len_frame)
|
| 888 |
+
|
| 889 |
+
hidden_states = self.block_forward(
|
| 890 |
+
hidden_states,
|
| 891 |
+
encoder_hidden_states,
|
| 892 |
+
timestep=timestep,
|
| 893 |
+
rope_positions=[frame, height, width],
|
| 894 |
+
attn_mask=attn_mask,
|
| 895 |
+
parallel=self.parallel
|
| 896 |
+
)
|
| 897 |
+
|
| 898 |
+
hidden_states = rearrange(hidden_states, 'b (f l) d -> (b f) l d', b=bsz, f=frame, l=len_frame)
|
| 899 |
+
|
| 900 |
+
embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d', f=frame).contiguous()
|
| 901 |
+
|
| 902 |
+
shift, scale = (self.scale_shift_table[None].to(dtype=embedded_timestep.dtype, device=embedded_timestep.device) + embedded_timestep[:, None]).chunk(2, dim=1)
|
| 903 |
+
hidden_states = self.norm_out(hidden_states)
|
| 904 |
+
# Modulation
|
| 905 |
+
hidden_states = hidden_states * (1 + scale) + shift
|
| 906 |
+
hidden_states = self.proj_out(hidden_states)
|
| 907 |
+
|
| 908 |
+
# unpatchify
|
| 909 |
+
hidden_states = hidden_states.reshape(
|
| 910 |
+
shape=(-1, height, width, self.patch_size, self.patch_size, self.out_channels)
|
| 911 |
+
)
|
| 912 |
+
|
| 913 |
+
hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')
|
| 914 |
+
output = hidden_states.reshape(
|
| 915 |
+
shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size)
|
| 916 |
+
)
|
| 917 |
+
|
| 918 |
+
output = rearrange(output, '(b f) c h w -> b f c h w', f=frame)
|
| 919 |
+
|
| 920 |
+
if return_dict:
|
| 921 |
+
return {'x': output}
|
| 922 |
+
return output
|
| 923 |
+
|
| 924 |
+
@staticmethod
|
| 925 |
+
def state_dict_converter():
|
| 926 |
+
return StepVideoDiTStateDictConverter()
|
| 927 |
+
|
| 928 |
+
|
| 929 |
+
class StepVideoDiTStateDictConverter:
|
| 930 |
+
def __init__(self):
|
| 931 |
+
super().__init__()
|
| 932 |
+
|
| 933 |
+
def from_diffusers(self, state_dict):
|
| 934 |
+
return state_dict
|
| 935 |
+
|
| 936 |
+
def from_civitai(self, state_dict):
|
| 937 |
+
return state_dict
|
| 938 |
+
|
| 939 |
+
|
| 940 |
+
|