taylorseer testing

This commit is contained in:
kijai
2025-03-30 18:55:19 +03:00
parent 1b4164f5b3
commit 7aeb82499c
18 changed files with 1166 additions and 183 deletions
@@ -454,6 +454,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
image_cond_latents: Optional[torch.Tensor] = None,
riflex_freq_index: Optional[int] = None,
i2v_stability=True,
taylorseer: Optional[dict] = None,
**kwargs,
):
r"""
@@ -689,9 +690,17 @@ class HunyuanVideoPipeline(DiffusionPipeline):
callback = prepare_callback(self.comfy_model, num_inference_steps)
#print(self.scheduler.sigmas)
tseercache_dict, tseer_current = None, None
if taylorseer:
print(taylorseer)
from ...modules.cache_functions import cache_init
tseercache_dict, tseer_current = cache_init(self._num_timesteps, cache_device=taylorseer["cache_device"], compute_device=taylorseer["compute_device"])
tseercache_dict["max_order"] = taylorseer["max_order"]
tseercache_dict["fresh_threshold"] = taylorseer["fresh_threshold"]
logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
comfy_pbar = ProgressBar(len(timesteps))
with self.progress_bar(total=len(timesteps)) as progress_bar:
for i, t in enumerate(timesteps):
@@ -709,6 +718,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_enabled = False
current_step_percentage = i / len(timesteps)
if taylorseer:
tseer_current['step'] = i
if self.do_spatio_temporal_guidance:
if stg_start_percent <= current_step_percentage <= stg_end_percent:
stg_enabled = True
@@ -851,6 +864,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx,
stg_mode=stg_mode,
return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer
tseer_current = tseer_current, #taylorseer
)["x"]
else:
uncond = self.transformer(
@@ -865,6 +880,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx,
stg_mode=stg_mode,
return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer
tseer_current = tseer_current, #taylorseer
)["x"]
cond = self.transformer(
latent_model_input[1].unsqueeze(0),
@@ -878,6 +895,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx,
stg_mode=stg_mode,
return_dict=True,
tseercache_dict = tseercache_dict, #taylorseer
tseer_current = tseer_current, #taylorseer
)["x"]
# perform guidance
@@ -0,0 +1,12 @@
from .cache_cutfresh import cache_cutfresh
from .fresh_ratio_scheduler import fresh_ratio_scheduler
from .score_evaluate import score_evaluate
from .global_force_fresh import global_force_fresh
from .cache_cutfresh import cache_cutfresh
from .update_cache import update_cache
from .force_init import force_init
from .attention import cached_attention_forward
from .cache_init import cache_init
from .cal_type import cal_type
from .force_scheduler import force_scheduler
from .support_set_selection import support_set_selection
@@ -0,0 +1,31 @@
# Besides, re-arrange the attention module
from torch.jit import Final
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Union
#from xformers.ops.fmha.attn_bias import BlockDiagonalMask
def cached_attention_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
#attn_bias: Optional[Union[torch.Tensor, BlockDiagonalMask]] = None,
attn_bias,
p: float = 0.0,
scale: Optional[float] = None
) -> torch.Tensor:
scale = 1.0 / query.shape[-1] ** 0.5
query = query * scale
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn = query @ key.transpose(-2, -1)
if attn_bias is not None:
attn_bias = attn_bias.materialize(shape= attn.shape, dtype= attn.dtype, device= attn.device)
attn = attn + attn_bias
#out_map = attn
attn_map = attn.softmax(-1)
attn = F.dropout(attn_map, p)
attn = attn @ value
return attn.transpose(1, 2).contiguous(), attn_map.mean(dim=1)
@@ -0,0 +1,75 @@
from .fresh_ratio_scheduler import fresh_ratio_scheduler
from .score_evaluate import score_evaluate
#from .token_merge import token_merge
from .support_set_selection import support_set_selection
import torch
def cache_cutfresh(cache_dic, tokens, current):
'''
Cut fresh tokens from the input tokens and update the cache counter.
cache_dic: dict, the cache dictionary containing cache(main extra memory cost), indices and some other information.
tokens: torch.Tensor, the input tokens to be cut.
current: dict, the current step, layer, and module information. Particularly convenient for debugging.
'''
step = current['step']
layer = current['layer']
stream = current['stream']
module = current['module']
fresh_ratio = fresh_ratio_scheduler(cache_dic, current)
fresh_ratio = torch.clamp(torch.tensor(fresh_ratio, device = tokens.device), min=0, max=1)
# Generate the index tensor for fresh tokens
score = score_evaluate(cache_dic, tokens, current) # s1, s2, s3 mentioned in the paper
#score = local_selection_with_bonus(score, 0.4, 4) # Uniform Spatial Distribution s4 mentioned in the paper
indices = score.argsort(dim=-1, descending=True)
topk = int(fresh_ratio * score.shape[1])
fresh_indices = indices[:, :topk]
stale_indices = indices[:, topk:]
#fresh_indices = support_set_selection(tokens, fresh_ratio, 0.4, current, cache_dic) # (B, fresh_ratio * N) # 0.4
# (B, fresh_ratio *N)
# Updating the Cache Frequency Score s3 mentioned in the paper
# stale tokens index + 1 in each ***module***, fresh tokens index = 0
###cache_dic['cache_index'][-1][layer][module] += 1
###cache_dic['cache_index'][-1][layer][module].scatter_(dim=1, index=fresh_indices,
### src = torch.zeros_like(fresh_indices, dtype=torch.int, device=fresh_indices.device))
#cache_dic['cache_index']['layer_index'][module] += 1
#cache_dic['cache_index']['layer_index'][module].scatter_(dim=1, index=fresh_indices,
# src = torch.zeros_like(fresh_indices, dtype=torch.int, device=fresh_indices.device))
fresh_indices_expand = fresh_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1])
fresh_tokens = torch.gather(input = tokens, dim = 1, index = fresh_indices_expand)
return fresh_indices, fresh_tokens
def local_selection_with_bonus(score, bonus_ratio, grid_size=2):
batch_size, num_tokens = score.shape
image_size = int(num_tokens ** 0.5)
block_size = grid_size * grid_size
assert num_tokens % block_size == 0, "The number of tokens must be divisible by the block size."
# Step 1: Reshape score to group it by blocks
score_reshaped = score.view(batch_size, image_size // grid_size, grid_size, image_size // grid_size, grid_size)
score_reshaped = score_reshaped.permute(0, 1, 3, 2, 4).contiguous()
score_reshaped = score_reshaped.view(batch_size, -1, block_size) # [batch_size, num_blocks, block_size]
# Step 2: Find the max token in each block
max_scores, max_indices = score_reshaped.max(dim=-1, keepdim=True) # [batch_size, num_blocks, 1]
# Step 3: Create a mask to identify max score tokens
mask = torch.zeros_like(score_reshaped)
mask.scatter_(-1, max_indices, 1) # Set mask to 1 at the max indices
# Step 4: Apply the bonus only to the max score tokens
score_reshaped = score_reshaped + (mask * max_scores * bonus_ratio) # Apply bonus only to max tokens
# Step 5: Reshape the score back to its original shape
score_modified = score_reshaped.view(batch_size, image_size // grid_size, image_size // grid_size, grid_size, grid_size)
score_modified = score_modified.permute(0, 1, 3, 2, 4).contiguous()
score_modified = score_modified.view(batch_size, num_tokens)
return score_modified
@@ -0,0 +1,124 @@
import torch
def cache_init(num_steps, model_kwargs=None, cache_device=torch.device("cpu"), compute_device=torch.device("cuda")):
'''
Initialization for cache.
'''
cache_dic = {}
cache = {}
cache_index = {}
cache[-1]={}
cache_index[-1]={}
cache_index['layer_index']={}
cache_dic['attn_map'] = {}
cache_dic['attn_map'][-1] = {}
cache_dic['attn_map'][-1]['double_stream'] = {}
cache_dic['attn_map'][-1]['single_stream'] = {}
cache_dic['k-norm'] = {}
cache_dic['k-norm'][-1] = {}
cache_dic['k-norm'][-1]['double_stream'] = {}
cache_dic['k-norm'][-1]['single_stream'] = {}
cache_dic['v-norm'] = {}
cache_dic['v-norm'][-1] = {}
cache_dic['v-norm'][-1]['double_stream'] = {}
cache_dic['v-norm'][-1]['single_stream'] = {}
cache_dic['cross_attn_map'] = {}
cache_dic['cross_attn_map'][-1] = {}
cache[-1]['double_stream']={}
cache[-1]['single_stream']={}
cache_dic['cache_counter'] = 0
cache_dic['cache_device'] = cache_device
cache_dic['compute_device'] = compute_device
for j in range(20):
cache[-1]['double_stream'][j] = {}
cache_index[-1][j] = {}
cache_dic['attn_map'][-1]['double_stream'][j] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['total'] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['attn_map'][-1]['double_stream'][j]['img_mlp'] = {}
cache_dic['k-norm'][-1]['double_stream'][j] = {}
cache_dic['k-norm'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['k-norm'][-1]['double_stream'][j]['img_mlp'] = {}
cache_dic['v-norm'][-1]['double_stream'][j] = {}
cache_dic['v-norm'][-1]['double_stream'][j]['txt_mlp'] = {}
cache_dic['v-norm'][-1]['double_stream'][j]['img_mlp'] = {}
for j in range(40):
cache[-1]['single_stream'][j] = {}
cache_index[-1][j] = {}
cache_dic['attn_map'][-1]['single_stream'][j] = {}
cache_dic['attn_map'][-1]['single_stream'][j]['total'] = {}
cache_dic['k-norm'][-1]['single_stream'][j] = {}
cache_dic['k-norm'][-1]['single_stream'][j]['total'] = {}
cache_dic['v-norm'][-1]['single_stream'][j] = {}
cache_dic['v-norm'][-1]['single_stream'][j]['total'] = {}
cache_dic['taylor_cache'] = False
cache_dic['duca'] = False
cache_dic['test_FLOPs'] = False
mode = 'Taylor'
if mode == 'original':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.0
cache_dic['fresh_threshold'] = 1
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
elif mode == 'ToCa':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.10
cache_dic['fresh_threshold'] = 5
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
cache_dic['duca'] = False
elif mode == 'DuCa':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.10
cache_dic['fresh_threshold'] = 5
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['max_order'] = 0
cache_dic['first_enhance'] = 1
cache_dic['duca'] = True
elif mode == 'Taylor':
cache_dic['cache_type'] = 'random'
cache_dic['cache_index'] = cache_index
cache_dic['cache'] = cache
cache_dic['fresh_ratio_schedule'] = 'ToCa'
cache_dic['fresh_ratio'] = 0.0
cache_dic['fresh_threshold'] = 5
cache_dic['max_order'] = 1
cache_dic['force_fresh'] = 'global'
cache_dic['soft_fresh_weight'] = 0.0
cache_dic['taylor_cache'] = True
cache_dic['first_enhance'] = 1
current = {}
current['num_steps'] = num_steps
current['activated_steps'] = [0]
return cache_dic, current
@@ -0,0 +1,49 @@
from .force_scheduler import force_scheduler
def cal_type(cache_dic, current):
'''
Determine calculation type for this step
'''
if (cache_dic['fresh_ratio'] == 0.0) and (not cache_dic['taylor_cache']):
# FORA:Uniform
first_step = (current['step'] == 0)
else:
# ToCa: First enhanced
first_step = (current['step'] < cache_dic['first_enhance'])
#first_step = (current['step'] <= 3)
force_fresh = cache_dic['force_fresh']
if not first_step:
fresh_interval = cache_dic['cal_threshold']
else:
fresh_interval = cache_dic['fresh_threshold']
if (first_step) or (cache_dic['cache_counter'] == fresh_interval - 1 ):
current['type'] = 'full'
cache_dic['cache_counter'] = 0
current['activated_steps'].append(current['step'])
#current['activated_times'].append(current['t'])
force_scheduler(cache_dic, current)
elif (cache_dic['taylor_cache']):
cache_dic['cache_counter'] += 1
current['type'] = 'taylor_cache'
else:
cache_dic['cache_counter'] += 1
if (cache_dic['duca']):
if (cache_dic['cache_counter'] % 2 == 1): # 0: ToCa-Aggresive-ToCa, 1: Aggresive-ToCa-Aggresive
current['type'] = 'ToCa'
# 'cache_noise' 'ToCa' 'FORA'
else:
current['type'] = 'aggressive'
else:
current['type'] = 'ToCa'
#if current['step'] < 25:
# current['type'] = 'FORA'
#else:
# current['type'] = 'aggressive'
######################################################################
#if (current['step'] in [3,2,1,0]):
# current['type'] = 'full'
@@ -0,0 +1,10 @@
import torch
def force_init(cache_dic, current, tokens):
'''
Initialization for Force Activation step.
'''
cache_dic['cache_index'][-1][current['layer']][current['module']] = torch.zeros(tokens.shape[0], tokens.shape[1], dtype=torch.int, device=tokens.device)
#if current['layer'] == 0:
# cache_dic['cache_index']['layer_index'][current['module']] = torch.zeros(tokens.shape[0], tokens.shape[1], dtype=torch.int, device=tokens.device)
@@ -0,0 +1,16 @@
import torch
def force_scheduler(cache_dic, current):
if cache_dic['fresh_ratio'] == 0:
# FORA
linear_step_weight = 0.0
else:
# TokenCache
linear_step_weight = 0.0
step_factor = torch.tensor(1 - linear_step_weight + 2 * linear_step_weight * current['step'] / current['num_steps'])
threshold = torch.round(cache_dic['fresh_threshold'] / step_factor)
# no force constrain for sensitive steps, cause the performance is good enough.
# you may have a try.
cache_dic['cal_threshold'] = threshold
#return threshold
@@ -0,0 +1,59 @@
import torch
def fresh_ratio_scheduler(cache_dic, current):
'''
Return the fresh ratio for the current step.
'''
fresh_ratio = cache_dic['fresh_ratio']
fresh_ratio_schedule = cache_dic['fresh_ratio_schedule']
step = current['step']
num_steps = current['num_steps']
threshold = cache_dic['fresh_threshold']
weight = 0.9
if fresh_ratio_schedule == 'constant':
return fresh_ratio
elif fresh_ratio_schedule == 'linear':
return fresh_ratio * (1 + weight - 2 * weight * step / num_steps)
elif fresh_ratio_schedule == 'exp':
#return 0.5 * (0.052 ** (step/num_steps))
return fresh_ratio * (weight ** (step / num_steps))
elif fresh_ratio_schedule == 'linear-mode':
mode = (step % threshold)/threshold - 0.5
mode_weight = 0.1
return fresh_ratio * (1 + weight - 2 * weight * step / num_steps + mode_weight * mode)
elif fresh_ratio_schedule == 'layerwise':
return fresh_ratio * (1 + weight - 2 * weight * current['layer'] / 27)
elif fresh_ratio_schedule == 'linear-layerwise':
step_weight = -0.9 #0.9
step_factor = 1 - step_weight + 2 * step_weight * step / num_steps
#if current['layer'] == 2:
# return 1.0
#sigmoid
#sigmoid_weight = 0.13
#layer_factor = 2 * torch.sigmoid(torch.tensor([sigmoid_weight * (13.5 - current['layer'])]))
layer_weight = 0.6
layer_factor = 1 + layer_weight - 2 * layer_weight * current['layer'] / 27
module_weight = 1.0 #TokenCache N=8 2.5 N=6 2.5 #N=4 2.1
module_time_weight = 0.6
module_factor = (1 - (1-module_time_weight) * module_weight) if current['module']=='cross-attn' else (1 + module_time_weight * module_weight)
return fresh_ratio * layer_factor * step_factor * module_factor
elif fresh_ratio_schedule == 'ToCa':
step_weight = 0.0 #0.9
step_factor = 1 - step_weight + 2 * step_weight * step / num_steps
layer_weight = 0.5
layer_factor = 1 + layer_weight - 2 * layer_weight * current['layer'] / 27
#module_weight = 1.0
#module_time_weight = 0.6
# this means 60*x% cross-attn computation, and 160*x% mlp computation. This is designed for cross-attn has best temporal redundancy, and mlp has worse.
# so cross-attn compute less and mlp compute more.
#module_factor = (1 - (1-module_time_weight) * module_weight) if current['module']=='cross-attn' else (1 + module_time_weight * module_weight)
stream_weight = 0.6
stream_factor = (1 - stream_weight) if current['stream']=='double_stream' else (1 + stream_weight)
return fresh_ratio * layer_factor * step_factor * stream_factor #* module_factor
else:
raise ValueError("unrecognized fresh ratio schedule", fresh_ratio_schedule)
@@ -0,0 +1,21 @@
from .force_scheduler import force_scheduler
def global_force_fresh(cache_dic, current):
'''
Return whether to force fresh tokens globally.
'''
first_step = (current['step'] == 0)
second_step = (current['step'] == 1)
force_fresh = cache_dic['force_fresh']
if not first_step:
fresh_threshold = cache_dic['cal_threshold']
else:
fresh_threshold = cache_dic['fresh_threshold']
if force_fresh == 'global':
return (first_step or (current['step']% fresh_threshold == 0))
elif force_fresh == 'local':
return first_step
elif force_fresh == 'none':
return first_step
else:
raise ValueError("unrecognized force fresh strategy", force_fresh)
@@ -0,0 +1,60 @@
import torch
import torch.nn as nn
from .scores import attn_score, similarity_score, norm_score, k_norm_score, v_norm_score
def score_evaluate(cache_dic, tokens, current) -> torch.Tensor:
'''
Return the score tensor (B, N) for the given tokens.
'''
#if ((not current['is_force_fresh']) and (cache_dic['force_fresh'] == 'local')):
# # abandoned branch, if you want to explore the local force fresh strategy, this may help.
# force_fresh_mask = torch.as_tensor((cache_dic['cache_index'][-1][current['layer']][current['module']] >= 2 * cache_dic['fresh_threshold']), dtype = int) # 2 because the threshold is for step, not module
# force_len = force_fresh_mask.sum(dim=1)
# force_indices = force_fresh_mask.argsort(dim = -1, descending = True)[:, :force_len.min()]
# force_indices = force_indices[:, torch.randperm(force_indices.shape[1])]
# Just see more explanation in the version of DiT-ToCa if needed.
if cache_dic['cache_type'] == 'random':
score = torch.rand(tokens.shape[0], tokens.shape[1], device=tokens.device)
elif cache_dic['cache_type'] == 'straight':
score = torch.ones(tokens.shape[0], tokens.shape[1]).to(tokens.device)
elif cache_dic['cache_type'] == 'attention':
# cache_dic['attn_map'][step][layer] (B, N, N), the last dimention has get softmaxed
score = attn_score(cache_dic, current)
#score = score + 0.0 * torch.rand_like(score, device= score.device)
elif cache_dic['cache_type'] == 'similarity':
score = similarity_score(cache_dic, current, tokens)
elif cache_dic['cache_type'] == 'norm':
score = norm_score(cache_dic, current, tokens)
elif cache_dic['cache_type'] == 'k-norm':
score = k_norm_score(cache_dic, current)
elif cache_dic['cache_type'] == 'v-norm':
score = v_norm_score(cache_dic, current)
elif cache_dic['cache_type'] == 'compress':
score1 = torch.rand(int(tokens.shape[0]*0.5), tokens.shape[1])
score1 = torch.cat([score1, score1], dim=0).to(tokens.device)
score2 = cache_dic['attn_map'][-1][current['layer']].sum(dim=1)#.mean(dim=0) # (B, N)
# normalize
score2 = score2 / score2.max(dim=1, keepdim=True)[0]
score = 0.5 * score1 + 0.5 * score2
# abandoned the branch, if you want to explore the local force fresh strategy, this may help.
#if ((not current['is_force_fresh']) and (cache_dic['force_fresh'] == 'local')): # current['is_force_fresh'] is False, cause when it is True, no cut and fresh are needed
# #print(torch.ones_like(force_indices, dtype=float, device=force_indices.device).dtype)
# score.scatter_(dim=1, index=force_indices, src=torch.ones_like(force_indices, dtype=torch.float32,
# device=force_indices.device))
###if (True and (cache_dic['force_fresh'] == 'global')):
### soft_step_score = cache_dic['cache_index'][-1][current['layer']][current['module']].float() / (cache_dic['fresh_threshold'])
### #soft_layer_score = cache_dic['cache_index']['layer_index'][current['module']].float() / (27)
### score = score + cache_dic['soft_fresh_weight'] * soft_step_score #+ 0.1 *soft_layer_score
return score.to(tokens.device)
+77
View File
@@ -0,0 +1,77 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
def attn_score(cache_dic, current):
#self_attn_score = 1- cache_dic['attn_map'][-1][current['layer']].diagonal(dim1=1, dim2=2)
#self_attn_score = F.normalize(self_attn_score, dim=1, p=2)
#attention_score = F.normalize(cache_dic['attn_map'][-1][current['layer']].sum(dim=1), dim=1, p=2)
#cross_attn_map = F.threshold(cache_dic['cross_attn_map'][-1][current['layer']],threshold=0.0, value=0.0)
#cross_attention_score = F.normalize(cross_attn_map.sum(dim=-1), dim=-1, p=2)
# Note: It is important to give a same selection method for cfg and no cfg.
# Because the influence of **Cross-Attention** in text-contidional models makes cfg and no cfg a BIG difference.
# Same selection for cfg and no cfg
#cond_cmap, uncond_cmap = torch.split(cache_dic['attn_map'][-1][current['layer']], len(cache_dic['cross_attn_map'][-1][current['layer']]) // 2, dim=0)
#cond_weight = 0.5
#cmap = cond_weight * cond_cmap + (1 - cond_weight) * uncond_cmap
## Entropy score
#cross_attention_entropy = -torch.sum(cmap * torch.log(cmap + 1e-7), dim=-1)
#cross_attention_score = F.normalize(1 + cross_attention_entropy, dim=1, p=2) # Note here "1" does not influence the sorted sequence, but provie stability.
#score = cross_attention_score.repeat(2, 1)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['attn_map'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['attn_map'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
# You can try conbining the self_attention_score (s1) and cross_attention_score (s2) as the final score, there exists a balance.
#cross_weight = 0.0
#score = (1-cross_weight) * attention_score + cross_weight * cross_attention_score
return score
def similarity_score(cache_dic, current, tokens):
cosine_sim = F.cosine_similarity(tokens, cache_dic['cache'][-1][current['layer']][current['module']], dim=-1)
return F.normalize(1- cosine_sim, dim=-1, p=2)
def norm_score(cache_dic, current, tokens):
norm = tokens.norm(dim=-1, p=2)
return F.normalize(norm, dim=-1, p=2)
def kv_norm_score(cache_dic, current):
# (B, N, num_heads)
#cond_k_norm, uncond_k_norm = torch.split(cache_dic['cache'][-1][current['layer']]['k_norm'], len(cache_dic['cache'][-1][current['layer']]['k_norm']) // 2, dim=0)
cond_v_norm, uncond_v_norm = torch.split(cache_dic['cache'][-1][current['layer']]['v_norm'], len(cache_dic['cache'][-1][current['layer']]['v_norm']) // 2, dim=0)
cond_weight = 0.5
#k_norm = cond_weight * cond_k_norm + (1 - cond_weight) * uncond_k_norm
v_norm = cond_weight * cond_v_norm + (1 - cond_weight) * uncond_v_norm
kv_norm = 1 -v_norm
## 计算 (B/2, N) 张量在 N 维度上的每个元素与均值的绝对值差
#kv_norm_mean = kv_norm.mean(dim=-2, keepdim=True)
#kv_norm_diff = torch.abs(kv_norm - kv_norm_mean)
return F.normalize(kv_norm.sum(dim=-1), p=2).repeat(2, 1)
def k_norm_score(cache_dic, current):
# (B, N)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['k-norm'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['k-norm'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
return score
def v_norm_score(cache_dic, current):
# (B, N)
if current['stream'] == 'double_stream':
score = F.normalize(cache_dic['v-norm'][-1][current['stream']][current['layer']][current['module']], dim=-1, p=2)
elif current['stream'] == 'single_stream':
score = F.normalize(cache_dic['v-norm'][-1][current['stream']][current['layer']]['total'], dim=-1, p=2)
return score
@@ -0,0 +1,52 @@
import torch
from typing import Dict
def support_set_selection(x: torch.Tensor, fresh_ratio: float, base_ratio: float, current: Dict, cache_dic: Dict) -> torch.Tensor:
#selection_start = 0
#
#if current['stream'] == 'single_stream':
# # only select from the img tokens
# x = x[:, cache_dic['txt_shape'] :]
# selection_start = cache_dic['txt_shape']
B, N, H = x.shape
num_total = int(fresh_ratio * N) # 最终每个 batch 选取的 token 数
base_count = int(base_ratio * num_total) # 随机选取的 token 数
#base_count = 1
add_count = num_total - base_count # 需要从候选集中选取的 token 数
# 1. 随机选取 (B, base_count) 个 token
random_indices = torch.randperm(N, device=x.device)
base_indices = random_indices[:base_count]
other_indices = random_indices[base_count:]
base_tokens = x.gather(dim=1, index=base_indices.unsqueeze(-1).expand(B, -1, H))
#other_tokens = x.gather(dim=1, index=other_indices.unsqueeze(-1).expand(-1, -1, H))
# 2. 计算余下 token 与已选 token 的相似度
# normaize
base_tokens = base_tokens / base_tokens.norm(dim=-1, keepdim=True)
#other_tokens = other_tokens / other_tokens.norm(dim=-1, keepdim=True)
x_norm = x / x.norm(dim=-1, keepdim=True)
# 计算余下 token 与已选 token 的相似度
similarity = torch.einsum('bnd,bmd->bnm', base_tokens, x_norm)
# 计算每列最小值
min_similarity = similarity.min(dim=1).values
#min_similarity = similarity.max(dim=1).values
# 3. 选取相似度最小的 token
_, min_indices = min_similarity.topk(add_count, largest=False)
#_, min_indices = min_similarity.topk(add_count, largest=True)
# 4. 合并 base_indices 和 min_indices
#indices = torch.cat([base_indices, other_indices[min_indices]], dim=-1)
indices = torch.cat([base_indices.expand(B, -1), min_indices], dim=-1) #+ selection_start
return indices
@@ -0,0 +1,28 @@
import torch
def token_merge(cache_dic, tokens, current, fresh_indices, stale_indices):
'''
An abandoned branch in exploring if token merge helps. The answer is no, at least no for training-free strategy.
'''
if (current['layer'] % 1 == 0):
fresh_tokens = torch.gather(input = tokens, dim = 1, index = fresh_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1]))
stale_tokens = torch.gather(input = tokens, dim = 1, index = stale_indices.unsqueeze(-1).expand(-1, -1, tokens.shape[-1]))
method = 'similarity'
if method == 'distance':
descending = False
distance = torch.cdist(stale_tokens, fresh_tokens, p=1)
stale_fresh_dist, stale_fresh_indices_allstale = torch.min(distance, dim=2)
elif method == 'similarity':
descending = True
fresh_tokens = torch.nn.functional.normalize(fresh_tokens, p=2, dim=-1)
stale_tokens = torch.nn.functional.normalize(stale_tokens, p=2, dim=-1)
similarity = stale_tokens @ fresh_tokens.transpose(1, 2)
stale_fresh_dist, stale_fresh_indices_allstale = torch.max(similarity, dim=2)
saved_topk_stale = int((stale_fresh_dist > 0.995).sum(dim=1).min())
merged_stale_sequence = torch.sort(stale_fresh_dist, dim=1, descending=descending)[1][:,:saved_topk_stale]
stale_fresh_indices = stale_fresh_indices_allstale.gather(1, merged_stale_sequence)
merged_stale_sequence = stale_indices.gather(1, merged_stale_sequence)
merged_stale_fresh_indices = fresh_indices.gather(1, stale_fresh_indices)
cache_dic['merged_stale_fresh_indices'] = merged_stale_fresh_indices
cache_dic['merged_stale_sequence'] = merged_stale_sequence
@@ -0,0 +1,19 @@
import torch
def update_cache(fresh_indices, fresh_tokens, cache_dic, current, fresh_attn_map=None):
'''
Update the cache with the fresh tokens.
'''
step = current['step']
layer = current['layer']
module = current['module']
# Update the cached tokens at the positions
indices = fresh_indices
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][0].scatter_(dim=1, index=indices.unsqueeze(-1).expand(-1, -1, fresh_tokens.shape[-1]), src=fresh_tokens)
+420 -179
View File
@@ -21,6 +21,9 @@ from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled_single, is_enhance_enabled_double, set_num_frames
from .norm_layers import RMSNorm
from .cache_functions import cal_type
from .taylor_utils import derivative_approximation, taylor_formula, taylor_cache_init
from contextlib import contextmanager
@contextmanager
@@ -200,6 +203,8 @@ class MMDoubleStreamBlock(nn.Module):
token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None,
condition_type: str = None,
cache_dic: Optional[Dict] = None,
current: Optional[Dict] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if condition_type == "token_replace":
img_mod1, token_replace_img_mod1 = self.img_mod(vec, condition_type=condition_type, \
@@ -236,108 +241,245 @@ class MMDoubleStreamBlock(nn.Module):
) = self.txt_mod(vec).chunk(6, dim=-1)
# Prepare image for attention.
img_modulated = self.img_norm1(img)
if condition_type == "token_replace":
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
first_frame_token_num=first_frame_token_num
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
if is_enhance_enabled_double():
feta_scores = get_feta_scores(img_q, img_k)
# Run actual attention.
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
if is_enhance_enabled_double():
img_attn *= feta_scores
# Calculate the img bloks.
if condition_type == "token_replace":
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
if cache_dic is None:
img_modulated = self.img_norm1(img)
if condition_type == "token_replace":
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
first_frame_token_num=first_frame_token_num
)
),
gate=txt_mod2_gate,
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
return img, txt
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
if is_enhance_enabled_double():
feta_scores = get_feta_scores(img_q, img_k)
# Run actual attention.
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
if is_enhance_enabled_double():
img_attn *= feta_scores
# Calculate the img bloks.
if condition_type == "token_replace":
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
),
gate=txt_mod2_gate,
)
return img, txt
else:
if current['type'] == 'full':
current['module'] = 'attn'
img_modulated = self.img_norm1(img)
if condition_type == "token_replace":
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
first_frame_token_num=first_frame_token_num
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
if is_enhance_enabled_double():
feta_scores = get_feta_scores(img_q, img_k)
# Run actual attention.
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
v = torch.cat((img_v, txt_v), dim=1)
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=img_k.shape[0],
attn_mask=attn_mask
)
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
if is_enhance_enabled_double():
img_attn *= feta_scores
# Calculate the img blocks
current['module'] = 'img_attn'
taylor_cache_init(cache_dic, current)
if condition_type == "token_replace":
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
#img attn
img_attn_out = self.img_attn_proj(img_attn)
img = img + apply_gate(img_attn_out, gate=img_mod1_gate)
derivative_approximation(cache_dic, current, img_attn_out)
#img mlp
current['module'] = 'img_mlp'
taylor_cache_init(cache_dic, current)
img_mlp_out = self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
)
img = img + apply_gate(img_mlp_out, gate=img_mod2_gate)
derivative_approximation(cache_dic, current, img_mlp_out)
# Calculate the txt blocks
current['module'] = 'txt_attn'
taylor_cache_init(cache_dic, current)
txt_attn_out = self.txt_attn_proj(txt_attn)
txt = txt + apply_gate(txt_attn_out, gate=txt_mod1_gate)
derivative_approximation(cache_dic, current, txt_attn_out)
current['module'] = 'txt_mlp'
taylor_cache_init(cache_dic, current)
txt_mlp_out = self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
)
txt = txt + apply_gate(txt_mlp_out, gate=txt_mod2_gate)
derivative_approximation(cache_dic, current, txt_mlp_out)
elif current['type'] == 'taylor_cache':
current['module'] = 'img_attn'
img = img + apply_gate(taylor_formula(cache_dic, current), gate=img_mod1_gate)
current['module'] = 'img_mlp'
img = img + apply_gate(taylor_formula(cache_dic, current), gate=img_mod2_gate)
current['module'] = 'txt_attn'
txt = txt + apply_gate(taylor_formula(cache_dic, current), gate=txt_mod1_gate)
current['module'] = 'txt_mlp'
txt = txt + apply_gate(taylor_formula(cache_dic, current),gate=txt_mod2_gate)
return img, txt
#region single block
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers as described in
@@ -426,6 +568,8 @@ class MMSingleStreamBlock(nn.Module):
token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None,
condition_type: str = None,
cache_dic: Optional[Dict] = None,
current: Optional[Dict] = None,
stg_mode: Optional[str] = None,
) -> torch.Tensor:
@@ -441,59 +585,81 @@ class MMSingleStreamBlock(nn.Module):
tr_mod_gate) = tr_mod.chunk(3, dim=-1)
else:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
if condition_type == "token_replace":
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type,
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num)
else:
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
if cache_dic is None:
if condition_type == "token_replace":
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type,
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num)
else:
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
if is_enhance_enabled_single():
feta_scores = get_feta_scores(img_q, img_k)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
# Compute attention.
#assert (
# cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1
#), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}"
if stg_mode is not None:
if stg_mode == "STG-A":
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
do_stg=True,
txt_len=txt_len,
attn_mask=attn_mask
)
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
elif stg_mode == "STG-R":
if is_enhance_enabled_single():
feta_scores = get_feta_scores(img_q, img_k)
# Compute attention.
#assert (
# cu_seqlens_q.shape[0] == 2 * x.shape[0] + 1
#), f"cu_seqlens_q.shape:{cu_seqlens_q.shape}, x.shape[0]:{x.shape[0]}"
if stg_mode is not None:
if stg_mode == "STG-A":
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
do_stg=True,
txt_len=txt_len,
attn_mask=attn_mask
)
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
elif stg_mode == "STG-R":
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
attn_mask=attn_mask
)
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
output = apply_gate(output, gate=mod_gate)
batch_size = output.shape[0]
output[:batch_size-1, :, :] = 0
return x + output
else:
attn = attention(
q,
k,
@@ -507,32 +673,88 @@ class MMSingleStreamBlock(nn.Module):
batch_size=x.shape[0],
attn_mask=attn_mask
)
if is_enhance_enabled_single():
attn *= feta_scores
#attn[:, :-txt_len, :] *= feta_scores
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
output = apply_gate(output, gate=mod_gate)
batch_size = output.shape[0]
output[:batch_size-1, :, :] = 0
return x + output
if condition_type == "token_replace":
output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type,
tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num)
return output
else:
return x + apply_gate(output, gate=mod_gate)
else:
attn = attention(
q,
k,
v,
heads = self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
attn_mask=attn_mask
)
if is_enhance_enabled_single():
attn *= feta_scores
#attn[:, :-txt_len, :] *= feta_scores
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
if current['type'] == 'full':
#current['module'] = 'mlp'
#taylor_cache_init(cache_dic, current)
if condition_type == "token_replace":
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type,
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num)
else:
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
current['module'] = 'attn'
taylor_cache_init(cache_dic, current)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
# Apply RoPE if needed.
if freqs_cis is not None:
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope)
# assert (
# img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
# ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
q = torch.cat((img_q, txt_q), dim=1)
k = torch.cat((img_k, txt_k), dim=1)
if is_enhance_enabled_single():
feta_scores = get_feta_scores(img_q, img_k)
# Compute attention.
attn = attention(
q,
k,
v,
heads=self.heads_num,
mode=self.attention_mode,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_kv=cu_seqlens_kv,
max_seqlen_q=max_seqlen_q,
max_seqlen_kv=max_seqlen_kv,
batch_size=x.shape[0],
attn_mask=attn_mask
)
if is_enhance_enabled_single():
attn *= feta_scores
#attn[:, :-txt_len, :] *= feta_scores
derivative_approximation(cache_dic, current, attn)
current['module'] = 'total'
taylor_cache_init(cache_dic, current)
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
derivative_approximation(cache_dic, current, output)
elif current['type'] == 'taylor_cache':
current['module'] = 'total'
output = taylor_formula(cache_dic, current)
if condition_type == "token_replace":
output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type,
tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num)
@@ -950,26 +1172,32 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
stg_mode: str = None,
stg_block_idx: int = -1,
return_dict: bool = True,
tseercache_dict = None,
tseer_current = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
def _process_double_blocks(img, txt, vec, block_args):
def _process_double_blocks(img, txt, vec, block_args, tseercache_dict=None, tseer_current=None):
for b, block in enumerate(self.double_blocks):
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
block.to(self.main_device)
img, txt = block(img, txt, vec, *block_args)
if tseer_current is not None:
tseer_current['layer'] = b
img, txt = block(img, txt, vec, *block_args, tseercache_dict, tseer_current)
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True)
return img, txt
def _process_single_blocks(x, vec, txt_seq_len, block_args, stg_mode=None, stg_block_idx=None):
def _process_single_blocks(x, vec, txt_seq_len, block_args, tseercache_dict=None, tseer_current=None, stg_mode=None, stg_block_idx=None):
for b, block in enumerate(self.single_blocks):
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
block.to(self.main_device)
curr_stg_mode = stg_mode if b == stg_block_idx else None
x = block(x, vec, txt_seq_len, *block_args, curr_stg_mode)
if tseer_current is not None:
tseer_current['layer'] = b
x = block(x, vec, txt_seq_len, *block_args, tseercache_dict, tseer_current, curr_stg_mode)
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=True)
@@ -1125,11 +1353,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
img = x[:, :img_seq_len, ...]
self.previous_residual = (img - ori_img).to(self.teacache_device)
else:
# Pass through DiT blocks
img, txt = _process_double_blocks(img, txt, vec, block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
# TaylorSeer
if tseercache_dict is not None:
cal_type(tseercache_dict, tseer_current)
tseer_current['compute'] = not (tseer_current['type'] == 'aggressive')
if tseer_current['compute']:
tseer_current['stream'] = 'double_stream'
img, txt = _process_double_blocks(img, txt, vec, block_args, tseercache_dict, tseer_current)
x = torch.cat((img, txt), 1)
tseer_current['stream'] = 'single_stream'
x = _process_single_blocks(x, vec, txt.shape[1], block_args, tseercache_dict=tseercache_dict, tseer_current=tseer_current,
stg_mode=stg_mode, stg_block_idx=stg_block_idx)
else:
x = tseercache_dict['aggressive_feature']
else:
img, txt = _process_double_blocks(img, txt, vec, block_args)
x = torch.cat((img, txt), 1)
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode=stg_mode, stg_block_idx=stg_block_idx)
img = x[:, :img_seq_len, ...]
# ---------------------------- Final layer ------------------------------
+52
View File
@@ -0,0 +1,52 @@
from typing import Dict
import torch
import math
@torch.compiler.disable()
def derivative_approximation(cache_dic: Dict, current: Dict, feature: torch.Tensor):
"""
Compute derivative approximation
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
difference_distance = current['activated_steps'][-1] - current['activated_steps'][-2]
updated_taylor_factors = {}
updated_taylor_factors[0] = feature.to(cache_dic['cache_device'])
for i in range(cache_dic['max_order']):
if (cache_dic['cache'][-1][current['stream']][current['layer']][current['module']].get(i, None) is not None) and (current['step'] > cache_dic['first_enhance'] - 2):
updated_factor = updated_taylor_factors[i].to(cache_dic['compute_device'])
cached = cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][i].to(cache_dic['compute_device'], non_blocking=True)
updated_taylor_factors[i + 1] = ((updated_factor - cached) / difference_distance).to(cache_dic['cache_device'], non_blocking=True)
else:
break
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']] = updated_taylor_factors
@torch.compiler.disable()
def taylor_formula(cache_dic: Dict, current: Dict) -> torch.Tensor:
"""
Compute Taylor expansion error
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
x = current['step'] - current['activated_steps'][-1]
#x = current['t'] - current['activated_times'][-1]
output = 0
for i in range(len(cache_dic['cache'][-1][current['stream']][current['layer']][current['module']])):
cached = cache_dic['cache'][-1][current['stream']][current['layer']][current['module']][i]
cached = cached.to(cache_dic['compute_device'], non_blocking=True)
output = output + (1 / math.factorial(i)) * cached * (x ** i)
return output
@torch.compiler.disable()
def taylor_cache_init(cache_dic: Dict, current: Dict):
"""
Initialize Taylor cache, expanding storage areas for Taylor series derivatives
:param cache_dic: Cache dictionary
:param current: Information of the current step
"""
if current['step'] == 0:
cache_dic['cache'][-1][current['stream']][current['layer']][current['module']] = {}
+41 -3
View File
@@ -450,6 +450,7 @@ class HyVideoModelLoader:
#compile
if compile_args is not None:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
if compile_args["compile_single_blocks"]:
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
@@ -628,8 +629,10 @@ class HyVideoTorchCompileSettings:
"compile_txt_in": ("BOOLEAN", {"default": False, "tooltip": "Compile txt_in layers"}),
"compile_vector_in": ("BOOLEAN", {"default": False, "tooltip": "Compile vector_in layers"}),
"compile_final_layer": ("BOOLEAN", {"default": False, "tooltip": "Compile final layer"}),
},
"optional": {
"dynamo_recompile_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
}
}
RETURN_TYPES = ("COMPILEARGS",)
RETURN_NAMES = ("torch_compile_args",)
@@ -637,7 +640,7 @@ class HyVideoTorchCompileSettings:
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended"
def loadmodel(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_single_blocks, compile_double_blocks, compile_txt_in, compile_vector_in, compile_final_layer):
def loadmodel(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_single_blocks, compile_double_blocks, compile_txt_in, compile_vector_in, compile_final_layer, dynamo_recompile_limit=64):
compile_args = {
"backend": backend,
@@ -645,6 +648,7 @@ class HyVideoTorchCompileSettings:
"mode": mode,
"dynamic": dynamic,
"dynamo_cache_size_limit": dynamo_cache_size_limit,
"dynamo_recompile_limit": dynamo_recompile_limit,
"compile_single_blocks": compile_single_blocks,
"compile_double_blocks": compile_double_blocks,
"compile_txt_in": compile_txt_in,
@@ -1266,6 +1270,36 @@ class HyVideoContextOptions:
}
return (context_options,)
class HyVideoTaylorSeerOptions:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"cache_device": (["main_device", "offload_device"], {"default": "offload_device"}),
"max_order": ("INT", {"default": 1, "min": 0, "max": 10, "step": 1, "tooltip": "Maximum order of the Taylor series expansion"}),
"fresh_threshold": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "A higher fresh_threshold results in faster inference but may reduce generation quality."}),
}
}
RETURN_TYPES = ("TAYLORSEERARGS", )
RETURN_NAMES = ("taylorseer_args",)
FUNCTION = "passargs"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "https://github.com/Shenyi-Z/TaylorSeer"
def passargs(self, cache_device, max_order, fresh_threshold):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
args = {
"cache_device": cache_device,
"compute_device": mm.get_torch_device(),
"max_order": max_order,
"fresh_threshold": fresh_threshold,
}
return (args,)
#region Sampler
class HyVideoSampler:
@classmethod
@@ -1298,6 +1332,7 @@ class HyVideoSampler:
}),
"riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}),
"i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}),
"taylorseer_args": ("TAYLORSEERARGS", ),
}
}
@@ -1308,7 +1343,7 @@ class HyVideoSampler:
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames,
samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None,
teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability"):
teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", taylorseer_args=None):
model = model.model
device = mm.get_torch_device()
@@ -1462,6 +1497,7 @@ class HyVideoSampler:
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
riflex_freq_index = riflex_freq_index,
i2v_stability = i2v_stability,
taylorseer = taylorseer_args,
)
print_memory(device)
@@ -1879,6 +1915,7 @@ NODE_CLASS_MAPPINGS = {
"HyVideoI2VEncode": HyVideoI2VEncode,
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
"HyVideoTaylorSeerOptions": HyVideoTaylorSeerOptions,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler",
@@ -1906,4 +1943,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoI2VEncode": "HyVideo I2V Encode",
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
"HyVideoTaylorSeerOptions": "HyVideo TaylorSeer Options",
}