Compare commits
10
Commits
will/design
...
ms-recipe
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d9d4bb5392 | ||
|
|
1d9a593ba9 | ||
|
|
1d9364c29c | ||
|
|
68d012210c | ||
|
|
89ed16efa4 | ||
|
|
b2a581c45d | ||
|
|
b267d0d041 | ||
|
|
782dae2739 | ||
|
|
722c47932f | ||
|
|
377d1607ba |
@@ -45,6 +45,32 @@ The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
```
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
## 🎯 STA mask search pipeline
|
||||
### Overview
|
||||
|
||||
The STA mask search pipeline consists of three sequential steps:
|
||||
|
||||
1. **Searching**: Choose sparse attention mask candidates and do searching
|
||||
2. **Tuning**: Use L2 loss to determine optimal mask strategy
|
||||
3. **Inference**: Apply selected strategy for fast video generation
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan.sh # Inference stepvideo with STA
|
||||
```
|
||||
The only thing you need to do is to specify ```--STA_mode``` with original hunyuan inference script.
|
||||
#### Step 1: Searching
|
||||
Run with ```--STA_mode STA_searching```, and this step generates a folder containing mask search results in JSON format for each prompt.
|
||||
#### Step 2: Tuning
|
||||
Run with ```--STA_mode STA_tuning```. During this step, the system will:
|
||||
1. Reads all JSON files from the search results folder
|
||||
2. Averages L2 distances across different masks to determine the optimal mask strategy per attention head. (First 12-15 steps will be full mask to get better quality)
|
||||
3. Generates accelerated videos for evaluation
|
||||
4. Saves the best strategy to a single json file
|
||||
#### Step 3: Inference
|
||||
After determining the optimal strategy, run with ```--STA_mode STA_inference```. This step reads the strategy file and runs inference with the optimized settings.
|
||||
#### Configuration
|
||||
You can modify various STA configuration parameters in:
|
||||
```fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py```
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
@@ -550,7 +550,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
enable_vae_sp: bool = False,
|
||||
n_tokens: Optional[int] = None,
|
||||
embedded_guidance_scale: Optional[float] = None,
|
||||
mask_strategy: Optional[Dict[str, list]] = None,
|
||||
STA_mode: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
@@ -782,6 +782,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
img_size = latents.shape[-3:]
|
||||
img_size = (img_size[0], img_size[1] // 2, img_size[2] // 2)
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
@@ -801,26 +804,36 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32) and not self.args.disable_autocast
|
||||
|
||||
# STA
|
||||
from fastvideo.utils.STA_configuration import configure_sta
|
||||
mask_search_final_result = []
|
||||
sparse_mask_candidates = ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1"]
|
||||
full_mask = ["5,6,10"]
|
||||
STA_param = None
|
||||
if STA_mode == 'STA_searching':
|
||||
STA_param = configure_sta(
|
||||
mode='STA_searching',
|
||||
mask_candidates=sparse_mask_candidates +
|
||||
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
|
||||
)
|
||||
elif STA_mode == 'STA_tuning':
|
||||
STA_param = configure_sta(
|
||||
mode='STA_tuning',
|
||||
mask_search_files_path='output/mask_search_result/',
|
||||
mask_candidates=sparse_mask_candidates,
|
||||
skip_time_steps=15, # Use full attention for first 15 steps
|
||||
save_dir='output/mask_strategy' # Custom save directory
|
||||
)
|
||||
elif STA_mode == 'STA_inference':
|
||||
STA_param = configure_sta(mode='STA_inference', load_path='output/mask_strategy/mask_strategy.json')
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
# if is_progress_bar:
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
@@ -841,15 +854,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
encoder_hidden_states = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
|
||||
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
noise_pred, _, mask_search_result = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input,
|
||||
encoder_hidden_states,
|
||||
t_expand,
|
||||
prompt_mask,
|
||||
mask_strategy=mask_strategy[i],
|
||||
STA_param=STA_param[i],
|
||||
guidance=guidance_expand,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
)
|
||||
mask_search_final_result.append(mask_search_result)
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
@@ -888,6 +902,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
if STA_mode == 'STA_searching':
|
||||
from fastvideo.utils.STA_configuration import save_mask_search_results
|
||||
save_mask_search_results(mask_search_final_result,
|
||||
prompt=prompt,
|
||||
mask_strategies=sparse_mask_candidates,
|
||||
output_dir='output/mask_search_result_test/')
|
||||
|
||||
if not output_type == "latent":
|
||||
expand_temporal_dim = False
|
||||
if len(latents.shape) == 4:
|
||||
|
||||
@@ -338,7 +338,7 @@ class HunyuanVideoSampler(Inference):
|
||||
embedded_guidance_scale=None,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
mask_strategy=None,
|
||||
STA_mode=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -471,7 +471,7 @@ class HunyuanVideoSampler(Inference):
|
||||
vae_ver=self.args.vae,
|
||||
enable_tiling=self.args.vae_tiling,
|
||||
enable_vae_sp=self.args.vae_sp,
|
||||
mask_strategy=mask_strategy,
|
||||
STA_mode=STA_mode,
|
||||
)[0]
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
@@ -58,7 +58,7 @@ def untile(x, sp_size):
|
||||
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
|
||||
|
||||
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=None):
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, STA_param=None):
|
||||
query, encoder_query = q
|
||||
key, encoder_key = k
|
||||
value, encoder_value = v
|
||||
@@ -81,18 +81,45 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
|
||||
sequence_length = query.size(1)
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
loss_result = None
|
||||
if STA_param[0] is not None:
|
||||
query = torch.cat([tile(query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2)
|
||||
key = torch.cat([tile(key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2)
|
||||
value = torch.cat([tile(value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2)
|
||||
|
||||
head_num = query.size(1)
|
||||
current_rank = nccl_info.rank_within_group
|
||||
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 len(STA_param) < 24: # searching mode; thus do not use more than 24 mask candidates
|
||||
sparse_attn_hidden_states_all = []
|
||||
full_mask_window = STA_param[-1]
|
||||
for window_size in STA_param[:-1]:
|
||||
hidden_states = sliding_tile_attention(query, key, value, [window_size] * head_num,
|
||||
text_length).transpose(1, 2)
|
||||
sparse_attn_hidden_states_all.append(hidden_states)
|
||||
|
||||
hidden_states = sliding_tile_attention(query, key, value, [full_mask_window] * head_num,
|
||||
text_length).transpose(1, 2) # torch.Size([1, 115456, 24, 128])
|
||||
|
||||
attn_L2_loss = []
|
||||
attn_L1_loss = []
|
||||
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
|
||||
# L2 loss
|
||||
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() - hidden_states.float())**2,
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
|
||||
attn_L2_loss.append(attn_L2_loss_)
|
||||
# L1 loss
|
||||
attn_L1_loss_ = torch.mean(torch.abs(sparse_attn_hidden_states.float() - hidden_states.float()),
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
|
||||
attn_L1_loss.append(attn_L1_loss_)
|
||||
|
||||
loss_result = [attn_L2_loss, attn_L1_loss]
|
||||
else:
|
||||
current_rank = nccl_info.rank_within_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [STA_param[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)
|
||||
else:
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
@@ -106,7 +133,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes((sequence_length, encoder_sequence_length),
|
||||
dim=1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
if STA_param[0] is not None:
|
||||
hidden_states = untile(hidden_states, nccl_info.sp_size)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
@@ -121,4 +148,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, loss_result
|
||||
|
||||
@@ -109,7 +109,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
STA_param=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
@@ -163,16 +163,23 @@ 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)
|
||||
|
||||
attn = parallel_attention(
|
||||
attn, loss_result = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
(img_k, txt_k),
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
STA_param=STA_param,
|
||||
)
|
||||
|
||||
if loss_result is not None:
|
||||
layer_loss_save = {
|
||||
"L2_loss": loss_result[0],
|
||||
"L1_loss": loss_result[1],
|
||||
}
|
||||
else:
|
||||
layer_loss_save = None
|
||||
# attention computation end
|
||||
|
||||
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
|
||||
@@ -190,7 +197,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
self.txt_mlp(modulate(self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale)),
|
||||
gate=txt_mod2_gate,
|
||||
)
|
||||
return img, txt
|
||||
return img, txt, layer_loss_save
|
||||
|
||||
|
||||
class MMSingleStreamBlock(nn.Module):
|
||||
@@ -259,7 +266,7 @@ class MMSingleStreamBlock(nn.Module):
|
||||
txt_len: int,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
STA_param=None,
|
||||
) -> torch.Tensor:
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
|
||||
@@ -288,21 +295,28 @@ 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
|
||||
|
||||
attn = parallel_attention(
|
||||
attn, loss_result = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
(img_k, txt_k),
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
STA_param=STA_param,
|
||||
)
|
||||
|
||||
if loss_result is not None:
|
||||
layer_loss_save = {
|
||||
"L2_loss": loss_result[0],
|
||||
"L1_loss": loss_result[1],
|
||||
}
|
||||
else:
|
||||
layer_loss_save = None
|
||||
# attention computation end
|
||||
|
||||
# Compute activation in mlp stream, cat again and run second linear layer.
|
||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||
return x + apply_gate(output, gate=mod_gate)
|
||||
return x + apply_gate(output, gate=mod_gate), layer_loss_save
|
||||
|
||||
|
||||
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
@@ -514,7 +528,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
mask_strategy=None,
|
||||
STA_param=None,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
@@ -523,9 +537,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
if mask_strategy is None:
|
||||
mask_strategy = [[None] * len(self.heads_num)
|
||||
for _ in range(len(self.double_blocks) + len(self.single_blocks))]
|
||||
if STA_param is None:
|
||||
STA_param = [[None] * len(self.heads_num) for _ in range(len(self.double_blocks) + len(self.single_blocks))]
|
||||
img = x = hidden_states
|
||||
text_mask = encoder_attention_mask
|
||||
t = timestep
|
||||
@@ -567,10 +580,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
|
||||
mask_search_result_save = []
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
|
||||
img, txt = block(*double_block_args)
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, STA_param[index]]
|
||||
img, txt, layer_loss_save = block(*double_block_args)
|
||||
mask_search_result_save.append(layer_loss_save)
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
if output_features:
|
||||
@@ -583,9 +597,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
mask_strategy[index + len(self.double_blocks)],
|
||||
STA_param[index + len(self.double_blocks)],
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
x, layer_loss_save = block(*single_block_args)
|
||||
mask_search_result_save.append(layer_loss_save)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
features_list.append(x[:, :img_seq_len, ...])
|
||||
|
||||
@@ -600,7 +615,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
features_list = torch.stack(features_list, dim=0)
|
||||
else:
|
||||
features_list = None
|
||||
return (img, features_list)
|
||||
return (img, features_list, mask_search_result_save)
|
||||
|
||||
def unpatchify(self, x, t, h, w):
|
||||
"""
|
||||
|
||||
@@ -61,6 +61,7 @@ def main(args):
|
||||
flow_shift=args.flow_shift,
|
||||
batch_size=args.batch_size,
|
||||
embedded_guidance_scale=args.embedded_cfg_scale,
|
||||
STA_mode=args.STA_mode,
|
||||
)
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
@@ -202,6 +203,10 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--text-states-dim-2", type=int, default=768)
|
||||
parser.add_argument("--tokenizer-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-len-2", type=int, default=77)
|
||||
parser.add_argument("--STA_mode",
|
||||
type=str,
|
||||
default="STA_inference",
|
||||
help="STA_modes should be one of ['STA_searching', 'STA_tuning', 'STA_inference']")
|
||||
|
||||
args = parser.parse_args()
|
||||
# process for vae sequence parallel
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def configure_sta(mode='STA_searching', **kwargs):
|
||||
"""
|
||||
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
|
||||
|
||||
Parameters:
|
||||
----------
|
||||
mode : str
|
||||
The STA mode to use. Options are:
|
||||
- 'STA_searching': Generate a set of mask candidates for initial search
|
||||
- 'STA_tuning': Select best mask strategy based on previously saved results
|
||||
- 'STA_inference': Load and use a previously tuned mask strategy
|
||||
|
||||
**kwargs : dict
|
||||
Mode-specific parameters:
|
||||
|
||||
For 'STA_searching':
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
|
||||
For 'STA_tuning':
|
||||
- mask_search_files_path: str, required, path to mask search results
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
- skip_time_steps: int, optional, number of time steps to use full attention (default 15)
|
||||
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
|
||||
|
||||
For 'STA_inference':
|
||||
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
|
||||
|
||||
Returns:
|
||||
-------
|
||||
list
|
||||
The configured STA parameter (STA_param) for the specified mode
|
||||
"""
|
||||
valid_modes = ['STA_searching', 'STA_tuning', 'STA_inference']
|
||||
if mode not in valid_modes:
|
||||
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
|
||||
|
||||
if mode == 'STA_searching':
|
||||
# Get parameters with defaults
|
||||
mask_candidates = kwargs.get('mask_candidates', ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1", "5,6,10"])
|
||||
mask_selected = kwargs.get('mask_selected', list(range(len(mask_candidates))))
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks = []
|
||||
for index in mask_selected:
|
||||
mask = mask_candidates[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Create 3D mask structure with fixed dimensions (t=50, l=60)
|
||||
masks_3d = []
|
||||
for i in range(50): # Fixed t dimension = 50
|
||||
row = []
|
||||
for j in range(60): # Fixed l dimension = 60
|
||||
row.append(selected_masks) # Add all masks at each position
|
||||
masks_3d.append(row)
|
||||
|
||||
return masks_3d
|
||||
|
||||
elif mode == 'STA_tuning':
|
||||
# Get required parameters
|
||||
mask_search_files_path = kwargs.get('mask_search_files_path')
|
||||
if not mask_search_files_path:
|
||||
raise ValueError("mask_search_files_path is required for STA_tuning mode")
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates = kwargs.get('mask_candidates', ["1,6,10", "3,3,5", "5,1,10", "5,3,3", "5,6,1"])
|
||||
mask_selected = kwargs.get('mask_selected', list(range(len(mask_candidates))))
|
||||
skip_time_steps = kwargs.get('skip_time_steps', 15)
|
||||
save_dir = kwargs.get('save_dir', "output/mask_strategy")
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks = []
|
||||
for index in mask_selected:
|
||||
mask = mask_candidates[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Read JSON results
|
||||
results = read_specific_json_files(mask_search_files_path)
|
||||
averaged_results = average_head_losses(results, selected_masks)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask = kwargs.get('full_attention_mask', [5, 6, 10])
|
||||
selected_masks.append(full_attention_mask)
|
||||
|
||||
# Select best mask strategy
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(averaged_results, selected_masks,
|
||||
skip_time_steps)
|
||||
|
||||
# Save mask strategy
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
file_path = os.path.join(save_dir, 'mask_strategy.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = 50 * 60 * 24 # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)")
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
else: # STA_inference
|
||||
# Get parameters with defaults
|
||||
load_path = kwargs.get('load_path', os.path.join("mask_candidates", 'mask_strategy.json'))
|
||||
|
||||
# Load previously saved mask strategy
|
||||
with open(load_path, 'r') as f:
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
|
||||
# Helper functions
|
||||
|
||||
|
||||
def read_specific_json_files(folder_path):
|
||||
"""Read and parse JSON files containing mask search results."""
|
||||
json_contents = []
|
||||
|
||||
# List files only in the current directory (no walk)
|
||||
files = os.listdir(folder_path)
|
||||
# Filter files
|
||||
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
|
||||
print(f"Found {len(matching_files)} matching files: {matching_files}")
|
||||
|
||||
for file_name in matching_files:
|
||||
file_path = os.path.join(folder_path, file_name)
|
||||
with open(file_path, 'r') as file:
|
||||
data = json.load(file)
|
||||
json_contents.append(data)
|
||||
|
||||
return json_contents
|
||||
|
||||
|
||||
def average_head_losses(results, selected_masks):
|
||||
"""Average losses across all prompts for each mask strategy."""
|
||||
# Initialize a dictionary to store the averaged results
|
||||
averaged_losses = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get all loss types (e.g., 'L2_loss')
|
||||
averaged_losses[loss_type] = {}
|
||||
|
||||
for mask in selected_masks:
|
||||
mask_str = str(mask)
|
||||
data_shape = np.array(results[0][loss_type][mask_str]).shape
|
||||
accumulated_data = np.zeros(data_shape)
|
||||
|
||||
# Sum across all prompts
|
||||
for prompt_result in results:
|
||||
accumulated_data += np.array(prompt_result[loss_type][mask_str])
|
||||
|
||||
# Average by dividing by number of prompts
|
||||
averaged_data = accumulated_data / len(results)
|
||||
averaged_losses[loss_type][mask_str] = averaged_data
|
||||
|
||||
return averaged_losses
|
||||
|
||||
|
||||
def select_best_mask_strategy(averaged_results, selected_masks, skip_time_steps=15):
|
||||
"""Select the best mask strategy for each head based on loss minimization."""
|
||||
best_mask_strategy = {}
|
||||
loss_type = 'L2_loss'
|
||||
|
||||
# Get the shape of time steps and layers
|
||||
time_steps = len(averaged_results[loss_type][str(selected_masks[0])])
|
||||
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
|
||||
|
||||
# Counter for sparsity calculation
|
||||
total_tokens = 0 # total number of masked tokens
|
||||
total_length = 0 # total sequence length
|
||||
|
||||
strategy_counts = {str(strategy): 0 for strategy in selected_masks}
|
||||
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
|
||||
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
|
||||
|
||||
for t in range(time_steps):
|
||||
for l in range(layers):
|
||||
for h in range(24):
|
||||
if t < skip_time_steps: # First steps use full attention
|
||||
strategy = full_attn_strategy
|
||||
else:
|
||||
# Get losses for this head across all strategies
|
||||
head_losses = []
|
||||
for strategy in selected_masks[:-1]: # Exclude full attention
|
||||
head_losses.append(averaged_results[loss_type][str(strategy)][t][l][h])
|
||||
|
||||
# Find which strategy gives minimum loss
|
||||
best_strategy_idx = np.argmin(head_losses)
|
||||
strategy = selected_masks[best_strategy_idx]
|
||||
|
||||
best_mask_strategy[f'{t}_{l}_{h}'] = strategy
|
||||
|
||||
# Calculate sparsity
|
||||
nums = strategy # strategy is already a list of numbers
|
||||
total_tokens += nums[0] * nums[1] * nums[2] # masked tokens for chosen strategy
|
||||
total_length += 300 # total length always 5*6*10=300
|
||||
|
||||
# Count strategy usage
|
||||
strategy_counts[str(strategy)] += 1
|
||||
|
||||
overall_sparsity = 1 - total_tokens / total_length
|
||||
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy):
|
||||
"""Convert a mask strategy dictionary to a 3D list structure with fixed dimensions."""
|
||||
# Fixed dimensions for t, l, h (50, 60, 24)
|
||||
result = [[[None for _ in range(24)] for _ in range(60)] for _ in range(50)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(mask_search_final_result,
|
||||
prompt,
|
||||
mask_strategies,
|
||||
output_dir='output/mask_search_result/'):
|
||||
if not mask_search_final_result:
|
||||
print("No mask search results to save")
|
||||
return None
|
||||
|
||||
# Create result dictionary with defaultdict for nested lists
|
||||
mask_search_dict = {"L2_loss": defaultdict(list), "L1_loss": defaultdict(list)}
|
||||
|
||||
mask_selected = list(range(len(mask_strategies)))
|
||||
selected_masks = []
|
||||
for index in mask_selected:
|
||||
mask = mask_strategies[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Process each mask strategy
|
||||
for i, mask_strategy in enumerate(selected_masks):
|
||||
mask_strategy = str(mask_strategy)
|
||||
# Process L2 loss
|
||||
step_results = []
|
||||
for step_data in mask_search_final_result:
|
||||
layer_losses = [layer_data["L2_loss"][i] for layer_data in step_data]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L2_loss"][mask_strategy] = step_results
|
||||
|
||||
step_results = []
|
||||
for step_data in mask_search_final_result:
|
||||
layer_losses = [layer_data["L1_loss"][i] for layer_data in step_data]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L1_loss"][mask_strategy] = step_results
|
||||
|
||||
# Create the output directory if it doesn't exist
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# Create a filename based on the first 20 characters of the prompt
|
||||
filename = prompt[0][:20].replace(" ", "_")
|
||||
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
|
||||
|
||||
# Save the results to a JSON file
|
||||
with open(filepath, 'w') as f:
|
||||
json.dump(mask_search_dict, f, indent=4)
|
||||
|
||||
print(f"Successfully saved mask research results to {filepath}")
|
||||
|
||||
return filepath
|
||||
+1
-1
@@ -66,7 +66,7 @@ skip ="./data,./wandb,./csrc/sliding_tile_attention/tk"
|
||||
"fastvideo/models/stepvideo/__init__.py" = ["F403"]
|
||||
"fastvideo/models/stepvideo/utils/__init__.py" = ["F403"]
|
||||
# Ignore all files that end in `_test.py`.
|
||||
"fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py" = ["E741"]
|
||||
"fastvideo/utils/STA_configuration.py" = ["E741"]
|
||||
|
||||
|
||||
[tool.yapf]
|
||||
|
||||
@@ -38,3 +38,24 @@ torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
--model_path $MODEL_BASE \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt \
|
||||
--vae-sp
|
||||
|
||||
# Mask search for STA
|
||||
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 768 \
|
||||
--width 1280 \
|
||||
--num_frames 117 \
|
||||
--num_inference_steps 1 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 7 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 1024 \
|
||||
--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 \
|
||||
--STA_mode STA_searching \
|
||||
Reference in New Issue
Block a user