Compare commits

...
10 Commits
9 changed files with 447 additions and 48 deletions
+26
View File
@@ -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:
+2 -2
View File
@@ -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
+37 -10
View File
@@ -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
+33 -18
View File
@@ -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):
"""
+5
View File
@@ -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
+284
View File
@@ -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
View File
@@ -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]
+21
View File
@@ -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 \