Skip to content

[do not merge] refactor Attention class for HunYuan DIT - #8265

Closed
yiyixuxu wants to merge 2 commits into
mainfrom
hunyuan-yiyi
Closed

[do not merge] refactor Attention class for HunYuan DIT #8265
yiyixuxu wants to merge 2 commits into
mainfrom
hunyuan-yiyi

Conversation

@yiyixuxu

Copy link
Copy Markdown
Collaborator
# integration test (hunyuan dit)importtorchfromdiffusersimportHunyuanDiTPipelineimporttorchfromhuggingface_hubimporthf_hub_downloadfromdiffusersimportHunyuanDiT2DModelimportsafetensors.torchdevice="cuda"model_config=HunyuanDiT2DModel.load_config("XCLiu/HunyuanDiT-0523", subfolder="transformer")
model=HunyuanDiT2DModel.from_config(model_config).to(device)
ckpt_path=hf_hub_download(
"XCLiu/HunyuanDiT-0523",
filename="diffusion_pytorch_model.safetensors",
subfolder="transformer",
)
state_dict=safetensors.torch.load_file(ckpt_path)
num_layers=40foriinrange(num_layers):
# attn1# Wkqv -> to_q, to_k, to_vq, k, v=torch.chunk(state_dict[f"blocks.{i}.attn1.Wqkv.weight"], 3, dim=0)
q_bias, k_bias, v_bias=torch.chunk(state_dict[f"blocks.{i}.attn1.Wqkv.bias"], 3, dim=0)
state_dict[f"blocks.{i}.attn1.to_q.weight"] =qstate_dict[f"blocks.{i}.attn1.to_q.bias"] =q_biasstate_dict[f"blocks.{i}.attn1.to_k.weight"] =kstate_dict[f"blocks.{i}.attn1.to_k.bias"] =k_biasstate_dict[f"blocks.{i}.attn1.to_v.weight"] =vstate_dict[f"blocks.{i}.attn1.to_v.bias"] =v_biasstate_dict.pop(f"blocks.{i}.attn1.Wqkv.weight")
state_dict.pop(f"blocks.{i}.attn1.Wqkv.bias")
# q_norm, k_norm -> norm_q, norm_kstate_dict[f"blocks.{i}.attn1.norm_q.weight"] =state_dict[f"blocks.{i}.attn1.q_norm.weight"]
state_dict[f"blocks.{i}.attn1.norm_q.bias"] =state_dict[f"blocks.{i}.attn1.q_norm.bias"]
state_dict[f"blocks.{i}.attn1.norm_k.weight"] =state_dict[f"blocks.{i}.attn1.k_norm.weight"]
state_dict[f"blocks.{i}.attn1.norm_k.bias"] =state_dict[f"blocks.{i}.attn1.k_norm.bias"]
state_dict.pop(f"blocks.{i}.attn1.q_norm.weight")
state_dict.pop(f"blocks.{i}.attn1.q_norm.bias")
state_dict.pop(f"blocks.{i}.attn1.k_norm.weight")
state_dict.pop(f"blocks.{i}.attn1.k_norm.bias")
# out_proj -> to_outstate_dict[f"blocks.{i}.attn1.to_out.0.weight"] =state_dict[f"blocks.{i}.attn1.out_proj.weight"]
state_dict[f"blocks.{i}.attn1.to_out.0.bias"] =state_dict[f"blocks.{i}.attn1.out_proj.bias"]
state_dict.pop(f"blocks.{i}.attn1.out_proj.weight")
state_dict.pop(f"blocks.{i}.attn1.out_proj.bias")
# attn2# kq_proj -> to_k, to_vk, v=torch.chunk(state_dict[f"blocks.{i}.attn2.kv_proj.weight"], 2, dim=0)
k_bias, v_bias=torch.chunk(state_dict[f"blocks.{i}.attn2.kv_proj.bias"], 2, dim=0)
state_dict[f"blocks.{i}.attn2.to_k.weight"] =kstate_dict[f"blocks.{i}.attn2.to_k.bias"] =k_biasstate_dict[f"blocks.{i}.attn2.to_v.weight"] =vstate_dict[f"blocks.{i}.attn2.to_v.bias"] =v_biasstate_dict.pop(f"blocks.{i}.attn2.kv_proj.weight")
state_dict.pop(f"blocks.{i}.attn2.kv_proj.bias")
# q_proj -> to_qstate_dict[f"blocks.{i}.attn2.to_q.weight"] =state_dict[f"blocks.{i}.attn2.q_proj.weight"]
state_dict[f"blocks.{i}.attn2.to_q.bias"] =state_dict[f"blocks.{i}.attn2.q_proj.bias"]
state_dict.pop(f"blocks.{i}.attn2.q_proj.weight")
state_dict.pop(f"blocks.{i}.attn2.q_proj.bias")
# q_norm, k_norm -> norm_q, norm_kstate_dict[f"blocks.{i}.attn2.norm_q.weight"] =state_dict[f"blocks.{i}.attn2.q_norm.weight"]
state_dict[f"blocks.{i}.attn2.norm_q.bias"] =state_dict[f"blocks.{i}.attn2.q_norm.bias"]
state_dict[f"blocks.{i}.attn2.norm_k.weight"] =state_dict[f"blocks.{i}.attn2.k_norm.weight"]
state_dict[f"blocks.{i}.attn2.norm_k.bias"] =state_dict[f"blocks.{i}.attn2.k_norm.bias"]
state_dict.pop(f"blocks.{i}.attn2.q_norm.weight")
state_dict.pop(f"blocks.{i}.attn2.q_norm.bias")
state_dict.pop(f"blocks.{i}.attn2.k_norm.weight")
state_dict.pop(f"blocks.{i}.attn2.k_norm.bias")
# out_proj -> to_outstate_dict[f"blocks.{i}.attn2.to_out.0.weight"] =state_dict[f"blocks.{i}.attn2.out_proj.weight"]
state_dict[f"blocks.{i}.attn2.to_out.0.bias"] =state_dict[f"blocks.{i}.attn2.out_proj.bias"]
state_dict.pop(f"blocks.{i}.attn2.out_proj.weight")
state_dict.pop(f"blocks.{i}.attn2.out_proj.bias")
model.load_state_dict(state_dict)
pipe=HunyuanDiTPipeline.from_pretrained("XCLiu/HunyuanDiT-0523", transformer=model, torch_dtype=torch.float32)
pipe.to('cuda')
### NOTE: HunyuanDiT supports both Chinese and English inputsprompt="一个宇航员在骑马"#prompt = "An astronaut riding a horse"image=pipe(height=1024, width=1024, prompt=prompt).images[0]
image.save("yiyi_test_5_out.png")

yiyi_test_5_out

LoRAAttnAddedKVProcessor,
]

from typing import Tuple

Copy link
Copy Markdown
CollaboratorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

put it here for now (I don't think they belong here)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hi I tested this branch and it gives the same outputs as the old pipeline with the same generator. The code looks good to me.

Sign up for freeto join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants

@yiyixuxu@gnobitab