Files
aigc-apps-VideoX-Fun/videox_fun/utils/cfg_optimization.py
T
Bubbliiiing a1e6ea6335 Pai sparse attention (#211)
* pai fuser sparse test

* make pai fuser more clear

* make pai fuser more clear

* make pai fuser more clear

* make pai fuser more clear

* Update Readme

* disable compile in rope for less error info

* Fix sage attention backward bug

* Fix checkpoint bugs

* Fix Attention

* Update predict and fast rope

* Fix bug in Sparse

* Fix bug in Sparse

* Fix bug in loras load

* Delete useless import

* Fix bug in ui

* Fix bug in ui
2025-05-26 16:00:58 +08:00

39 lines
1.3 KiB
Python

import numpy as np
import torch
def cfg_skip():
def decorator(func):
def wrapper(self, x, *args, **kwargs):
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
bs = len(x)
bs_half = int(bs // 2)
new_x = x[bs_half:]
new_args = []
for arg in args:
if isinstance(arg, (torch.Tensor, list, tuple, np.ndarray)):
new_args.append(arg[bs_half:])
else:
new_args.append(arg)
new_kwargs = {}
for key, content in kwargs.items():
if isinstance(content, (torch.Tensor, list, tuple, np.ndarray)):
new_kwargs[key] = content[bs_half:]
else:
new_kwargs[key] = content
else:
new_x = x
new_args = args
new_kwargs = kwargs
result = func(self, new_x, *new_args, **new_kwargs)
if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
result = torch.cat([result, result], dim=0)
return result
return wrapper
return decorator