Compare commits
15
Commits
kernels
...
hangliang_sw
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b435db8271 | ||
|
|
2e66366024 | ||
|
|
78bb902ab9 | ||
|
|
24a72543b8 | ||
|
|
1fa986b4a4 | ||
|
|
69714a58b0 | ||
|
|
86412d8ec6 | ||
|
|
dc072bc989 | ||
|
|
c530f8188b | ||
|
|
39350417e2 | ||
|
|
664cf4591d | ||
|
|
92475e19d8 | ||
|
|
5cc8daf0da | ||
|
|
99dbdaacac | ||
|
|
dbe71057d5 |
@@ -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)
|
||||
])
|
||||
|
||||
@@ -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,
|
||||
|
||||
File diff suppressed because one or more lines are too long
+681
File diff suppressed because one or more lines are too long
@@ -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
|
||||
Reference in New Issue
Block a user