98 lines
3.7 KiB
Python
Executable File
98 lines
3.7 KiB
Python
Executable File
import inspect
|
|
|
|
import numpy as np
|
|
import torch
|
|
|
|
|
|
def cfg_skip():
|
|
def decorator(func):
|
|
def wrapper(self, *args, **kwargs):
|
|
if torch.is_grad_enabled():
|
|
return func(self, *args, **kwargs)
|
|
|
|
if 'hidden_states' in kwargs and kwargs['hidden_states'] is not None:
|
|
main_input = kwargs['hidden_states']
|
|
elif 'x' in kwargs and kwargs['x'] is not None:
|
|
main_input = kwargs['x']
|
|
elif len(args) > 0:
|
|
main_input = args[0]
|
|
else:
|
|
raise ValueError("No input tensor found in args or kwargs")
|
|
|
|
bs = len(main_input)
|
|
if bs >= 2 and self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio):
|
|
bs_half = int(bs // 2)
|
|
|
|
new_x = main_input[bs_half:]
|
|
new_args = [
|
|
arg[bs_half:] if
|
|
isinstance(arg,
|
|
(torch.Tensor, list, tuple, np.ndarray)) and
|
|
len(arg) == bs else arg for arg in args
|
|
]
|
|
|
|
new_kwargs = {
|
|
k: (v[bs_half:] if
|
|
isinstance(v,
|
|
(torch.Tensor, list, tuple,
|
|
np.ndarray)) and len(v) == bs else v
|
|
) for k, v in kwargs.items()
|
|
}
|
|
else:
|
|
new_x = main_input
|
|
new_args = args
|
|
new_kwargs = kwargs
|
|
|
|
sig = inspect.signature(func)
|
|
|
|
new_bs = len(new_x)
|
|
new_bs_half = int(new_bs // 2)
|
|
if new_bs >= 2:
|
|
# cond
|
|
args_i = [
|
|
arg[new_bs_half:] if
|
|
isinstance(arg,
|
|
(torch.Tensor, list, tuple, np.ndarray)) and
|
|
len(arg) == new_bs else arg for arg in new_args
|
|
]
|
|
kwargs_i = {
|
|
k: (v[new_bs_half:] if
|
|
isinstance(v,
|
|
(torch.Tensor, list, tuple,
|
|
np.ndarray)) and len(v) == new_bs else v
|
|
) for k, v in new_kwargs.items()
|
|
}
|
|
if 'cond_flag' in sig.parameters:
|
|
kwargs_i["cond_flag"] = True
|
|
|
|
cond_out = func(self, *args_i, **kwargs_i)
|
|
|
|
# uncond
|
|
uncond_args_i = [
|
|
arg[:new_bs_half] if
|
|
isinstance(arg,
|
|
(torch.Tensor, list, tuple, np.ndarray)) and
|
|
len(arg) == new_bs else arg for arg in new_args
|
|
]
|
|
uncond_kwargs_i = {
|
|
k: (v[:new_bs_half] if
|
|
isinstance(v,
|
|
(torch.Tensor, list, tuple,
|
|
np.ndarray)) and len(v) == new_bs else v
|
|
) for k, v in new_kwargs.items()
|
|
}
|
|
if 'cond_flag' in sig.parameters:
|
|
uncond_kwargs_i["cond_flag"] = False
|
|
uncond_out = func(self, *uncond_args_i,
|
|
**uncond_kwargs_i)
|
|
|
|
result = torch.cat([uncond_out, cond_out], dim=0)
|
|
else:
|
|
result = func(self, *new_args, **new_kwargs)
|
|
|
|
if bs >= 2 and 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 |