Compare commits

...
Author SHA1 Message Date
rlsu9 805978d0da stage bash diff 2025-03-13 17:11:55 +00:00
rlsu9 797cb2f659 recover sample hunyuan 2025-03-13 17:08:02 +00:00
rlsu9 ce749878be fix attention gap 2025-03-13 17:04:53 +00:00
rlsu9 8c844ba8ac stage change 2025-03-13 16:42:43 +00:00
rlsu9 69acec287e stage sh 2025-03-13 16:19:44 +00:00
rlsu9 5a2fd28f95 fix cache lime error 2025-03-13 08:15:18 +00:00
rlsu9 3133dbf69d remove unused item 2025-03-13 07:25:47 +00:00
rlsu9 0d3a05d58e remove unused difference 2025-03-13 07:16:39 +00:00
rlsu9 f1fc13b28d remove unused files 2025-03-13 07:05:39 +00:00
rlsu9 0b4033decc upload debug code 2025-03-13 07:02:11 +00:00
rlsu9 6b119a9f76 stage 2025-03-13 05:38:29 +00:00
rlsu9 15fb1d9a57 add tk debug 2025-03-07 01:29:22 +00:00
rlsu9 66060dd22f update kernel benchmark 2025-03-06 23:59:00 +00:00
rlsu9 9fe8704537 stage kernel-flops benchmark 2025-03-06 23:23:50 +00:00
rlsu9 31bbdcd2a4 stage 2025-03-04 00:51:21 +00:00
4 changed files with 88 additions and 7 deletions
@@ -813,7 +813,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
torch._dynamo.config.cache_size_limit = 128
mask_strategy = dict_to_3d_list(mask_strategy)
# if is_progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
+85 -4
View File
@@ -6,12 +6,68 @@ try:
from st_attn import sliding_tile_attention
except ImportError:
print("Could not load Sliding Tile Attention.")
sliding_tile_attention = None
sliding_tile_attention = None
from functools import lru_cache
from torch.nn.attention.flex_attention import flex_attention
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from csrc.sliding_tile_attention.test.flex_sta_ref import get_sliding_tile_attention_mask
@lru_cache(maxsize=32)
def get_compiled_flex_attention(strategy, tile_size, image_size, text_length, device):
"""
Create and compile flex attention with a specific sliding block mask.
This function is cached to avoid recompiling for the same parameters.
Args:
strategy (tuple): A tuple (t, h, w) defining the strategy
tile_size (tuple): A tuple (ts_t, ts_h, ts_w) defining the tile size
image_size (tuple): A tuple (n_t, n_h, n_w) defining the image size
text_length (int): The text length
device (str): The device to use
Returns:
function: A compiled flex attention function with the specified mask
"""
# Convert strategy to the required format (ceil(t*3/2), h*2, w)
adjusted_strategy = strategy
# Get the sliding block attention mask
mask = get_sliding_tile_attention_mask(
adjusted_strategy,
tile_size,
image_size,
text_length,
device
)
def flex_attn_with_mask(q, k, v, scale=None):
return flex_attention(q, k, v, block_mask=mask, scale=scale)
# Compile the wrapper function
compiled_flex_attn = torch.compile(flex_attn_with_mask)
return compiled_flex_attn
def flex_sliding_tile_attention(q_all, k_all, v_all, strategy, tile_size,
image_size, text_length, scale=None):
device = q_all.device
# Get the compiled flex attention function (cached if called with same parameters)
compiled_flex_attn = get_compiled_flex_attention(
strategy,
tile_size,
image_size,
text_length,
device
)
# Apply the compiled flex attention
output = compiled_flex_attn(q_all, k_all, v_all, scale=scale)
return output
def attention(
q,
@@ -92,7 +148,32 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
start_head = current_rank * head_num
windows = [mask_strategy[head_idx + start_head] for head_idx in range(head_num)]
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
if sliding_tile_attention is not None:
hidden_states = sliding_tile_attention(query, key, value, windows, text_length).transpose(1, 2)
else:
print("Sliding Tile Attention not available. Using Flex Sliding Tile Attention.")
hidden_states = torch.empty_like(query)
strategy_to_heads = {}
for head_index in range(head_num):
strategy = tuple(windows[head_index]) # Convert list to tuple for dict key
if strategy not in strategy_to_heads:
strategy_to_heads[strategy] = []
strategy_to_heads[strategy].append(head_index)
for strategy, heads in strategy_to_heads.items():
# Gather all heads with this strategy
query_heads = torch.cat([query[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
key_heads = torch.cat([key[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
value_heads = torch.cat([value[:, head_idx:head_idx + 1, :, :] for head_idx in heads], dim=1)
# Process all heads with this strategy at once
# processed_heads = selected_attn_processor[processor_idx](query_heads, key_heads, value_heads)
processed_heads = flex_sliding_tile_attention(query_heads, key_heads, value_heads, strategy, (6, 8, 8), (30, 48, 80), text_length)
# Distribute results back to the correct positions
for i, head_idx in enumerate(heads):
hidden_states[:, head_idx:head_idx + 1, :, :] = processed_heads[:, i:i + 1, :, :]
hidden_states = hidden_states.transpose(1, 2)
else:
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
@@ -121,4 +202,4 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
b, s, a, d = attn.shape
attn = attn.reshape(b, s, -1)
return attn
return attn
+1 -1
View File
@@ -207,4 +207,4 @@ if __name__ == "__main__":
# process for vae sequence parallel
if args.vae_sp and not args.vae_tiling:
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
main(args)
main(args)
+1 -1
View File
@@ -37,4 +37,4 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
--output_path outputs_video/hunyuan/vae_sp/ \
--model_path $MODEL_BASE \
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
--vae-sp
--vae-sp