Compare commits

...
Author SHA1 Message Date
foreverpiano b435db8271 pd 2025-01-21 13:30:29 +00:00
foreverpiano 2e66366024 upd 2025-01-21 13:17:29 +00:00
foreverpiano 78bb902ab9 upd 2025-01-20 07:02:09 +00:00
foreverpiano 24a72543b8 upd 2025-01-20 06:49:06 +00:00
foreverpiano 1fa986b4a4 upd 2025-01-20 06:42:52 +00:00
foreverpiano 69714a58b0 upd 2025-01-20 06:41:20 +00:00
foreverpiano 86412d8ec6 upd 2025-01-20 06:37:19 +00:00
foreverpiano dc072bc989 upd 2025-01-19 15:11:09 +00:00
foreverpiano c530f8188b upd 2025-01-19 14:49:39 +00:00
foreverpiano 39350417e2 udp 2025-01-19 07:11:39 +00:00
foreverpiano 664cf4591d upd 2025-01-18 11:16:14 +00:00
foreverpiano 92475e19d8 upd 2025-01-18 07:53:12 +00:00
foreverpiano 5cc8daf0da upd 2025-01-17 17:05:17 +00:00
foreverpiano 99dbdaacac upd 2025-01-17 16:39:46 +00:00
foreverpiano dbe71057d5 update window 2025-01-17 16:20:44 +00:00
6 changed files with 1070 additions and 2 deletions
+95 -2
View File
@@ -32,7 +32,7 @@ def attention(
return out
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, sample_step, layer_id):
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
query, encoder_query = q
@@ -57,6 +57,90 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# attn_map plot
num_heads = query.size(2)
map_q = query.transpose(1, 2) # [B, H, S, D]
map_k = torch.cat([key, encoder_key], dim=1).transpose(1, 2) # [B, H, S+T, D]
d_k = map_q.size(-1)
shape = (16, 45, 45)
q_coords = torch.tensor([get_block(i, shape=shape) for i in range(sequence_length)],
device='cuda', dtype=torch.int16)
k_coords = torch.tensor([get_block(i, shape=shape) for i in range(img_kv_len)],
device='cuda', dtype=torch.int16)
diffs = (q_coords.unsqueeze(1) - k_coords.unsqueeze(0)).abs() # [seq_len, kv_len, 3]
mask_t_diff = 6
mask_x_diff = 12
mask_y_diff = 12
mask_t = diffs[..., 0] <= mask_t_diff
mask_x = diffs[..., 1] <= mask_x_diff
mask_y = diffs[..., 2] <= mask_y_diff
valid_mask_2d = mask_t & mask_x & mask_y # [seq_len, kv_len]
del mask_t, mask_x, mask_y, diffs
# Pad mask for text part (all True)
full_valid_mask = F.pad(valid_mask_2d, (0, encoder_sequence_length), value=True) # [seq_len, kv_len+text_len]
img_mask_density = valid_mask_2d.float().mean().item()
full_mask_density = full_valid_mask.float().mean().item()
del valid_mask_2d
chunk_size = 512
save_attn = False #(sample_step == 1) and (layer_id == 59)
if save_attn:
attn_map_cumulated = torch.zeros((sequence_length, img_kv_len + encoder_sequence_length),
dtype=torch.float32, device='cuda')
# 逐head计算
for head_idx in range(num_heads):
current_q = map_q[:, head_idx:head_idx+1].to(dtype=torch.float32) # [B, 1, S, D]
current_k = map_k[:, head_idx:head_idx+1].to(dtype=torch.float32) # [B, 1, S+T, D]
valid_score_total = 0.0
all_score_total = 0.0
if save_attn:
head_attn_map = torch.zeros((sequence_length, img_kv_len + encoder_sequence_length),
dtype=torch.float32, device='cuda')
# 分块计算
for i in range(0, current_q.size(2), chunk_size):
chunk_end = min(i + chunk_size, current_q.size(2))
q_chunk = current_q[:, :, i:chunk_end] # [B, 1, chunk_size, D]
scores_32 = torch.matmul(
q_chunk,
current_k.transpose(-2, -1)
) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32, device='cuda'))
attn_weights = F.softmax(scores_32, dim=-1)
chunk_valid_mask = full_valid_mask[i:chunk_end].unsqueeze(0).unsqueeze(0)
valid_score_sum = (attn_weights * chunk_valid_mask).sum()
all_score_sum = attn_weights.sum()
valid_score_total += valid_score_sum.item()
all_score_total += all_score_sum.item()
# For last layer, accumulate attention weights
if save_attn:
head_attn_map[i:chunk_end] = attn_weights.squeeze(0).squeeze(0)
del scores_32, attn_weights
torch.cuda.empty_cache()
recall = valid_score_total / (all_score_total + 1e-9)
print(f"step{sample_step}_layer{layer_id}_head{head_idx}_window_diff_{mask_t_diff}_{mask_x_diff}_{mask_y_diff}_shape_{shape[0]}_{shape[1]}_{shape[2]}_img{img_mask_density:.4f}_full{full_mask_density:.4f}'s recall:", recall)
if save_attn:
attn_map_cumulated += head_attn_map
if save_attn:
attn_map_avg = attn_map_cumulated / num_heads
torch.save(attn_map_avg, f'logs/attn_map_step{sample_step}_layer{layer_id}.pt')
# Hint: please check encoder_query.shape
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
@@ -70,7 +154,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
causal=False,
dropout_p=0.0,
softmax_scale=None)
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1)
if get_sequence_parallel_state():
@@ -88,3 +172,12 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
attn = attn.reshape(b, s, -1)
return attn
def get_block(idx, shape=(16, 30, 30)):
t_size, x_size, y_size = shape
xy = x_size * y_size
t = idx // xy
r = idx % xy
x = r // y_size
y = r % y_size
return t, x, y
@@ -38,6 +38,7 @@ class MMDoubleStreamBlock(nn.Module):
qkv_bias: bool = False,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
layer_id: int = 0,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
@@ -124,6 +125,8 @@ class MMDoubleStreamBlock(nn.Module):
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
self.sample_step = 0
self.layer_id = layer_id
def enable_deterministic(self):
self.deterministic = True
@@ -207,6 +210,10 @@ class MMDoubleStreamBlock(nn.Module):
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
print("DOUBLE====")
print(img_q.shape, txt_q.shape)
self.sample_step += 1
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
@@ -214,6 +221,8 @@ class MMDoubleStreamBlock(nn.Module):
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
sample_step=self.sample_step,
layer_id=self.layer_id,
)
# attention computation end
@@ -264,6 +273,7 @@ class MMSingleStreamBlock(nn.Module):
qk_scale: float = None,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
layer_id: int = 0,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
@@ -304,6 +314,8 @@ class MMSingleStreamBlock(nn.Module):
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
self.sample_step = 0
self.layer_id = layer_id
def enable_deterministic(self):
self.deterministic = True
@@ -354,6 +366,7 @@ class MMSingleStreamBlock(nn.Module):
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
self.sample_step += 1
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
@@ -361,6 +374,8 @@ class MMSingleStreamBlock(nn.Module):
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
sample_step=self.sample_step,
layer_id=self.layer_id,
)
# attention computation end
@@ -521,6 +536,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
layer_id=_,
**factory_kwargs,
) for _ in range(mm_double_blocks_depth)
])
@@ -534,6 +550,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
layer_id=_+mm_double_blocks_depth,
**factory_kwargs,
) for _ in range(mm_single_blocks_depth)
])
+1
View File
@@ -50,6 +50,7 @@ def main(args):
prompts = f.readlines()
for prompt in prompts:
print("#####PROMPT#####", prompt)
outputs = hunyuan_video_sampler.predict(
prompt=prompt,
height=args.height,
+256
View File
File diff suppressed because one or more lines are too long
+681
View File
File diff suppressed because one or more lines are too long
+20
View File
@@ -0,0 +1,20 @@
#!/bin/bash
num_gpus=1
export MODEL_BASE=data/hunyuan
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/sample/sample_t2v_hunyuan.py \
--height 720 \
--width 720 \
--num_frames 61 \
--num_inference_steps 6 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 17 \
--flow-reverse \
--prompt ./assets/prompt.txt \
--seed 1024 \
--output_path outputs_video/hunyuan/sw/ \
--model_path $MODEL_BASE \
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
--vae-sp