feat(CRITICAL): 从 GitHub 扫描搬运 ixformer SDK + xllm 完整 GDN/MoE 代码
来源:
1. Chranos/ixformer (GitHub) → ixformer_sdk/ (230 files, 70K lines)
- inference/functions/vllm.py: vllm_moe_topk_softmax 完整实现 (2033 lines)
- inference/functions/moe.py: MoE ops 完整实现 (1380 lines)
- contrib/vllm_flash_attn/: FA2 Python 接口 (1018 lines)
- contrib/tgi/fused_moe.py: TGI fused MoE (429 lines)
- csrc/include/ixformer/: C++ kernel headers + cmake
2. Deep-Spark/xllm (GitHub) → upstream_ref/xllm_latest/ (+15 files)
- npu_torch/qwen3_5_decoder_layer_impl.cpp/.h
- npu_torch/qwen3_5_gated_delta_net.cpp/.h
- npu_torch/qwen3_next_*.cpp/.h (6 files)
- npu_torch/attention.cpp/.h + fused_moe.cpp/.h + CMakeLists.txt
- models/llm/qwen3_5.h + qwen3_5_mtp.h + qwen3_next.h
- models/vlm/qwen3_5.h
调用链完整性:
ixformer_sdk/inference/functions/vllm.py
→ ops.infer.moe_topk_softmax() (C++ 层)
→ 这就是 base 镜像 libixformer.so 里的实现
upstream_ref/xllm_latest/core/layers/ilu/fused_moe.cpp
→ ixformer::infer::topk_softmax() (直接 C++ 调用)
→ ixformer::infer::group_gemm() → 完整 7-step MoE pipeline
This commit is contained in:
5
ixformer_sdk/contrib/DeepCache/__init__.py
Normal file
5
ixformer_sdk/contrib/DeepCache/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
from .sd.pipeline_stable_diffusion import StableDiffusionPipeline
|
||||
from .sdxl.pipeline_stable_diffusion_xl import StableDiffusionXLPipeline
|
||||
from .sdxl.pipeline_stable_diffusion_xl_img2img import StableDiffusionXLImg2ImgPipeline
|
||||
|
||||
from .sd.pipeline_text_to_video_zero import TextToVideoZeroPipeline
|
||||
0
ixformer_sdk/contrib/DeepCache/ddpm/__init__.py
Normal file
0
ixformer_sdk/contrib/DeepCache/ddpm/__init__.py
Normal file
180
ixformer_sdk/contrib/DeepCache/ddpm/ddim.py
Normal file
180
ixformer_sdk/contrib/DeepCache/ddpm/ddim.py
Normal file
@@ -0,0 +1,180 @@
|
||||
import argparse
|
||||
import traceback
|
||||
import shutil
|
||||
import logging
|
||||
import yaml
|
||||
import random
|
||||
import sys
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from ddpm.utils.logging import Logger, EmptyLogger
|
||||
from ddpm.utils.tools import set_random_seed
|
||||
from accelerate import Accelerator, DistributedDataParallelKwargs
|
||||
|
||||
torch.set_printoptions(sci_mode=False)
|
||||
|
||||
def dict2namespace(config):
|
||||
namespace = argparse.Namespace()
|
||||
for key, value in config.items():
|
||||
if isinstance(value, dict):
|
||||
new_value = dict2namespace(value)
|
||||
else:
|
||||
new_value = value
|
||||
setattr(namespace, key, new_value)
|
||||
return namespace
|
||||
|
||||
def parse_args_and_config():
|
||||
parser = argparse.ArgumentParser(description=globals()["__doc__"])
|
||||
|
||||
parser.add_argument(
|
||||
"--config", type=str, required=True, help="Path to the config file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed", type=int, default=1234, help="Random seed")
|
||||
parser.add_argument(
|
||||
"--exp", type=str, default="exp", help="Path for saving running related data."
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--test", action="store_true", help="Whether to test the model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sample", action="store_true", help="Whether to produce samples from the model",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image_folder", type=str, default="images", help="folder name for storing the sampled images"
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--fid", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--interpolation", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume_training", action="store_true", help="Whether to resume training"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ni", action="store_true", help="No interaction. Suitable for Slurm Job launcher",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_pretrained", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sample_type", type=str, default="generalized", help="sampling approach (generalized or ddpm_noisy)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip_type", type=str, default="uniform", help="skip according to (uniform or quadratic)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--timesteps", type=int, default=1000, help="number of steps involved"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--eta", type=float, default=0.0, help="eta used to control the variances of sigma",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dyn", action="store_true", help="whether to activate the dynamic train/inference"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--select_step", type=int, default=None
|
||||
)
|
||||
parser.add_argument(
|
||||
"--select_depth", type=int, default=None
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--cache", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache_interval", type=int, default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--non_uniform", action="store_true"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pow", type=float, default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--center", type=int, default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--branch", type=int, default=None,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
# parse config file
|
||||
with open(args.config, "r") as f:
|
||||
config = yaml.safe_load(f)
|
||||
new_config = dict2namespace(config)
|
||||
new_config.select_step = args.select_step
|
||||
new_config.select_depth = args.select_depth
|
||||
|
||||
torch.backends.cudnn.benchmark = True
|
||||
|
||||
return args, new_config
|
||||
|
||||
|
||||
def main():
|
||||
args, config = parse_args_and_config()
|
||||
|
||||
if args.dyn:
|
||||
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
|
||||
accelerator = Accelerator(kwargs_handlers=[ddp_kwargs])
|
||||
else:
|
||||
accelerator = Accelerator()
|
||||
args.accelerator = accelerator
|
||||
|
||||
#log_root_dir = "{}_runtime_log".format(args.config[8:-4])
|
||||
log_root_dir = "runtime_log"
|
||||
dataset = args.config[8:-4]
|
||||
if args.cache:
|
||||
if args.non_uniform:
|
||||
sub_dir_name = "{}_{}_cache_{}_pow_{}_center_{}".format(dataset, args.exp, args.cache_interval, args.pow, args.center)
|
||||
else:
|
||||
sub_dir_name = "{}_{}_cache_{}".format(dataset, args.exp, args.cache_interval)
|
||||
else:
|
||||
sub_dir_name = "{}".format(args.exp)
|
||||
|
||||
if accelerator.is_main_process:
|
||||
logger = Logger(
|
||||
root_dir=log_root_dir,
|
||||
sub_name=sub_dir_name,
|
||||
config=args.__dict__,
|
||||
append=(args.sample == True)
|
||||
)
|
||||
args.logger = logger
|
||||
|
||||
args.logger.log("Writing log file to {}".format(args.logger.sub_dir))
|
||||
args.logger.log("Exp instance PID = {}".format(os.getpid()))
|
||||
else:
|
||||
args.logger = EmptyLogger(
|
||||
root_dir=log_root_dir,
|
||||
sub_name=sub_dir_name,
|
||||
)
|
||||
|
||||
args.image_folder = args.logger.setup_image_folder("{}".format(args.image_folder))
|
||||
|
||||
args.seed += accelerator.process_index
|
||||
# set random seed
|
||||
set_random_seed(args.seed)
|
||||
try:
|
||||
if args.cache:
|
||||
from ddpm.runners.deepcache import Diffusion
|
||||
runner = Diffusion(args, config)
|
||||
runner.sample()
|
||||
else:
|
||||
from ddpm.runners.diffusion import Diffusion
|
||||
runner = Diffusion(args, config)
|
||||
runner.sample()
|
||||
except Exception:
|
||||
logging.error(traceback.format_exc())
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
361
ixformer_sdk/contrib/DeepCache/ddpm/fid.py
Normal file
361
ixformer_sdk/contrib/DeepCache/ddpm/fid.py
Normal file
@@ -0,0 +1,361 @@
|
||||
"""Calculates the Frechet Inception Distance (FID) to evalulate GANs
|
||||
|
||||
The FID metric calculates the distance between two distributions of images.
|
||||
Typically, we have summary statistics (mean & covariance matrix) of one
|
||||
of these distributions, while the 2nd distribution is given by a GAN.
|
||||
|
||||
When run as a stand-alone program, it compares the distribution of
|
||||
images that are stored as PNG/JPEG at a specified location with a
|
||||
distribution given by summary statistics (in pickle format).
|
||||
|
||||
The FID is calculated by assuming that X_1 and X_2 are the activations of
|
||||
the pool_3 layer of the inception net for generated samples and real world
|
||||
samples respectively.
|
||||
|
||||
See --help to see further details.
|
||||
|
||||
Code apapted from https://github.com/bioinf-jku/TTUR to use PyTorch instead
|
||||
of Tensorflow
|
||||
|
||||
Copyright 2018 Institute of Bioinformatics, JKU Linz
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import os
|
||||
import pathlib
|
||||
from argparse import ArgumentDefaultsHelpFormatter, ArgumentParser
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms as TF
|
||||
from PIL import Image
|
||||
from scipy import linalg
|
||||
from torch.nn.functional import adaptive_avg_pool2d
|
||||
|
||||
try:
|
||||
from tqdm import tqdm
|
||||
except ImportError:
|
||||
# If tqdm is not available, provide a mock version of it
|
||||
def tqdm(x):
|
||||
return x
|
||||
|
||||
from pytorch_fid.inception import InceptionV3
|
||||
|
||||
parser = ArgumentParser(formatter_class=ArgumentDefaultsHelpFormatter)
|
||||
parser.add_argument('--batch-size', type=int, default=50,
|
||||
help='Batch size to use')
|
||||
parser.add_argument('--dataset_name', type=str, default=None)
|
||||
parser.add_argument('--num-workers', type=int,
|
||||
help=('Number of processes to use for data loading. '
|
||||
'Defaults to `min(8, num_cpus)`'))
|
||||
parser.add_argument('--device', type=str, default=None,
|
||||
help='Device to use. Like cuda, cuda:0 or cpu')
|
||||
parser.add_argument('--dims', type=int, default=2048,
|
||||
choices=list(InceptionV3.BLOCK_INDEX_BY_DIM),
|
||||
help=('Dimensionality of Inception features to use. '
|
||||
'By default, uses pool3 features'))
|
||||
parser.add_argument('--num_samples', type=int, default=None,
|
||||
help=('Number of samples for FID estimation'))
|
||||
parser.add_argument('--res', type=int, default=None,
|
||||
help=('Resolutions of samples for FID estimation'))
|
||||
parser.add_argument('--save-stats', action='store_true',
|
||||
help=('Generate an npz archive from a directory of samples. '
|
||||
'The first path is used as input and the second as output.'))
|
||||
|
||||
parser.add_argument('--path', type=str, nargs=2,
|
||||
help=('Paths to the generated images or '
|
||||
'to .npz statistic files'))
|
||||
|
||||
|
||||
IMAGE_EXTENSIONS = {'bmp', 'jpg', 'jpeg', 'pgm', 'png', 'ppm',
|
||||
'tif', 'tiff', 'webp'}
|
||||
|
||||
|
||||
class ImagePathDataset(torch.utils.data.Dataset):
|
||||
def __init__(self, files, transforms=None):
|
||||
self.files = files
|
||||
self.transforms = transforms
|
||||
|
||||
def __len__(self):
|
||||
return len(self.files)
|
||||
|
||||
def __getitem__(self, i):
|
||||
path = self.files[i]
|
||||
img = Image.open(path).convert('RGB')
|
||||
if self.transforms is not None:
|
||||
img = self.transforms(img)
|
||||
return img
|
||||
|
||||
|
||||
def get_activations(files, model, batch_size=50, dims=2048, device='cpu',
|
||||
num_workers=1, res=None, dataset_name=None):
|
||||
"""Calculates the activations of the pool_3 layer for all images.
|
||||
|
||||
Params:
|
||||
-- files : List of image files paths
|
||||
-- model : Instance of inception model
|
||||
-- batch_size : Batch size of images for the model to process at once.
|
||||
Make sure that the number of samples is a multiple of
|
||||
the batch size, otherwise some samples are ignored. This
|
||||
behavior is retained to match the original FID score
|
||||
implementation.
|
||||
-- dims : Dimensionality of features returned by Inception
|
||||
-- device : Device to run calculations
|
||||
-- num_workers : Number of parallel dataloader workers
|
||||
|
||||
Returns:
|
||||
-- A numpy array of dimension (num images, dims) that contains the
|
||||
activations of the given tensor when feeding inception with the
|
||||
query tensor.
|
||||
"""
|
||||
model.eval()
|
||||
|
||||
if batch_size > len(files):
|
||||
print(('Warning: batch size is bigger than the data size. '
|
||||
'Setting batch size to data size'))
|
||||
batch_size = len(files)
|
||||
|
||||
if res is None:
|
||||
trans = TF.ToTensor()
|
||||
else:
|
||||
if dataset_name == 'celeba':
|
||||
from switchable_diffusion.datasets import Crop
|
||||
print("In crop image: {}, {}".format(res, dataset_name))
|
||||
cx = 89
|
||||
cy = 121
|
||||
x1 = cy - 64
|
||||
x2 = cy + 64
|
||||
y1 = cx - 64
|
||||
y2 = cx + 64
|
||||
trans = TF.Compose([
|
||||
Crop(x1, x2, y1, y2),
|
||||
TF.Resize(res),
|
||||
TF.ToTensor(),
|
||||
])
|
||||
else:
|
||||
trans = TF.Compose([
|
||||
TF.Resize(res),
|
||||
TF.CenterCrop(res),
|
||||
TF.ToTensor()
|
||||
])
|
||||
|
||||
dataset = ImagePathDataset(files, transforms=trans)
|
||||
dataloader = torch.utils.data.DataLoader(dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=False,
|
||||
drop_last=False,
|
||||
num_workers=num_workers)
|
||||
|
||||
pred_arr = np.empty((len(files), dims))
|
||||
|
||||
start_idx = 0
|
||||
|
||||
for batch in tqdm(dataloader):
|
||||
batch = batch.to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
pred = model(batch)[0]
|
||||
|
||||
# If model output is not scalar, apply global spatial average pooling.
|
||||
# This happens if you choose a dimensionality not equal 2048.
|
||||
if pred.size(2) != 1 or pred.size(3) != 1:
|
||||
pred = adaptive_avg_pool2d(pred, output_size=(1, 1))
|
||||
|
||||
pred = pred.squeeze(3).squeeze(2).cpu().numpy()
|
||||
|
||||
pred_arr[start_idx:start_idx + pred.shape[0]] = pred
|
||||
|
||||
start_idx = start_idx + pred.shape[0]
|
||||
|
||||
return pred_arr
|
||||
|
||||
|
||||
def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
|
||||
"""Numpy implementation of the Frechet Distance.
|
||||
The Frechet distance between two multivariate Gaussians X_1 ~ N(mu_1, C_1)
|
||||
and X_2 ~ N(mu_2, C_2) is
|
||||
d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
|
||||
|
||||
Stable version by Dougal J. Sutherland.
|
||||
|
||||
Params:
|
||||
-- mu1 : Numpy array containing the activations of a layer of the
|
||||
inception net (like returned by the function 'get_predictions')
|
||||
for generated samples.
|
||||
-- mu2 : The sample mean over activations, precalculated on an
|
||||
representative data set.
|
||||
-- sigma1: The covariance matrix over activations for generated samples.
|
||||
-- sigma2: The covariance matrix over activations, precalculated on an
|
||||
representative data set.
|
||||
|
||||
Returns:
|
||||
-- : The Frechet Distance.
|
||||
"""
|
||||
|
||||
mu1 = np.atleast_1d(mu1)
|
||||
mu2 = np.atleast_1d(mu2)
|
||||
|
||||
sigma1 = np.atleast_2d(sigma1)
|
||||
sigma2 = np.atleast_2d(sigma2)
|
||||
|
||||
assert mu1.shape == mu2.shape, \
|
||||
'Training and test mean vectors have different lengths'
|
||||
assert sigma1.shape == sigma2.shape, \
|
||||
'Training and test covariances have different dimensions'
|
||||
|
||||
diff = mu1 - mu2
|
||||
|
||||
# Product might be almost singular
|
||||
covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
|
||||
if not np.isfinite(covmean).all():
|
||||
msg = ('fid calculation produces singular product; '
|
||||
'adding %s to diagonal of cov estimates') % eps
|
||||
print(msg)
|
||||
offset = np.eye(sigma1.shape[0]) * eps
|
||||
covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
|
||||
|
||||
# Numerical error might give slight imaginary component
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
m = np.max(np.abs(covmean.imag))
|
||||
raise ValueError('Imaginary component {}'.format(m))
|
||||
covmean = covmean.real
|
||||
|
||||
tr_covmean = np.trace(covmean)
|
||||
|
||||
return (diff.dot(diff) + np.trace(sigma1)
|
||||
+ np.trace(sigma2) - 2 * tr_covmean)
|
||||
|
||||
|
||||
def calculate_activation_statistics(files, model, batch_size=50, dims=2048,
|
||||
device='cpu', num_workers=1, res=None, dataset_name=None):
|
||||
"""Calculation of the statistics used by the FID.
|
||||
Params:
|
||||
-- files : List of image files paths
|
||||
-- model : Instance of inception model
|
||||
-- batch_size : The images numpy array is split into batches with
|
||||
batch size batch_size. A reasonable batch size
|
||||
depends on the hardware.
|
||||
-- dims : Dimensionality of features returned by Inception
|
||||
-- device : Device to run calculations
|
||||
-- num_workers : Number of parallel dataloader workers
|
||||
|
||||
Returns:
|
||||
-- mu : The mean over samples of the activations of the pool_3 layer of
|
||||
the inception model.
|
||||
-- sigma : The covariance matrix of the activations of the pool_3 layer of
|
||||
the inception model.
|
||||
"""
|
||||
act = get_activations(files, model, batch_size, dims, device, num_workers, res=res, dataset_name=dataset_name)
|
||||
mu = np.mean(act, axis=0)
|
||||
sigma = np.cov(act, rowvar=False)
|
||||
return mu, sigma
|
||||
|
||||
|
||||
def compute_statistics_of_path(path, model, batch_size, dims, device,
|
||||
num_workers=1, num_samples=None, res=None, dataset_name=None):
|
||||
if path.endswith('.npz'):
|
||||
with np.load(path) as f:
|
||||
m, s = f['mu'][:], f['sigma'][:]
|
||||
else:
|
||||
path = pathlib.Path(path)
|
||||
|
||||
files = sorted([file for ext in IMAGE_EXTENSIONS
|
||||
for file in path.glob('**/*.{}'.format(ext))])
|
||||
if num_samples is not None:
|
||||
#import random
|
||||
#files = random.sample(files, num_samples)
|
||||
files = files[:num_samples]
|
||||
print("Found %d files." % len(files))
|
||||
m, s = calculate_activation_statistics(files, model, batch_size,
|
||||
dims, device, num_workers, res=res, dataset_name=dataset_name)
|
||||
|
||||
return m, s
|
||||
|
||||
|
||||
def calculate_fid_given_paths(paths, batch_size, device, dims, num_workers=1, num_samples=None, res=None, dataset_name=None):
|
||||
"""Calculates the FID of two paths"""
|
||||
for p in paths:
|
||||
if not os.path.exists(p):
|
||||
raise RuntimeError('Invalid path: %s' % p)
|
||||
|
||||
block_idx = InceptionV3.BLOCK_INDEX_BY_DIM[dims]
|
||||
|
||||
model = InceptionV3([block_idx]).to(device)
|
||||
|
||||
m1, s1 = compute_statistics_of_path(paths[0], model, batch_size,
|
||||
dims, device, num_workers, num_samples=num_samples, res=res, dataset_name=dataset_name)
|
||||
m2, s2 = compute_statistics_of_path(paths[1], model, batch_size,
|
||||
dims, device, num_workers, num_samples=num_samples, res=res, dataset_name=dataset_name)
|
||||
fid_value = calculate_frechet_distance(m1, s1, m2, s2)
|
||||
|
||||
return fid_value
|
||||
|
||||
|
||||
def save_fid_stats(paths, batch_size, device, dims, num_workers=1, num_samples=None, res=None, dataset_name=None):
|
||||
"""Calculates the FID of two paths"""
|
||||
if not os.path.exists(paths[0]):
|
||||
raise RuntimeError('Invalid path: %s' % paths[0])
|
||||
|
||||
if os.path.exists(paths[1]):
|
||||
raise RuntimeError('Existing output file: %s' % paths[1])
|
||||
|
||||
block_idx = InceptionV3.BLOCK_INDEX_BY_DIM[dims]
|
||||
|
||||
model = InceptionV3([block_idx]).to(device)
|
||||
|
||||
print(f"Saving statistics for {paths[0]}")
|
||||
|
||||
m1, s1 = compute_statistics_of_path(paths[0], model, batch_size,
|
||||
dims, device, num_workers, num_samples=num_samples, res=res, dataset_name=dataset_name)
|
||||
|
||||
np.savez_compressed(paths[1], mu=m1, sigma=s1)
|
||||
|
||||
|
||||
def main():
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.device is None:
|
||||
device = torch.device('cuda' if (torch.cuda.is_available()) else 'cpu')
|
||||
else:
|
||||
device = torch.device(args.device)
|
||||
|
||||
if args.num_workers is None:
|
||||
try:
|
||||
num_cpus = len(os.sched_getaffinity(0))
|
||||
except AttributeError:
|
||||
# os.sched_getaffinity is not available under Windows, use
|
||||
# os.cpu_count instead (which may not return the *available* number
|
||||
# of CPUs).
|
||||
num_cpus = os.cpu_count()
|
||||
|
||||
num_workers = min(num_cpus, 8) if num_cpus is not None else 0
|
||||
else:
|
||||
num_workers = args.num_workers
|
||||
|
||||
if args.save_stats:
|
||||
save_fid_stats(args.path, args.batch_size, device, args.dims, num_workers, num_samples=args.num_samples, res=args.res, dataset_name=args.dataset_name)
|
||||
return
|
||||
|
||||
fid_value = calculate_fid_given_paths(args.path,
|
||||
args.batch_size,
|
||||
device,
|
||||
args.dims,
|
||||
num_workers,
|
||||
num_samples=args.num_samples,
|
||||
res = args.res, dataset_name=args.dataset_name)
|
||||
print('FID: ', fid_value)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
559
ixformer_sdk/contrib/DeepCache/flops.py
Normal file
559
ixformer_sdk/contrib/DeepCache/flops.py
Normal file
@@ -0,0 +1,559 @@
|
||||
'''
|
||||
This opcounter is adapted from https://github.com/sovrasov/flops-counter.pytorch and https://github.com/Lyken17/pytorch-OpCounter
|
||||
|
||||
Copyright (C) 2021 Sovrasov V. - All Rights Reserved
|
||||
* You may use, distribute and modify this code under the
|
||||
* terms of the MIT license.
|
||||
* You should have received a copy of the MIT license with
|
||||
* this file. If not visit https://opensource.org/licenses/MIT
|
||||
'''
|
||||
import os
|
||||
import yaml
|
||||
import numpy as np
|
||||
import torch.nn as nn
|
||||
import torch
|
||||
has_timm = False
|
||||
|
||||
from diffusers.models.lora import LoRACompatibleLinear, LoRACompatibleConv
|
||||
|
||||
@torch.no_grad()
|
||||
def count_ops_and_params(model, example_inputs, layer_wise=False):
|
||||
global CUSTOM_MODULES_MAPPING
|
||||
ori_model = model
|
||||
model = copy.deepcopy(model) # deepcopy to avoid changing the original model
|
||||
flops_model = add_flops_counting_methods(model)
|
||||
flops_model.eval()
|
||||
flops_model.start_flops_count(ost=sys.stdout, verbose=False,
|
||||
ignore_list=[])
|
||||
if isinstance(example_inputs, (tuple, list)):
|
||||
_ = flops_model(*example_inputs)
|
||||
elif isinstance(example_inputs, dict):
|
||||
_ = flops_model(**example_inputs)
|
||||
else:
|
||||
_ = flops_model(example_inputs)
|
||||
flops_count, params_count, _layer_flops, _layer_params = flops_model.compute_average_flops_cost()
|
||||
layer_flops = {}
|
||||
layer_params = {}
|
||||
|
||||
for m_name, m in model.named_modules():
|
||||
layer_flops[m_name] = _layer_flops.get(m)
|
||||
layer_params[m_name] = _layer_params.get(m)
|
||||
if layer_wise:
|
||||
space = 30 - len(m_name)
|
||||
print("Layer {}: {} MACs = {:.4f} G, Params = {:.4f} M, MACs% = {:.2f}".format(
|
||||
m_name, ' ' * space, layer_flops[m_name]/1e9, layer_params[m_name] / 1e6, 100 * layer_flops[m_name] / flops_count
|
||||
))
|
||||
|
||||
flops_model.stop_flops_count()
|
||||
CUSTOM_MODULES_MAPPING = {}
|
||||
#if layer_wise:
|
||||
# return flops_count, params_count, layer_flops, layer_params
|
||||
return flops_count, params_count
|
||||
|
||||
def empty_flops_counter_hook(module, input, output):
|
||||
module.__flops__ += 0
|
||||
|
||||
|
||||
def upsample_flops_counter_hook(module, input, output):
|
||||
output_size = output[0]
|
||||
batch_size = output_size.shape[0]
|
||||
output_elements_count = batch_size
|
||||
for val in output_size.shape[1:]:
|
||||
output_elements_count *= val
|
||||
module.__flops__ += int(output_elements_count)
|
||||
|
||||
|
||||
def relu_flops_counter_hook(module, input, output):
|
||||
active_elements_count = output.numel()
|
||||
module.__flops__ += int(active_elements_count)
|
||||
|
||||
|
||||
def linear_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
# pytorch checks dimensions, so here we don't care much
|
||||
output_last_dim = output.shape[-1]
|
||||
bias_flops = output_last_dim if module.bias is not None else 0
|
||||
module.__flops__ += int(np.prod(input.shape) * output_last_dim + bias_flops)
|
||||
|
||||
|
||||
def pool_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
module.__flops__ += int(np.prod(input.shape))
|
||||
|
||||
|
||||
def bn_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
|
||||
batch_flops = np.prod(input.shape)
|
||||
if module.affine:
|
||||
batch_flops *= 2
|
||||
module.__flops__ += int(batch_flops)
|
||||
|
||||
def ln_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
batch_flops = np.prod(input.shape)
|
||||
if module.elementwise_affine:
|
||||
batch_flops *= 2
|
||||
module.__flops__ += int(batch_flops)
|
||||
|
||||
def conv_flops_counter_hook(conv_module, input, output):
|
||||
# Can have multiple inputs, getting the first one
|
||||
input = input[0]
|
||||
|
||||
batch_size = input.shape[0]
|
||||
output_dims = list(output.shape[2:])
|
||||
|
||||
kernel_dims = list(conv_module.kernel_size)
|
||||
in_channels = conv_module.in_channels
|
||||
out_channels = conv_module.out_channels
|
||||
groups = conv_module.groups
|
||||
|
||||
filters_per_channel = out_channels // groups
|
||||
conv_per_position_flops = int(np.prod(kernel_dims)) * \
|
||||
in_channels * filters_per_channel
|
||||
|
||||
active_elements_count = batch_size * int(np.prod(output_dims))
|
||||
|
||||
overall_conv_flops = conv_per_position_flops * active_elements_count
|
||||
|
||||
bias_flops = 0
|
||||
|
||||
if conv_module.bias is not None:
|
||||
|
||||
bias_flops = out_channels * active_elements_count
|
||||
|
||||
overall_flops = overall_conv_flops + bias_flops
|
||||
|
||||
conv_module.__flops__ += int(overall_flops)
|
||||
|
||||
|
||||
def rnn_flops(flops, rnn_module, w_ih, w_hh, input_size):
|
||||
# matrix matrix mult ih state and internal state
|
||||
flops += w_ih.shape[0]*w_ih.shape[1]
|
||||
# matrix matrix mult hh state and internal state
|
||||
flops += w_hh.shape[0]*w_hh.shape[1]
|
||||
if isinstance(rnn_module, (nn.RNN, nn.RNNCell)):
|
||||
# add both operations
|
||||
flops += rnn_module.hidden_size
|
||||
elif isinstance(rnn_module, (nn.GRU, nn.GRUCell)):
|
||||
# hadamard of r
|
||||
flops += rnn_module.hidden_size
|
||||
# adding operations from both states
|
||||
flops += rnn_module.hidden_size*3
|
||||
# last two hadamard product and add
|
||||
flops += rnn_module.hidden_size*3
|
||||
elif isinstance(rnn_module, (nn.LSTM, nn.LSTMCell)):
|
||||
# adding operations from both states
|
||||
flops += rnn_module.hidden_size*4
|
||||
# two hadamard product and add for C state
|
||||
flops += rnn_module.hidden_size + rnn_module.hidden_size + rnn_module.hidden_size
|
||||
# final hadamard
|
||||
flops += rnn_module.hidden_size + rnn_module.hidden_size + rnn_module.hidden_size
|
||||
return flops
|
||||
|
||||
|
||||
def rnn_flops_counter_hook(rnn_module, input, output):
|
||||
"""
|
||||
Takes into account batch goes at first position, contrary
|
||||
to pytorch common rule (but actually it doesn't matter).
|
||||
If sigmoid and tanh are hard, only a comparison FLOPS should be accurate
|
||||
"""
|
||||
flops = 0
|
||||
# input is a tuple containing a sequence to process and (optionally) hidden state
|
||||
inp = input[0]
|
||||
batch_size = inp[0].shape[0]
|
||||
seq_length = inp[0].shape[1]
|
||||
num_layers = rnn_module.num_layers
|
||||
|
||||
for i in range(num_layers):
|
||||
w_ih = rnn_module.__getattr__('weight_ih_l' + str(i))
|
||||
w_hh = rnn_module.__getattr__('weight_hh_l' + str(i))
|
||||
if i == 0:
|
||||
input_size = rnn_module.input_size
|
||||
else:
|
||||
input_size = rnn_module.hidden_size
|
||||
flops = rnn_flops(flops, rnn_module, w_ih, w_hh, input_size)
|
||||
if rnn_module.bias:
|
||||
b_ih = rnn_module.__getattr__('bias_ih_l' + str(i))
|
||||
b_hh = rnn_module.__getattr__('bias_hh_l' + str(i))
|
||||
flops += b_ih.shape[0] + b_hh.shape[0]
|
||||
|
||||
flops *= batch_size
|
||||
flops *= seq_length
|
||||
if rnn_module.bidirectional:
|
||||
flops *= 2
|
||||
rnn_module.__flops__ += int(flops)
|
||||
|
||||
|
||||
def rnn_cell_flops_counter_hook(rnn_cell_module, input, output):
|
||||
flops = 0
|
||||
inp = input[0]
|
||||
batch_size = inp.shape[0]
|
||||
w_ih = rnn_cell_module.__getattr__('weight_ih')
|
||||
w_hh = rnn_cell_module.__getattr__('weight_hh')
|
||||
input_size = inp.shape[1]
|
||||
flops = rnn_flops(flops, rnn_cell_module, w_ih, w_hh, input_size)
|
||||
if rnn_cell_module.bias:
|
||||
b_ih = rnn_cell_module.__getattr__('bias_ih')
|
||||
b_hh = rnn_cell_module.__getattr__('bias_hh')
|
||||
flops += b_ih.shape[0] + b_hh.shape[0]
|
||||
|
||||
flops *= batch_size
|
||||
rnn_cell_module.__flops__ += int(flops)
|
||||
|
||||
|
||||
def multihead_attention_counter_hook(multihead_attention_module, input, output):
|
||||
flops = 0
|
||||
q, k, v = input
|
||||
|
||||
batch_first = multihead_attention_module.batch_first \
|
||||
if hasattr(multihead_attention_module, 'batch_first') else False
|
||||
if batch_first:
|
||||
batch_size = q.shape[0]
|
||||
len_idx = 1
|
||||
else:
|
||||
batch_size = q.shape[1]
|
||||
len_idx = 0
|
||||
|
||||
dim_idx = 2
|
||||
|
||||
qdim = q.shape[dim_idx]
|
||||
kdim = k.shape[dim_idx]
|
||||
vdim = v.shape[dim_idx]
|
||||
|
||||
qlen = q.shape[len_idx]
|
||||
klen = k.shape[len_idx]
|
||||
vlen = v.shape[len_idx]
|
||||
|
||||
num_heads = multihead_attention_module.num_heads
|
||||
assert qdim == multihead_attention_module.embed_dim
|
||||
|
||||
if multihead_attention_module.kdim is None:
|
||||
assert kdim == qdim
|
||||
if multihead_attention_module.vdim is None:
|
||||
assert vdim == qdim
|
||||
|
||||
flops = 0
|
||||
|
||||
# Q scaling
|
||||
flops += qlen * qdim
|
||||
# Initial projections
|
||||
flops += (
|
||||
(qlen * qdim * qdim) # QW
|
||||
+ (klen * kdim * kdim) # KW
|
||||
+ (vlen * vdim * vdim) # VW
|
||||
)
|
||||
if multihead_attention_module.in_proj_bias is not None:
|
||||
flops += (qlen + klen + vlen) * qdim
|
||||
# attention heads: scale, matmul, softmax, matmul
|
||||
qk_head_dim = qdim // num_heads
|
||||
v_head_dim = vdim // num_heads
|
||||
|
||||
head_flops = (
|
||||
(qlen * klen * qk_head_dim) # QK^T
|
||||
+ (qlen * klen) # softmax
|
||||
+ (qlen * klen * v_head_dim) # AV
|
||||
)
|
||||
flops += num_heads * head_flops
|
||||
# final projection, bias is always enabled
|
||||
flops += qlen * vdim * (vdim + 1)
|
||||
flops *= batch_size
|
||||
multihead_attention_module.__flops__ += int(flops)
|
||||
|
||||
def timm_multihead_attention_counter_hook(multihead_attention_module, input, output):
|
||||
flops = 0
|
||||
|
||||
q, k, v = input[0], input[0], input[0]
|
||||
input_dim = input[0].shape[2]
|
||||
input_len = input[0].shape[1]
|
||||
batch_size = input[0].shape[0]
|
||||
|
||||
kdim = qdim = vdim = multihead_attention_module.qkv.out_features//3
|
||||
qlen = klen = vlen = input_len
|
||||
|
||||
num_heads = multihead_attention_module.num_heads
|
||||
assert qdim == multihead_attention_module.head_dim * multihead_attention_module.num_heads
|
||||
|
||||
flops = 0
|
||||
# Q scaling
|
||||
flops += qlen * qdim
|
||||
# Initial projections
|
||||
flops += (
|
||||
(qlen * input_dim * qdim) # QW
|
||||
+ (klen * input_dim * kdim) # KW
|
||||
+ (vlen * input_dim * vdim) # VW
|
||||
)
|
||||
|
||||
if multihead_attention_module.qkv.bias is not None:
|
||||
flops += (qlen + klen + vlen) * qdim
|
||||
# attention heads: scale, matmul, softmax, matmul
|
||||
qk_head_dim = qdim // num_heads
|
||||
v_head_dim = vdim // num_heads
|
||||
|
||||
head_flops = (
|
||||
(qlen * klen * qk_head_dim) # QK^T
|
||||
+ (qlen * klen) # softmax
|
||||
+ (qlen * klen * v_head_dim) # AV
|
||||
)
|
||||
flops += num_heads * head_flops
|
||||
# final projection, bias is always enabled
|
||||
flops += qlen * vdim * (vdim + 1)
|
||||
flops *= batch_size
|
||||
multihead_attention_module.__flops__ += int(flops)
|
||||
|
||||
|
||||
|
||||
CUSTOM_MODULES_MAPPING = {}
|
||||
|
||||
MODULES_MAPPING = {
|
||||
# convolutions
|
||||
nn.Conv1d: conv_flops_counter_hook,
|
||||
nn.Conv2d: conv_flops_counter_hook,
|
||||
nn.Conv3d: conv_flops_counter_hook,
|
||||
LoRACompatibleConv: conv_flops_counter_hook,
|
||||
# activations
|
||||
nn.ReLU: relu_flops_counter_hook,
|
||||
nn.PReLU: relu_flops_counter_hook,
|
||||
nn.ELU: relu_flops_counter_hook,
|
||||
nn.LeakyReLU: relu_flops_counter_hook,
|
||||
nn.ReLU6: relu_flops_counter_hook,
|
||||
# poolings
|
||||
nn.MaxPool1d: pool_flops_counter_hook,
|
||||
nn.AvgPool1d: pool_flops_counter_hook,
|
||||
nn.AvgPool2d: pool_flops_counter_hook,
|
||||
nn.MaxPool2d: pool_flops_counter_hook,
|
||||
nn.MaxPool3d: pool_flops_counter_hook,
|
||||
nn.AvgPool3d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool1d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool1d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool2d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool2d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool3d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool3d: pool_flops_counter_hook,
|
||||
# BNs
|
||||
nn.BatchNorm1d: bn_flops_counter_hook,
|
||||
nn.BatchNorm2d: bn_flops_counter_hook,
|
||||
nn.BatchNorm3d: bn_flops_counter_hook,
|
||||
|
||||
nn.InstanceNorm1d: bn_flops_counter_hook,
|
||||
nn.InstanceNorm2d: bn_flops_counter_hook,
|
||||
nn.InstanceNorm3d: bn_flops_counter_hook,
|
||||
nn.GroupNorm: bn_flops_counter_hook,
|
||||
nn.LayerNorm: ln_flops_counter_hook,
|
||||
# FC
|
||||
nn.Linear: linear_flops_counter_hook,
|
||||
LoRACompatibleLinear: linear_flops_counter_hook,
|
||||
# Upscale
|
||||
nn.Upsample: upsample_flops_counter_hook,
|
||||
# Deconvolution
|
||||
nn.ConvTranspose1d: conv_flops_counter_hook,
|
||||
nn.ConvTranspose2d: conv_flops_counter_hook,
|
||||
nn.ConvTranspose3d: conv_flops_counter_hook,
|
||||
# RNN
|
||||
nn.RNN: rnn_flops_counter_hook,
|
||||
nn.GRU: rnn_flops_counter_hook,
|
||||
nn.LSTM: rnn_flops_counter_hook,
|
||||
nn.RNNCell: rnn_cell_flops_counter_hook,
|
||||
nn.LSTMCell: rnn_cell_flops_counter_hook,
|
||||
nn.GRUCell: rnn_cell_flops_counter_hook,
|
||||
nn.MultiheadAttention: multihead_attention_counter_hook
|
||||
}
|
||||
|
||||
if has_timm:
|
||||
MODULES_MAPPING.update(
|
||||
{
|
||||
timm.models.vision_transformer.Attention: timm_multihead_attention_counter_hook,
|
||||
}
|
||||
)
|
||||
|
||||
if hasattr(nn, 'GELU'):
|
||||
MODULES_MAPPING[nn.GELU] = relu_flops_counter_hook
|
||||
|
||||
|
||||
import sys
|
||||
from functools import partial
|
||||
import torch.nn as nn
|
||||
import copy
|
||||
|
||||
def accumulate_flops(self, layer_flops):
|
||||
if is_supported_instance(self):
|
||||
layer_flops[self] = self.__flops__
|
||||
return self.__flops__
|
||||
else:
|
||||
sum = 0
|
||||
for m in self.children():
|
||||
sum += m.accumulate_flops(layer_flops)
|
||||
layer_flops[self] = sum
|
||||
return sum
|
||||
|
||||
|
||||
def get_model_parameters_number(model):
|
||||
params_num = sum(p.numel() for p in model.parameters())
|
||||
return params_num
|
||||
|
||||
|
||||
def add_flops_counting_methods(net_main_module):
|
||||
# adding additional methods to the existing module object,
|
||||
# this is done this way so that each function has access to self object
|
||||
net_main_module.start_flops_count = start_flops_count.__get__(net_main_module)
|
||||
net_main_module.stop_flops_count = stop_flops_count.__get__(net_main_module)
|
||||
net_main_module.reset_flops_count = reset_flops_count.__get__(net_main_module)
|
||||
net_main_module.compute_average_flops_cost = compute_average_flops_cost.__get__(
|
||||
net_main_module)
|
||||
|
||||
net_main_module.reset_flops_count()
|
||||
|
||||
return net_main_module
|
||||
|
||||
def compute_average_flops_cost(self):
|
||||
"""
|
||||
A method that will be available after add_flops_counting_methods() is called
|
||||
on a desired net object.
|
||||
Returns current mean flops consumption per image.
|
||||
"""
|
||||
|
||||
for m in self.modules():
|
||||
m.accumulate_flops = accumulate_flops.__get__(m)
|
||||
|
||||
layer_flops = {}
|
||||
flops_sum = self.accumulate_flops(layer_flops)
|
||||
|
||||
for m in self.modules():
|
||||
if hasattr(m, 'accumulate_flops'):
|
||||
del m.accumulate_flops
|
||||
|
||||
layer_params = {}
|
||||
for m in self.modules():
|
||||
layer_params[m] = get_model_parameters_number(m)
|
||||
|
||||
params_sum = get_model_parameters_number(self)
|
||||
return flops_sum / self.__batch_counter__, params_sum, layer_flops, layer_params
|
||||
|
||||
|
||||
def start_flops_count(self, **kwargs):
|
||||
"""
|
||||
A method that will be available after add_flops_counting_methods() is called
|
||||
on a desired net object.
|
||||
Activates the computation of mean flops consumption per image.
|
||||
Call it before you run the network.
|
||||
"""
|
||||
add_batch_counter_hook_function(self)
|
||||
|
||||
seen_types = set()
|
||||
|
||||
def add_flops_counter_hook_function(module, ost, verbose, ignore_list):
|
||||
if type(module) in ignore_list:
|
||||
seen_types.add(type(module))
|
||||
if is_supported_instance(module):
|
||||
module.__params__ = 0
|
||||
elif is_supported_instance(module):
|
||||
if hasattr(module, '__flops_handle__'):
|
||||
return
|
||||
if type(module) in CUSTOM_MODULES_MAPPING:
|
||||
handle = module.register_forward_hook(
|
||||
CUSTOM_MODULES_MAPPING[type(module)])
|
||||
else:
|
||||
handle = module.register_forward_hook(MODULES_MAPPING[type(module)])
|
||||
module.__flops_handle__ = handle
|
||||
seen_types.add(type(module))
|
||||
else:
|
||||
if verbose and not type(module) in (nn.Sequential, nn.ModuleList) and \
|
||||
not type(module) in seen_types:
|
||||
print('Warning: module ' + type(module).__name__ +
|
||||
' is treated as a zero-op.', file=ost)
|
||||
seen_types.add(type(module))
|
||||
|
||||
self.apply(partial(add_flops_counter_hook_function, **kwargs))
|
||||
|
||||
|
||||
def stop_flops_count(self):
|
||||
"""
|
||||
A method that will be available after add_flops_counting_methods() is called
|
||||
on a desired net object.
|
||||
Stops computing the mean flops consumption per image.
|
||||
Call whenever you want to pause the computation.
|
||||
"""
|
||||
remove_batch_counter_hook_function(self)
|
||||
self.apply(remove_flops_counter_hook_function)
|
||||
self.apply(remove_flops_counter_variables)
|
||||
|
||||
|
||||
def reset_flops_count(self):
|
||||
"""
|
||||
A method that will be available after add_flops_counting_methods() is called
|
||||
on a desired net object.
|
||||
Resets statistics computed so far.
|
||||
"""
|
||||
add_batch_counter_variables_or_reset(self)
|
||||
self.apply(add_flops_counter_variable_or_reset)
|
||||
|
||||
|
||||
# ---- Internal functions
|
||||
def batch_counter_hook(module, input, output):
|
||||
batch_size = 1
|
||||
if len(input) > 0:
|
||||
# Can have multiple inputs, getting the first one
|
||||
input = input[0]
|
||||
batch_size = len(input)
|
||||
else:
|
||||
pass
|
||||
print('Warning! No positional inputs found for a module,'
|
||||
' assuming batch size is 1.')
|
||||
module.__batch_counter__ += batch_size
|
||||
|
||||
|
||||
def add_batch_counter_variables_or_reset(module):
|
||||
|
||||
module.__batch_counter__ = 0
|
||||
|
||||
|
||||
def add_batch_counter_hook_function(module):
|
||||
if hasattr(module, '__batch_counter_handle__'):
|
||||
return
|
||||
|
||||
handle = module.register_forward_hook(batch_counter_hook)
|
||||
module.__batch_counter_handle__ = handle
|
||||
|
||||
|
||||
def remove_batch_counter_hook_function(module):
|
||||
if hasattr(module, '__batch_counter_handle__'):
|
||||
module.__batch_counter_handle__.remove()
|
||||
del module.__batch_counter_handle__
|
||||
|
||||
|
||||
def add_flops_counter_variable_or_reset(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops__') or hasattr(module, '__params__'):
|
||||
print('Warning: variables __flops__ or __params__ are already '
|
||||
'defined for the module' + type(module).__name__ +
|
||||
' ptflops can affect your code!')
|
||||
module.__ptflops_backup_flops__ = module.__flops__
|
||||
module.__ptflops_backup_params__ = module.__params__
|
||||
module.__flops__ = 0
|
||||
module.__params__ = get_model_parameters_number(module)
|
||||
|
||||
|
||||
def is_supported_instance(module):
|
||||
if type(module) in MODULES_MAPPING or type(module) in CUSTOM_MODULES_MAPPING:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def remove_flops_counter_hook_function(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops_handle__'):
|
||||
module.__flops_handle__.remove()
|
||||
del module.__flops_handle__
|
||||
|
||||
|
||||
def remove_flops_counter_variables(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops__'):
|
||||
del module.__flops__
|
||||
if hasattr(module, '__ptflops_backup_flops__'):
|
||||
module.__flops__ = module.__ptflops_backup_flops__
|
||||
if hasattr(module, '__params__'):
|
||||
del module.__params__
|
||||
if hasattr(module, '__ptflops_backup_params__'):
|
||||
module.__params__ = module.__ptflops_backup_params__
|
||||
|
||||
0
ixformer_sdk/contrib/DeepCache/sd/__init__.py
Normal file
0
ixformer_sdk/contrib/DeepCache/sd/__init__.py
Normal file
812
ixformer_sdk/contrib/DeepCache/sd/pipeline_stable_diffusion.py
Normal file
812
ixformer_sdk/contrib/DeepCache/sd/pipeline_stable_diffusion.py
Normal file
@@ -0,0 +1,812 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import time
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from packaging import version
|
||||
from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer
|
||||
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.loaders import FromSingleFileMixin, LoraLoaderMixin, TextualInversionLoaderMixin
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.models.lora import adjust_lora_scale_text_encoder
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
from diffusers.utils import (
|
||||
deprecate,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker
|
||||
|
||||
from .unet_2d_condition import UNet2DConditionModel
|
||||
from .pipeline_utils import DiffusionPipeline
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import StableDiffusionPipeline
|
||||
|
||||
>>> pipe = StableDiffusionPipeline.from_pretrained("runwayml/stable-diffusion-v1-5", torch_dtype=torch.float16)
|
||||
>>> pipe = pipe.to("cuda")
|
||||
|
||||
>>> prompt = "a photo of an astronaut riding a horse on mars"
|
||||
>>> image = pipe(prompt).images[0]
|
||||
```
|
||||
"""
|
||||
|
||||
def sample_gaussian_centered(n=1000, sample_size=100, std_dev=100):
|
||||
samples = []
|
||||
|
||||
while len(samples) < sample_size:
|
||||
# Sample from a Gaussian centered at n/2
|
||||
sample = int(np.random.normal(loc=n/2, scale=std_dev))
|
||||
|
||||
# Check if the sample is in bounds
|
||||
if 1 <= sample < n and sample not in samples:
|
||||
samples.append(sample)
|
||||
|
||||
return samples
|
||||
|
||||
def sample_from_quad(total_numbers, n_samples, pow=1.2):
|
||||
while pow > 1:
|
||||
# Generate linearly spaced values between 0 and a max value
|
||||
x_values = np.linspace(0, total_numbers**(1/pow), n_samples+1)
|
||||
|
||||
# Raise these values to the power of 1.5 to get a non-linear distribution
|
||||
indices = np.unique(np.int32(x_values**pow))[:-1]
|
||||
if len(indices) == n_samples:
|
||||
break
|
||||
pow -=0.02
|
||||
if pow <= 1:
|
||||
raise ValueError("Cannot find suitable pow. Please adjust n_samples or decrease center.")
|
||||
return indices, pow
|
||||
|
||||
def sample_from_quad_center(total_numbers, n_samples, center, pow=1.2):
|
||||
while pow > 1:
|
||||
# Generate linearly spaced values between 0 and a max value
|
||||
x_values = np.linspace((-center)**(1/pow), (total_numbers-center)**(1/pow), n_samples+1)
|
||||
indices = [0] + [x+center for x in np.unique(np.int32(x_values**pow))[1:-1]]
|
||||
if len(indices) == n_samples:
|
||||
break
|
||||
pow -=0.02
|
||||
if pow <= 1:
|
||||
raise ValueError("Cannot find suitable pow. Please adjust n_samples or decrease center.")
|
||||
return indices, pow
|
||||
|
||||
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
||||
"""
|
||||
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
|
||||
"""
|
||||
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
|
||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
||||
# rescale the results from guidance (fixes overexposure)
|
||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
||||
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
|
||||
noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
|
||||
return noise_cfg
|
||||
|
||||
|
||||
class StableDiffusionPipeline(DiffusionPipeline, TextualInversionLoaderMixin, LoraLoaderMixin, FromSingleFileMixin):
|
||||
r"""
|
||||
Pipeline for text-to-image generation using Stable Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
The pipeline also inherits the following loading methods:
|
||||
- [`~loaders.TextualInversionLoaderMixin.load_textual_inversion`] for loading textual inversion embeddings
|
||||
- [`~loaders.LoraLoaderMixin.load_lora_weights`] for loading LoRA weights
|
||||
- [`~loaders.LoraLoaderMixin.save_lora_weights`] for saving LoRA weights
|
||||
- [`~loaders.FromSingleFileMixin.from_single_file`] for loading `.ckpt` files
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`~transformers.CLIPTextModel`]):
|
||||
Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)).
|
||||
tokenizer ([`~transformers.CLIPTokenizer`]):
|
||||
A `CLIPTokenizer` to tokenize text.
|
||||
unet ([`UNet2DConditionModel`]):
|
||||
A `UNet2DConditionModel` to denoise the encoded image latents.
|
||||
scheduler ([`SchedulerMixin`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents. Can be one of
|
||||
[`DDIMScheduler`], [`LMSDiscreteScheduler`], or [`PNDMScheduler`].
|
||||
safety_checker ([`StableDiffusionSafetyChecker`]):
|
||||
Classification module that estimates whether generated images could be considered offensive or harmful.
|
||||
Please refer to the [model card](https://huggingface.co/runwayml/stable-diffusion-v1-5) for more details
|
||||
about a model's potential harms.
|
||||
feature_extractor ([`~transformers.CLIPImageProcessor`]):
|
||||
A `CLIPImageProcessor` to extract features from generated images; used as inputs to the `safety_checker`.
|
||||
"""
|
||||
model_cpu_offload_seq = "text_encoder->unet->vae"
|
||||
_optional_components = ["safety_checker", "feature_extractor"]
|
||||
_exclude_from_cpu_offload = ["safety_checker"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
unet: UNet2DConditionModel,
|
||||
scheduler: KarrasDiffusionSchedulers,
|
||||
safety_checker: StableDiffusionSafetyChecker,
|
||||
feature_extractor: CLIPImageProcessor,
|
||||
requires_safety_checker: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||
" file"
|
||||
)
|
||||
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["steps_offset"] = 1
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||
)
|
||||
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["clip_sample"] = False
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if safety_checker is None and requires_safety_checker:
|
||||
logger.warning(
|
||||
f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`. Ensure"
|
||||
" that you abide to the conditions of the Stable Diffusion license and do not expose unfiltered"
|
||||
" results in services or applications open to the public. Both the diffusers team and Hugging Face"
|
||||
" strongly recommend to keep the safety filter enabled in all public facing circumstances, disabling"
|
||||
" it only for use-cases that involve analyzing network behavior or auditing its results. For more"
|
||||
" information, please have a look at https://github.com/huggingface/diffusers/pull/254 ."
|
||||
)
|
||||
|
||||
if safety_checker is not None and feature_extractor is None:
|
||||
raise ValueError(
|
||||
"Make sure to define a feature extractor when loading {self.__class__} if you want to use the safety"
|
||||
" checker. If you do not want to use the safety checker, you can pass `'safety_checker=None'` instead."
|
||||
)
|
||||
|
||||
is_unet_version_less_0_9_0 = hasattr(unet.config, "_diffusers_version") and version.parse(
|
||||
version.parse(unet.config._diffusers_version).base_version
|
||||
) < version.parse("0.9.0.dev0")
|
||||
is_unet_sample_size_less_64 = hasattr(unet.config, "sample_size") and unet.config.sample_size < 64
|
||||
if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:
|
||||
deprecation_message = (
|
||||
"The configuration file of the unet has set the default `sample_size` to smaller than"
|
||||
" 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"
|
||||
" following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"
|
||||
" CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"
|
||||
" \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"
|
||||
" configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"
|
||||
" in the config might lead to incorrect results in future versions. If you have downloaded this"
|
||||
" checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"
|
||||
" the `unet/config.json` file"
|
||||
)
|
||||
deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)
|
||||
new_config = dict(unet.config)
|
||||
new_config["sample_size"] = 64
|
||||
unet._internal_dict = FrozenDict(new_config)
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
safety_checker=safety_checker,
|
||||
feature_extractor=feature_extractor,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
self.register_to_config(requires_safety_checker=requires_safety_checker)
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
def _encode_prompt(
|
||||
self,
|
||||
prompt,
|
||||
device,
|
||||
num_images_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt=None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
lora_scale: Optional[float] = None,
|
||||
):
|
||||
deprecation_message = "`_encode_prompt()` is deprecated and it will be removed in a future version. Use `encode_prompt()` instead. Also, be aware that the output format changed from a concatenated tensor to a tuple."
|
||||
deprecate("_encode_prompt()", "1.0.0", deprecation_message, standard_warn=False)
|
||||
|
||||
prompt_embeds_tuple = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||
negative_prompt=negative_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
|
||||
# concatenate for backwards comp
|
||||
prompt_embeds = torch.cat([prompt_embeds_tuple[1], prompt_embeds_tuple[0]])
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt,
|
||||
device,
|
||||
num_images_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt=None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
lora_scale: Optional[float] = None,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
device: (`torch.device`):
|
||||
torch device
|
||||
num_images_per_prompt (`int`):
|
||||
number of images that should be generated per prompt
|
||||
do_classifier_free_guidance (`bool`):
|
||||
whether to use classifier free guidance or not
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
lora_scale (`float`, *optional*):
|
||||
A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded.
|
||||
"""
|
||||
# set lora scale so that monkey patched LoRA
|
||||
# function of text encoder can correctly access it
|
||||
if lora_scale is not None and isinstance(self, LoraLoaderMixin):
|
||||
self._lora_scale = lora_scale
|
||||
|
||||
# dynamically adjust the LoRA scale
|
||||
adjust_lora_scale_text_encoder(self.text_encoder, lora_scale)
|
||||
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
# textual inversion: procecss multi-vector tokens if necessary
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
|
||||
text_input_ids, untruncated_ids
|
||||
):
|
||||
removed_text = self.tokenizer.batch_decode(
|
||||
untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1]
|
||||
)
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {self.tokenizer.model_max_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = text_inputs.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
prompt_embeds = self.text_encoder(
|
||||
text_input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
prompt_embeds = prompt_embeds[0]
|
||||
|
||||
if self.text_encoder is not None:
|
||||
prompt_embeds_dtype = self.text_encoder.dtype
|
||||
elif self.unet is not None:
|
||||
prompt_embeds_dtype = self.unet.dtype
|
||||
else:
|
||||
prompt_embeds_dtype = prompt_embeds.dtype
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
|
||||
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
|
||||
|
||||
# get unconditional embeddings for classifier free guidance
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
uncond_tokens: List[str]
|
||||
if negative_prompt is None:
|
||||
uncond_tokens = [""] * batch_size
|
||||
elif prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif isinstance(negative_prompt, str):
|
||||
uncond_tokens = [negative_prompt]
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
else:
|
||||
uncond_tokens = negative_prompt
|
||||
|
||||
# textual inversion: procecss multi-vector tokens if necessary
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
uncond_tokens = self.maybe_convert_prompt(uncond_tokens, self.tokenizer)
|
||||
|
||||
max_length = prompt_embeds.shape[1]
|
||||
uncond_input = self.tokenizer(
|
||||
uncond_tokens,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||
attention_mask = uncond_input.attention_mask.to(device)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
negative_prompt_embeds = self.text_encoder(
|
||||
uncond_input.input_ids.to(device),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
negative_prompt_embeds = negative_prompt_embeds[0]
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
|
||||
seq_len = negative_prompt_embeds.shape[1]
|
||||
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
|
||||
|
||||
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
|
||||
negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
def run_safety_checker(self, image, device, dtype):
|
||||
if self.safety_checker is None:
|
||||
has_nsfw_concept = None
|
||||
else:
|
||||
if torch.is_tensor(image):
|
||||
feature_extractor_input = self.image_processor.postprocess(image, output_type="pil")
|
||||
else:
|
||||
feature_extractor_input = self.image_processor.numpy_to_pil(image)
|
||||
safety_checker_input = self.feature_extractor(feature_extractor_input, return_tensors="pt").to(device)
|
||||
image, has_nsfw_concept = self.safety_checker(
|
||||
images=image, clip_input=safety_checker_input.pixel_values.to(dtype)
|
||||
)
|
||||
return image, has_nsfw_concept
|
||||
|
||||
def decode_latents(self, latents):
|
||||
deprecation_message = "The decode_latents method is deprecated and will be removed in 1.0.0. Please use VaeImageProcessor.postprocess(...) instead"
|
||||
deprecate("decode_latents", "1.0.0", deprecation_message, standard_warn=False)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
image = image.cpu().permute(0, 2, 3, 1).float().numpy()
|
||||
return image
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
height,
|
||||
width,
|
||||
callback_steps,
|
||||
negative_prompt=None,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
):
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
if (callback_steps is None) or (
|
||||
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
|
||||
if negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
|
||||
def prepare_latents(self, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None):
|
||||
shape = (batch_size, num_channels_latents, height // self.vae_scale_factor, width // self.vae_scale_factor)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: int = 1,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
guidance_rescale: float = 0.0,
|
||||
cache_interval: int = 1,
|
||||
cache_layer_id: int = None,
|
||||
cache_block_id: int = None,
|
||||
uniform: bool = True,
|
||||
pow: float = None,
|
||||
center: int = None,
|
||||
output_all_sequence: bool = False,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide image generation. If not defined, you need to pass `prompt_embeds`.
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.5):
|
||||
A higher guidance scale value encourages the model to generate images closely linked to the text
|
||||
`prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`.
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide what to not include in image generation. If not defined, you need to
|
||||
pass `negative_prompt_embeds` instead. Ignored when not using guidance (`guidance_scale < 1`).
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
eta (`float`, *optional*, defaults to 0.0):
|
||||
Corresponds to parameter eta (η) from the [DDIM](https://arxiv.org/abs/2010.02502) paper. Only applies
|
||||
to the [`~schedulers.DDIMScheduler`], and is ignored in other schedulers.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs (prompt weighting). If
|
||||
not provided, `negative_prompt_embeds` are generated from the `negative_prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||
plain tuple.
|
||||
callback (`Callable`, *optional*):
|
||||
A function that calls every `callback_steps` steps during inference. The function is called with the
|
||||
following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||
callback_steps (`int`, *optional*, defaults to 1):
|
||||
The frequency at which the `callback` function is called. If not specified, the callback is called at
|
||||
every step.
|
||||
cross_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the [`AttentionProcessor`] as defined in
|
||||
[`self.processor`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
guidance_rescale (`float`, *optional*, defaults to 0.7):
|
||||
Guidance rescale factor from [Common Diffusion Noise Schedules and Sample Steps are
|
||||
Flawed](https://arxiv.org/pdf/2305.08891.pdf). Guidance rescale factor should fix overexposure when
|
||||
using zero terminal SNR.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is a list with the generated images and the
|
||||
second element is a list of `bool`s indicating whether the corresponding generated image contains
|
||||
"not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt, height, width, callback_steps, negative_prompt, prompt_embeds, negative_prompt_embeds
|
||||
)
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input prompt
|
||||
text_encoder_lora_scale = (
|
||||
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
|
||||
)
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt,
|
||||
device,
|
||||
num_images_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
negative_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
lora_scale=text_encoder_lora_scale,
|
||||
)
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
if do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds])
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
prv_features = None
|
||||
latents_list = [latents]
|
||||
|
||||
if cache_interval == 1:
|
||||
interval_seq = list(range(num_inference_steps))
|
||||
else:
|
||||
if uniform:
|
||||
interval_seq = list(range(0, num_inference_steps, cache_interval))
|
||||
else:
|
||||
num_slow_step = num_inference_steps//cache_interval
|
||||
if num_inference_steps%cache_interval != 0:
|
||||
num_slow_step += 1
|
||||
|
||||
interval_seq, pow = sample_from_quad_center(num_inference_steps, num_slow_step, center=center, pow=pow)#[0, 3, 6, 9, 12, 16, 22, 28, 35, 43,]
|
||||
#interval_seq, pow = sample_from_quad(num_inference_steps, num_inference_steps//cache_interval, pow=pow)#[0, 3, 6, 9, 12, 16, 22, 28, 35, 43,]
|
||||
|
||||
interval_seq = sorted(interval_seq)
|
||||
#print(interval_seq, len(interval_seq), pow)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
#print("[INFO] Update Feature Interval = {}, Update Layer Number = {}, Update Block Number = {}".format(cache_interval, cache_layer_id, cache_block_id))
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
if i in interval_seq:
|
||||
prv_features = None
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred, prv_features = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
replicate_prv_feature=prv_features,
|
||||
quick_replicate= cache_interval>1,
|
||||
cache_layer_id=cache_layer_id,
|
||||
cache_block_id=cache_block_id,
|
||||
return_dict=False,
|
||||
)
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
if do_classifier_free_guidance and guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
latents_list.append(latents)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
callback(i, t, latents)
|
||||
|
||||
if not output_type == "latent":
|
||||
if output_all_sequence:
|
||||
image = [self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] for latents in latents_list]
|
||||
has_nsfw_concept = None #self.run_safety_checker(images[0], device, prompt_embeds.dtype)
|
||||
num_img = len(image)
|
||||
else:
|
||||
image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0]
|
||||
has_nsfw_concept = None
|
||||
num_img = image.shape[0]
|
||||
else:
|
||||
image = latents
|
||||
has_nsfw_concept = None
|
||||
|
||||
if has_nsfw_concept is None:
|
||||
do_denormalize = [True] * num_img
|
||||
else:
|
||||
do_denormalize = [not has_nsfw for has_nsfw in has_nsfw_concept]
|
||||
|
||||
if output_all_sequence:
|
||||
image = [self.image_processor.postprocess(img, output_type=output_type, do_denormalize=do_denormalize) for img in image]
|
||||
else:
|
||||
image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
if not return_dict:
|
||||
return (image, has_nsfw_concept,)
|
||||
|
||||
return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)
|
||||
741
ixformer_sdk/contrib/DeepCache/sd/pipeline_text_to_video_zero.py
Normal file
741
ixformer_sdk/contrib/DeepCache/sd/pipeline_text_to_video_zero.py
Normal file
@@ -0,0 +1,741 @@
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.functional import grid_sample
|
||||
from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer
|
||||
|
||||
from diffusers.models import AutoencoderKL
|
||||
from .unet_2d_condition import UNet2DConditionModel
|
||||
from .pipeline_stable_diffusion import StableDiffusionPipeline, StableDiffusionSafetyChecker
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
def sample_gaussian_centered(n=1000, sample_size=100, std_dev=100):
|
||||
samples = []
|
||||
|
||||
while len(samples) < sample_size:
|
||||
# Sample from a Gaussian centered at n/2
|
||||
sample = int(np.random.normal(loc=n/2, scale=std_dev))
|
||||
|
||||
# Check if the sample is in bounds
|
||||
if 1 <= sample < n and sample not in samples:
|
||||
samples.append(sample)
|
||||
|
||||
return samples
|
||||
|
||||
def sample_from_quad(total_numbers, n_samples, pow=1.2):
|
||||
while pow > 1:
|
||||
# Generate linearly spaced values between 0 and a max value
|
||||
x_values = np.linspace(0, total_numbers**(1/pow), n_samples+1)
|
||||
|
||||
# Raise these values to the power of 1.5 to get a non-linear distribution
|
||||
indices = np.unique(np.int32(x_values**pow))[:-1]
|
||||
if len(indices) == n_samples:
|
||||
break
|
||||
pow -=0.02
|
||||
if pow <= 1:
|
||||
raise ValueError("Cannot find suitable pow. Please adjust n_samples or decrease center.")
|
||||
return indices, pow
|
||||
|
||||
def sample_from_quad_center(total_numbers, n_samples, center, pow=1.2):
|
||||
while pow > 1:
|
||||
# Generate linearly spaced values between 0 and a max value
|
||||
x_values = np.linspace((-center)**(1/pow), (total_numbers-center)**(1/pow), n_samples+1)
|
||||
indices = [0] + [x+center for x in np.unique(np.int32(x_values**pow))[1:-1]]
|
||||
if len(indices) == n_samples:
|
||||
break
|
||||
pow -=0.02
|
||||
if pow <= 1:
|
||||
raise ValueError("Cannot find suitable pow. Please adjust n_samples or decrease center.")
|
||||
return indices, pow
|
||||
|
||||
def rearrange_0(tensor, f):
|
||||
F, C, H, W = tensor.size()
|
||||
tensor = torch.permute(torch.reshape(tensor, (F // f, f, C, H, W)), (0, 2, 1, 3, 4))
|
||||
return tensor
|
||||
|
||||
|
||||
def rearrange_1(tensor):
|
||||
B, C, F, H, W = tensor.size()
|
||||
return torch.reshape(torch.permute(tensor, (0, 2, 1, 3, 4)), (B * F, C, H, W))
|
||||
|
||||
|
||||
def rearrange_3(tensor, f):
|
||||
F, D, C = tensor.size()
|
||||
return torch.reshape(tensor, (F // f, f, D, C))
|
||||
|
||||
|
||||
def rearrange_4(tensor):
|
||||
B, F, D, C = tensor.size()
|
||||
return torch.reshape(tensor, (B * F, D, C))
|
||||
|
||||
|
||||
class CrossFrameAttnProcessor:
|
||||
"""
|
||||
Cross frame attention processor. Each frame attends the first frame.
|
||||
|
||||
Args:
|
||||
batch_size: The number that represents actual batch size, other than the frames.
|
||||
For example, calling unet with a single prompt and num_images_per_prompt=1, batch_size should be equal to
|
||||
2, due to classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(self, batch_size=2):
|
||||
self.batch_size = batch_size
|
||||
|
||||
def __call__(self, attn, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
is_cross_attention = encoder_hidden_states is not None
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
# Cross Frame Attention
|
||||
if not is_cross_attention:
|
||||
video_length = key.size()[0] // self.batch_size
|
||||
first_frame_index = [0] * video_length
|
||||
|
||||
# rearrange keys to have batch and frames in the 1st and 2nd dims respectively
|
||||
key = rearrange_3(key, video_length)
|
||||
key = key[:, first_frame_index]
|
||||
# rearrange values to have batch and frames in the 1st and 2nd dims respectively
|
||||
value = rearrange_3(value, video_length)
|
||||
value = value[:, first_frame_index]
|
||||
|
||||
# rearrange back to original shape
|
||||
key = rearrange_4(key)
|
||||
value = rearrange_4(value)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CrossFrameAttnProcessor2_0:
|
||||
"""
|
||||
Cross frame attention processor with scaled_dot_product attention of Pytorch 2.0.
|
||||
|
||||
Args:
|
||||
batch_size: The number that represents actual batch size, other than the frames.
|
||||
For example, calling unet with a single prompt and num_images_per_prompt=1, batch_size should be equal to
|
||||
2, due to classifier-free guidance.
|
||||
"""
|
||||
|
||||
def __init__(self, batch_size=2):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
self.batch_size = batch_size
|
||||
|
||||
def __call__(self, attn, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
inner_dim = hidden_states.shape[-1]
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
is_cross_attention = encoder_hidden_states is not None
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
# Cross Frame Attention
|
||||
if not is_cross_attention:
|
||||
video_length = max(1, key.size()[0] // self.batch_size)
|
||||
first_frame_index = [0] * video_length
|
||||
|
||||
# rearrange keys to have batch and frames in the 1st and 2nd dims respectively
|
||||
key = rearrange_3(key, video_length)
|
||||
key = key[:, first_frame_index]
|
||||
# rearrange values to have batch and frames in the 1st and 2nd dims respectively
|
||||
value = rearrange_3(value, video_length)
|
||||
value = value[:, first_frame_index]
|
||||
|
||||
# rearrange back to original shape
|
||||
key = rearrange_4(key)
|
||||
value = rearrange_4(value)
|
||||
|
||||
head_dim = inner_dim // attn.heads
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextToVideoPipelineOutput(BaseOutput):
|
||||
r"""
|
||||
Output class for zero-shot text-to-video pipeline.
|
||||
|
||||
Args:
|
||||
images (`[List[PIL.Image.Image]`, `np.ndarray`]):
|
||||
List of denoised PIL images of length `batch_size` or NumPy array of shape `(batch_size, height, width,
|
||||
num_channels)`.
|
||||
nsfw_content_detected (`[List[bool]]`):
|
||||
List indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content or
|
||||
`None` if safety checking could not be performed.
|
||||
"""
|
||||
|
||||
images: Union[List[PIL.Image.Image], np.ndarray]
|
||||
nsfw_content_detected: Optional[List[bool]]
|
||||
|
||||
|
||||
def coords_grid(batch, ht, wd, device):
|
||||
# Adapted from https://github.com/princeton-vl/RAFT/blob/master/core/utils/utils.py
|
||||
coords = torch.meshgrid(torch.arange(ht, device=device), torch.arange(wd, device=device))
|
||||
coords = torch.stack(coords[::-1], dim=0).float()
|
||||
return coords[None].repeat(batch, 1, 1, 1)
|
||||
|
||||
|
||||
def warp_single_latent(latent, reference_flow):
|
||||
"""
|
||||
Warp latent of a single frame with given flow
|
||||
|
||||
Args:
|
||||
latent: latent code of a single frame
|
||||
reference_flow: flow which to warp the latent with
|
||||
|
||||
Returns:
|
||||
warped: warped latent
|
||||
"""
|
||||
_, _, H, W = reference_flow.size()
|
||||
_, _, h, w = latent.size()
|
||||
coords0 = coords_grid(1, H, W, device=latent.device).to(latent.dtype)
|
||||
|
||||
coords_t0 = coords0 + reference_flow
|
||||
coords_t0[:, 0] /= W
|
||||
coords_t0[:, 1] /= H
|
||||
|
||||
coords_t0 = coords_t0 * 2.0 - 1.0
|
||||
coords_t0 = F.interpolate(coords_t0, size=(h, w), mode="bilinear")
|
||||
coords_t0 = torch.permute(coords_t0, (0, 2, 3, 1))
|
||||
|
||||
warped = grid_sample(latent, coords_t0, mode="nearest", padding_mode="reflection")
|
||||
return warped
|
||||
|
||||
|
||||
def create_motion_field(motion_field_strength_x, motion_field_strength_y, frame_ids, device, dtype):
|
||||
"""
|
||||
Create translation motion field
|
||||
|
||||
Args:
|
||||
motion_field_strength_x: motion strength along x-axis
|
||||
motion_field_strength_y: motion strength along y-axis
|
||||
frame_ids: indexes of the frames the latents of which are being processed.
|
||||
This is needed when we perform chunk-by-chunk inference
|
||||
device: device
|
||||
dtype: dtype
|
||||
|
||||
Returns:
|
||||
|
||||
"""
|
||||
seq_length = len(frame_ids)
|
||||
reference_flow = torch.zeros((seq_length, 2, 512, 512), device=device, dtype=dtype)
|
||||
for fr_idx in range(seq_length):
|
||||
reference_flow[fr_idx, 0, :, :] = motion_field_strength_x * (frame_ids[fr_idx])
|
||||
reference_flow[fr_idx, 1, :, :] = motion_field_strength_y * (frame_ids[fr_idx])
|
||||
return reference_flow
|
||||
|
||||
|
||||
def create_motion_field_and_warp_latents(motion_field_strength_x, motion_field_strength_y, frame_ids, latents):
|
||||
"""
|
||||
Creates translation motion and warps the latents accordingly
|
||||
|
||||
Args:
|
||||
motion_field_strength_x: motion strength along x-axis
|
||||
motion_field_strength_y: motion strength along y-axis
|
||||
frame_ids: indexes of the frames the latents of which are being processed.
|
||||
This is needed when we perform chunk-by-chunk inference
|
||||
latents: latent codes of frames
|
||||
|
||||
Returns:
|
||||
warped_latents: warped latents
|
||||
"""
|
||||
motion_field = create_motion_field(
|
||||
motion_field_strength_x=motion_field_strength_x,
|
||||
motion_field_strength_y=motion_field_strength_y,
|
||||
frame_ids=frame_ids,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
warped_latents = latents.clone().detach()
|
||||
for i in range(len(warped_latents)):
|
||||
warped_latents[i] = warp_single_latent(latents[i][None], motion_field[i][None])
|
||||
return warped_latents
|
||||
|
||||
|
||||
class TextToVideoZeroPipeline(StableDiffusionPipeline):
|
||||
r"""
|
||||
Pipeline for zero-shot text-to-video generation using Stable Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`CLIPTextModel`]):
|
||||
Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)).
|
||||
tokenizer (`CLIPTokenizer`):
|
||||
A [`~transformers.CLIPTokenizer`] to tokenize text.
|
||||
unet ([`UNet2DConditionModel`]):
|
||||
A [`UNet3DConditionModel`] to denoise the encoded video latents.
|
||||
scheduler ([`SchedulerMixin`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents. Can be one of
|
||||
[`DDIMScheduler`], [`LMSDiscreteScheduler`], or [`PNDMScheduler`].
|
||||
safety_checker ([`StableDiffusionSafetyChecker`]):
|
||||
Classification module that estimates whether generated images could be considered offensive or harmful.
|
||||
Please refer to the [model card](https://huggingface.co/runwayml/stable-diffusion-v1-5) for more details
|
||||
about a model's potential harms.
|
||||
feature_extractor ([`CLIPImageProcessor`]):
|
||||
A [`CLIPImageProcessor`] to extract features from generated images; used as inputs to the `safety_checker`.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
unet: UNet2DConditionModel,
|
||||
scheduler: KarrasDiffusionSchedulers,
|
||||
safety_checker: StableDiffusionSafetyChecker,
|
||||
feature_extractor: CLIPImageProcessor,
|
||||
requires_safety_checker: bool = True,
|
||||
):
|
||||
super().__init__(
|
||||
vae, text_encoder, tokenizer, unet, scheduler, safety_checker, feature_extractor, requires_safety_checker
|
||||
)
|
||||
processor = (
|
||||
CrossFrameAttnProcessor2_0(batch_size=2)
|
||||
if hasattr(F, "scaled_dot_product_attention")
|
||||
else CrossFrameAttnProcessor(batch_size=2)
|
||||
)
|
||||
self.unet.set_attn_processor(processor)
|
||||
|
||||
def forward_loop(self, x_t0, t0, t1, generator):
|
||||
"""
|
||||
Perform DDPM forward process from time t0 to t1. This is the same as adding noise with corresponding variance.
|
||||
|
||||
Args:
|
||||
x_t0:
|
||||
Latent code at time t0.
|
||||
t0:
|
||||
Timestep at t0.
|
||||
t1:
|
||||
Timestamp at t1.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
|
||||
Returns:
|
||||
x_t1:
|
||||
Forward process applied to x_t0 from time t0 to t1.
|
||||
"""
|
||||
eps = randn_tensor(x_t0.size(), generator=generator, dtype=x_t0.dtype, device=x_t0.device)
|
||||
alpha_vec = torch.prod(self.scheduler.alphas[t0:t1])
|
||||
x_t1 = torch.sqrt(alpha_vec) * x_t0 + torch.sqrt(1 - alpha_vec) * eps
|
||||
return x_t1
|
||||
|
||||
def backward_loop(
|
||||
self,
|
||||
latents,
|
||||
timesteps,
|
||||
prompt_embeds,
|
||||
guidance_scale,
|
||||
callback,
|
||||
callback_steps,
|
||||
num_warmup_steps,
|
||||
extra_step_kwargs,
|
||||
prv_features,
|
||||
interval_seq,
|
||||
cache_interval,
|
||||
cache_block_id,
|
||||
cache_layer_id,
|
||||
cross_attention_kwargs=None,
|
||||
):
|
||||
"""
|
||||
Perform backward process given list of time steps.
|
||||
|
||||
Args:
|
||||
latents:
|
||||
Latents at time timesteps[0].
|
||||
timesteps:
|
||||
Time steps along which to perform backward process.
|
||||
prompt_embeds:
|
||||
Pre-generated text embeddings.
|
||||
guidance_scale:
|
||||
A higher guidance scale value encourages the model to generate images closely linked to the text
|
||||
`prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`.
|
||||
callback (`Callable`, *optional*):
|
||||
A function that calls every `callback_steps` steps during inference. The function is called with the
|
||||
following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||
callback_steps (`int`, *optional*, defaults to 1):
|
||||
The frequency at which the `callback` function is called. If not specified, the callback is called at
|
||||
every step.
|
||||
extra_step_kwargs:
|
||||
Extra_step_kwargs.
|
||||
cross_attention_kwargs:
|
||||
A kwargs dictionary that if specified is passed along to the [`AttentionProcessor`] as defined in
|
||||
[`self.processor`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
num_warmup_steps:
|
||||
number of warmup steps.
|
||||
|
||||
Returns:
|
||||
latents:
|
||||
Latents of backward process output at time timesteps[-1].
|
||||
"""
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
num_steps = (len(timesteps) - num_warmup_steps) // self.scheduler.order
|
||||
with self.progress_bar(total=num_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
########
|
||||
if i in interval_seq:
|
||||
prv_features = None
|
||||
# predict the noise residual
|
||||
noise_pred, prv_features = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
replicate_prv_feature=prv_features,
|
||||
quick_replicate= cache_interval>1,
|
||||
cache_layer_id=cache_layer_id,
|
||||
cache_block_id=cache_block_id,
|
||||
return_dict=False,
|
||||
)
|
||||
########
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
return latents.clone().detach()
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
video_length: Optional[int] = 8,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
motion_field_strength_x: float = 12,
|
||||
motion_field_strength_y: float = 12,
|
||||
output_type: Optional[str] = "tensor",
|
||||
return_dict: bool = True,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: Optional[int] = 1,
|
||||
t0: int = 44,
|
||||
t1: int = 47,
|
||||
frame_ids: Optional[List[int]] = None,
|
||||
########
|
||||
cache_interval: int = 1,
|
||||
cache_layer_id: int = None,
|
||||
cache_block_id: int = None,
|
||||
uniform: bool = True,
|
||||
pow: float = None,
|
||||
center: int = None,
|
||||
output_all_sequence: bool = False,
|
||||
########
|
||||
):
|
||||
"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide image generation. If not defined, you need to pass `prompt_embeds`.
|
||||
video_length (`int`, *optional*, defaults to 8):
|
||||
The number of generated video frames.
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.5):
|
||||
A higher guidance scale value encourages the model to generate images closely linked to the text
|
||||
`prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`.
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide what to not include in video generation. If not defined, you need to
|
||||
pass `negative_prompt_embeds` instead. Ignored when not using guidance (`guidance_scale < 1`).
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of videos to generate per prompt.
|
||||
eta (`float`, *optional*, defaults to 0.0):
|
||||
Corresponds to parameter eta (η) from the [DDIM](https://arxiv.org/abs/2010.02502) paper. Only applies
|
||||
to the [`~schedulers.DDIMScheduler`], and is ignored in other schedulers.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for video
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
output_type (`str`, *optional*, defaults to `"numpy"`):
|
||||
The output format of the generated video. Choose between `"latent"` and `"numpy"`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a
|
||||
[`~pipelines.text_to_video_synthesis.pipeline_text_to_video_zero.TextToVideoPipelineOutput`] instead of
|
||||
a plain tuple.
|
||||
callback (`Callable`, *optional*):
|
||||
A function that calls every `callback_steps` steps during inference. The function is called with the
|
||||
following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||
callback_steps (`int`, *optional*, defaults to 1):
|
||||
The frequency at which the `callback` function is called. If not specified, the callback is called at
|
||||
every step.
|
||||
motion_field_strength_x (`float`, *optional*, defaults to 12):
|
||||
Strength of motion in generated video along x-axis. See the [paper](https://arxiv.org/abs/2303.13439),
|
||||
Sect. 3.3.1.
|
||||
motion_field_strength_y (`float`, *optional*, defaults to 12):
|
||||
Strength of motion in generated video along y-axis. See the [paper](https://arxiv.org/abs/2303.13439),
|
||||
Sect. 3.3.1.
|
||||
t0 (`int`, *optional*, defaults to 44):
|
||||
Timestep t0. Should be in the range [0, num_inference_steps - 1]. See the
|
||||
[paper](https://arxiv.org/abs/2303.13439), Sect. 3.3.1.
|
||||
t1 (`int`, *optional*, defaults to 47):
|
||||
Timestep t0. Should be in the range [t0 + 1, num_inference_steps - 1]. See the
|
||||
[paper](https://arxiv.org/abs/2303.13439), Sect. 3.3.1.
|
||||
frame_ids (`List[int]`, *optional*):
|
||||
Indexes of the frames that are being generated. This is used when generating longer videos
|
||||
chunk-by-chunk.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.text_to_video_synthesis.pipeline_text_to_video_zero.TextToVideoPipelineOutput`]:
|
||||
The output contains a `ndarray` of the generated video, when `output_type` != `"latent"`, otherwise a
|
||||
latent code of generated videos and a list of `bool`s indicating whether the corresponding generated
|
||||
video contains "not-safe-for-work" (nsfw) content..
|
||||
"""
|
||||
assert video_length > 0
|
||||
if frame_ids is None:
|
||||
frame_ids = list(range(video_length))
|
||||
assert len(frame_ids) == video_length
|
||||
|
||||
assert num_videos_per_prompt == 1
|
||||
|
||||
if isinstance(prompt, str):
|
||||
prompt = [prompt]
|
||||
if isinstance(negative_prompt, str):
|
||||
negative_prompt = [negative_prompt]
|
||||
|
||||
# Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
# Check inputs. Raise error if not correct
|
||||
self.check_inputs(prompt, height, width, callback_steps)
|
||||
|
||||
# Define call parameters
|
||||
batch_size = 1 if isinstance(prompt, str) else len(prompt)
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# Encode input prompt
|
||||
prompt_embeds = self._encode_prompt(
|
||||
prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt
|
||||
)
|
||||
|
||||
# Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
# Prepare extra step kwargs.
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
prv_features = None #record cache feature ****
|
||||
latents_list = [latents]
|
||||
|
||||
if cache_interval == 1:
|
||||
interval_seq = list(range(num_inference_steps))
|
||||
else:
|
||||
if uniform:
|
||||
interval_seq = list(range(0, num_inference_steps, cache_interval))
|
||||
else:
|
||||
num_slow_step = num_inference_steps//cache_interval
|
||||
if num_inference_steps%cache_interval != 0:
|
||||
num_slow_step += 1
|
||||
|
||||
interval_seq, pow = sample_from_quad_center(num_inference_steps, num_slow_step, center=center, pow=pow)#[0, 3, 6, 9, 12, 16, 22, 28, 35, 43,]
|
||||
#interval_seq, pow = sample_from_quad(num_inference_steps, num_inference_steps//cache_interval, pow=pow)#[0, 3, 6, 9, 12, 16, 22, 28, 35, 43,]
|
||||
|
||||
interval_seq = sorted(interval_seq)
|
||||
|
||||
# Perform the first backward process up to time T_1
|
||||
x_1_t1 = self.backward_loop(
|
||||
timesteps=timesteps[: -t1 - 1],
|
||||
prompt_embeds=prompt_embeds,
|
||||
latents=latents,
|
||||
guidance_scale=guidance_scale,
|
||||
callback=callback,
|
||||
callback_steps=callback_steps,
|
||||
extra_step_kwargs=extra_step_kwargs,
|
||||
num_warmup_steps=num_warmup_steps,
|
||||
prv_features=prv_features,
|
||||
interval_seq=interval_seq,
|
||||
cache_interval=cache_interval,
|
||||
cache_block_id=cache_block_id,
|
||||
cache_layer_id=cache_layer_id,
|
||||
)
|
||||
scheduler_copy = copy.deepcopy(self.scheduler)
|
||||
|
||||
# Perform the second backward process up to time T_0
|
||||
x_1_t0 = self.backward_loop(
|
||||
timesteps=timesteps[-t1 - 1 : -t0 - 1],
|
||||
prompt_embeds=prompt_embeds,
|
||||
latents=x_1_t1,
|
||||
guidance_scale=guidance_scale,
|
||||
callback=callback,
|
||||
callback_steps=callback_steps,
|
||||
extra_step_kwargs=extra_step_kwargs,
|
||||
num_warmup_steps=0,
|
||||
prv_features=prv_features,
|
||||
interval_seq=interval_seq,
|
||||
cache_interval=cache_interval,
|
||||
cache_block_id=cache_block_id,
|
||||
cache_layer_id=cache_layer_id,
|
||||
)
|
||||
|
||||
# Propagate first frame latents at time T_0 to remaining frames
|
||||
x_2k_t0 = x_1_t0.repeat(video_length - 1, 1, 1, 1)
|
||||
|
||||
# Add motion in latents at time T_0
|
||||
x_2k_t0 = create_motion_field_and_warp_latents(
|
||||
motion_field_strength_x=motion_field_strength_x,
|
||||
motion_field_strength_y=motion_field_strength_y,
|
||||
latents=x_2k_t0,
|
||||
frame_ids=frame_ids[1:],
|
||||
)
|
||||
|
||||
# Perform forward process up to time T_1
|
||||
x_2k_t1 = self.forward_loop(
|
||||
x_t0=x_2k_t0,
|
||||
t0=timesteps[-t0 - 1].item(),
|
||||
t1=timesteps[-t1 - 1].item(),
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
# Perform backward process from time T_1 to 0
|
||||
x_1k_t1 = torch.cat([x_1_t1, x_2k_t1])
|
||||
b, l, d = prompt_embeds.size()
|
||||
prompt_embeds = prompt_embeds[:, None].repeat(1, video_length, 1, 1).reshape(b * video_length, l, d)
|
||||
|
||||
self.scheduler = scheduler_copy
|
||||
x_1k_0 = self.backward_loop(
|
||||
timesteps=timesteps[-t1 - 1 :],
|
||||
prompt_embeds=prompt_embeds,
|
||||
latents=x_1k_t1,
|
||||
guidance_scale=guidance_scale,
|
||||
callback=callback,
|
||||
callback_steps=callback_steps,
|
||||
extra_step_kwargs=extra_step_kwargs,
|
||||
num_warmup_steps=0,
|
||||
prv_features=prv_features,
|
||||
interval_seq=interval_seq,
|
||||
cache_interval=cache_interval,
|
||||
cache_block_id=cache_block_id,
|
||||
cache_layer_id=cache_layer_id,
|
||||
)
|
||||
latents = x_1k_0
|
||||
|
||||
# manually for max memory savings
|
||||
if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None:
|
||||
self.unet.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if output_type == "latent":
|
||||
image = latents
|
||||
has_nsfw_concept = None
|
||||
else:
|
||||
image = self.decode_latents(latents)
|
||||
# Run safety checker
|
||||
image, has_nsfw_concept = self.run_safety_checker(image, device, prompt_embeds.dtype)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (image, has_nsfw_concept)
|
||||
|
||||
return TextToVideoPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)
|
||||
1839
ixformer_sdk/contrib/DeepCache/sd/pipeline_utils.py
Normal file
1839
ixformer_sdk/contrib/DeepCache/sd/pipeline_utils.py
Normal file
File diff suppressed because it is too large
Load Diff
3296
ixformer_sdk/contrib/DeepCache/sd/unet_2d_blocks.py
Normal file
3296
ixformer_sdk/contrib/DeepCache/sd/unet_2d_blocks.py
Normal file
File diff suppressed because it is too large
Load Diff
1257
ixformer_sdk/contrib/DeepCache/sd/unet_2d_condition.py
Normal file
1257
ixformer_sdk/contrib/DeepCache/sd/unet_2d_condition.py
Normal file
File diff suppressed because it is too large
Load Diff
0
ixformer_sdk/contrib/DeepCache/sdxl/__init__.py
Normal file
0
ixformer_sdk/contrib/DeepCache/sdxl/__init__.py
Normal file
1100
ixformer_sdk/contrib/DeepCache/sdxl/pipeline_stable_diffusion_xl.py
Normal file
1100
ixformer_sdk/contrib/DeepCache/sdxl/pipeline_stable_diffusion_xl.py
Normal file
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
1839
ixformer_sdk/contrib/DeepCache/sdxl/pipeline_utils.py
Normal file
1839
ixformer_sdk/contrib/DeepCache/sdxl/pipeline_utils.py
Normal file
File diff suppressed because it is too large
Load Diff
3339
ixformer_sdk/contrib/DeepCache/sdxl/unet_2d_blocks.py
Normal file
3339
ixformer_sdk/contrib/DeepCache/sdxl/unet_2d_blocks.py
Normal file
File diff suppressed because it is too large
Load Diff
1259
ixformer_sdk/contrib/DeepCache/sdxl/unet_2d_condition.py
Normal file
1259
ixformer_sdk/contrib/DeepCache/sdxl/unet_2d_condition.py
Normal file
File diff suppressed because it is too large
Load Diff
0
ixformer_sdk/contrib/DeepCache/svd/__init__.py
Normal file
0
ixformer_sdk/contrib/DeepCache/svd/__init__.py
Normal file
@@ -0,0 +1,659 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import torch
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.models import AutoencoderKLTemporalDecoder, UNetSpatioTemporalConditionModel
|
||||
from diffusers.schedulers import EulerDiscreteScheduler
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from .pipeline_utils import DiffusionPipeline
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def _append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f"input has {x.ndim} dims but target_dims is {target_dims}, which is less")
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def tensor2vid(video: torch.Tensor, processor, output_type="np"):
|
||||
# Based on:
|
||||
# https://github.com/modelscope/modelscope/blob/1509fdb973e5871f37148a4b5e5964cafd43e64d/modelscope/pipelines/multi_modal/text_to_video_synthesis_pipeline.py#L78
|
||||
|
||||
batch_size, channels, num_frames, height, width = video.shape
|
||||
outputs = []
|
||||
for batch_idx in range(batch_size):
|
||||
batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
batch_output = processor.postprocess(batch_vid, output_type)
|
||||
|
||||
outputs.append(batch_output)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableVideoDiffusionPipelineOutput(BaseOutput):
|
||||
r"""
|
||||
Output class for zero-shot text-to-video pipeline.
|
||||
|
||||
Args:
|
||||
frames (`[List[PIL.Image.Image]`, `np.ndarray`]):
|
||||
List of denoised PIL images of length `batch_size` or NumPy array of shape `(batch_size, height, width,
|
||||
num_channels)`.
|
||||
"""
|
||||
|
||||
frames: Union[List[PIL.Image.Image], np.ndarray]
|
||||
|
||||
|
||||
class StableVideoDiffusionPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline to generate video from an input image using Stable Video Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations.
|
||||
image_encoder ([`~transformers.CLIPVisionModelWithProjection`]):
|
||||
Frozen CLIP image-encoder ([laion/CLIP-ViT-H-14-laion2B-s32B-b79K](https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K)).
|
||||
unet ([`UNetSpatioTemporalConditionModel`]):cache_interval=5, cache_branch=0,
|
||||
A `UNetSpatioTemporalConditionModel` to denoise the encoded image latents.
|
||||
scheduler ([`EulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents.
|
||||
feature_extractor ([`~transformers.CLIPImageProcessor`]):
|
||||
A `CLIPImageProcessor` to extract features from generated images.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "image_encoder->unet->vae"
|
||||
_callback_tensor_inputs = ["latents"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae: AutoencoderKLTemporalDecoder,
|
||||
image_encoder: CLIPVisionModelWithProjection,
|
||||
unet: UNetSpatioTemporalConditionModel,
|
||||
scheduler: EulerDiscreteScheduler,
|
||||
feature_extractor: CLIPImageProcessor,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
|
||||
def _encode_image(self, image, device, num_videos_per_prompt, do_classifier_free_guidance):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.image_processor.pil_to_numpy(image)
|
||||
image = self.image_processor.numpy_to_pt(image)
|
||||
|
||||
# We normalize the image before resizing to match with the original implementation.
|
||||
# Then we unnormalize it after resizing.
|
||||
image = image * 2.0 - 1.0
|
||||
image = _resize_with_antialiasing(image, (224, 224))
|
||||
image = (image + 1.0) / 2.0
|
||||
|
||||
# Normalize the image with for CLIP input
|
||||
image = self.feature_extractor(
|
||||
images=image,
|
||||
do_normalize=True,
|
||||
do_center_crop=False,
|
||||
do_resize=False,
|
||||
do_rescale=False,
|
||||
return_tensors="pt",
|
||||
).pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeddings = self.image_encoder(image).image_embeds
|
||||
image_embeddings = image_embeddings.unsqueeze(1)
|
||||
|
||||
# duplicate image embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = image_embeddings.shape
|
||||
image_embeddings = image_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
image_embeddings = image_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_embeddings = torch.zeros_like(image_embeddings)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_embeddings = torch.cat([negative_image_embeddings, image_embeddings])
|
||||
|
||||
return image_embeddings
|
||||
|
||||
def _encode_vae_image(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
device,
|
||||
num_videos_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
):
|
||||
image = image.to(device=device)
|
||||
image_latents = self.vae.encode(image).latent_dist.mode()
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_latents = torch.zeros_like(image_latents)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_latents = torch.cat([negative_image_latents, image_latents])
|
||||
|
||||
# duplicate image_latents for each generation per prompt, using mps friendly method
|
||||
image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1)
|
||||
|
||||
return image_latents
|
||||
|
||||
def _get_add_time_ids(
|
||||
self,
|
||||
fps,
|
||||
motion_bucket_id,
|
||||
noise_aug_strength,
|
||||
dtype,
|
||||
batch_size,
|
||||
num_videos_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
):
|
||||
add_time_ids = [fps, motion_bucket_id, noise_aug_strength]
|
||||
|
||||
passed_add_embed_dim = self.unet.config.addition_time_embed_dim * len(add_time_ids)
|
||||
expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features
|
||||
|
||||
if expected_add_embed_dim != passed_add_embed_dim:
|
||||
raise ValueError(
|
||||
f"Model expects an added time embedding vector of length {expected_add_embed_dim}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
)
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
add_time_ids = add_time_ids.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
add_time_ids = torch.cat([add_time_ids, add_time_ids])
|
||||
|
||||
return add_time_ids
|
||||
|
||||
def decode_latents(self, latents, num_frames, decode_chunk_size=14):
|
||||
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||
latents = latents.flatten(0, 1)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
|
||||
accepts_num_frames = "num_frames" in set(inspect.signature(self.vae.forward).parameters.keys())
|
||||
|
||||
# decode decode_chunk_size frames at a time to avoid OOM
|
||||
frames = []
|
||||
for i in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[i : i + decode_chunk_size].shape[0]
|
||||
decode_kwargs = {}
|
||||
if accepts_num_frames:
|
||||
# we only pass num_frames_in if it's expected
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
frame = self.vae.decode(latents[i : i + decode_chunk_size], **decode_kwargs).sample
|
||||
frames.append(frame)
|
||||
frames = torch.cat(frames, dim=0)
|
||||
|
||||
# [batch*frames, channels, height, width] -> [batch, channels, frames, height, width]
|
||||
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
frames = frames.float()
|
||||
return frames
|
||||
|
||||
def check_inputs(self, image, height, width):
|
||||
if (
|
||||
not isinstance(image, torch.Tensor)
|
||||
and not isinstance(image, PIL.Image.Image)
|
||||
and not isinstance(image, list)
|
||||
):
|
||||
raise ValueError(
|
||||
"`image` has to be of type `torch.FloatTensor` or `PIL.Image.Image` or `List[PIL.Image.Image]` but is"
|
||||
f" {type(image)}"
|
||||
)
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size,
|
||||
num_frames,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
dtype,
|
||||
device,
|
||||
generator,
|
||||
latents=None,
|
||||
):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_frames,
|
||||
num_channels_latents // 2,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1 and self.unet.config.time_cond_proj_dim is None
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor],
|
||||
height: int = 576,
|
||||
width: int = 1024,
|
||||
num_frames: Optional[int] = None,
|
||||
num_inference_steps: int = 25,
|
||||
min_guidance_scale: float = 1.0,
|
||||
max_guidance_scale: float = 3.0,
|
||||
fps: int = 7,
|
||||
motion_bucket_id: int = 127,
|
||||
noise_aug_strength: int = 0.02,
|
||||
decode_chunk_size: Optional[int] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
cache_interval: Optional[int] = 1,
|
||||
cache_branch: Optional[int] = None,
|
||||
return_dict: bool = True,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image` or `List[PIL.Image.Image]` or `torch.FloatTensor`):
|
||||
Image or images to guide image generation. If you provide a tensor, it needs to be compatible with
|
||||
[`CLIPImageProcessor`](https://huggingface.co/lambdalabs/sd-image-variations-diffusers/blob/main/feature_extractor/preprocessor_config.json).
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, *optional*):
|
||||
The number of video frames to generate. Defaults to 14 for `stable-video-diffusion-img2vid` and to 25 for `stable-video-diffusion-img2vid-xt`
|
||||
num_inference_steps (`int`, *optional*, defaults to 25):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference. This parameter is modulated by `strength`.
|
||||
min_guidance_scale (`float`, *optional*, defaults to 1.0):
|
||||
The minimum guidance scale. Used for the classifier free guidance with first frame.
|
||||
max_guidance_scale (`float`, *optional*, defaults to 3.0):
|
||||
The maximum guidance scale. Used for the classifier free guidance with last frame.
|
||||
fps (`int`, *optional*, defaults to 7):
|
||||
Frames per second. The rate at which the generated images shall be exported to a video after generation.
|
||||
Note that Stable Diffusion Video's UNet was micro-conditioned on fps-1 during training.
|
||||
motion_bucket_id (`int`, *optional*, defaults to 127):
|
||||
The motion bucket ID. Used as conditioning for the generation. The higher the number the more motion will be in the video.
|
||||
noise_aug_strength (`int`, *optional*, defaults to 0.02):
|
||||
The amount of noise added to the init image, the higher it is the less the video will look like the init image. Increase it for more motion.
|
||||
decode_chunk_size (`int`, *optional*):
|
||||
The number of frames to decode at a time. The higher the chunk size, the higher the temporal consistency
|
||||
between frames, but also the higher the memory consumption. By default, the decoder will decode all frames at once
|
||||
for maximal quality. Reduce `decode_chunk_size` to reduce memory usage.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||
plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is a list of list with the generated frames.
|
||||
|
||||
Examples:
|
||||
|
||||
```py
|
||||
from diffusers import StableVideoDiffusionPipeline
|
||||
from diffusers.utils import load_image, export_to_video
|
||||
|
||||
pipe = StableVideoDiffusionPipeline.from_pretrained("stabilityai/stable-video-diffusion-img2vid-xt", torch_dtype=torch.float16, variant="fp16")
|
||||
pipe.to("cuda")
|
||||
|
||||
image = load_image("https://lh3.googleusercontent.com/y-iFOHfLTwkuQSUegpwDdgKmOjRSTvPxat63dQLB25xkTs4lhIbRUFeNBWZzYf370g=s1200")
|
||||
image = image.resize((1024, 576))
|
||||
|
||||
frames = pipe(image, num_frames=25, decode_chunk_size=8).frames[0]
|
||||
export_to_video(frames, "generated.mp4", fps=7)
|
||||
```
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
num_frames = num_frames if num_frames is not None else self.unet.config.num_frames
|
||||
decode_chunk_size = decode_chunk_size if decode_chunk_size is not None else num_frames
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(image, height, width)
|
||||
|
||||
# 2. Define call parameters
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
batch_size = 1
|
||||
elif isinstance(image, list):
|
||||
batch_size = len(image)
|
||||
else:
|
||||
batch_size = image.shape[0]
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = max_guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input image
|
||||
image_embeddings = self._encode_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||
|
||||
# NOTE: Stable Diffusion Video was conditioned on fps - 1, which
|
||||
# is why it is reduced here.
|
||||
# See: https://github.com/Stability-AI/generative-models/blob/ed0997173f98eaf8f4edf7ba5fe8f15c6b877fd3/scripts/sampling/simple_video_sample.py#L188
|
||||
fps = fps - 1
|
||||
|
||||
# 4. Encode input image using VAE
|
||||
image = self.image_processor.preprocess(image, height=height, width=width)
|
||||
noise = randn_tensor(image.shape, generator=generator, device=image.device, dtype=image.dtype)
|
||||
image = image + noise_aug_strength * noise
|
||||
|
||||
needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=torch.float32)
|
||||
|
||||
image_latents = self._encode_vae_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||
image_latents = image_latents.to(image_embeddings.dtype)
|
||||
|
||||
# cast back to fp16 if needed
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=torch.float16)
|
||||
|
||||
# Repeat the image latents for each frame so we can concatenate them with the noise
|
||||
# image_latents [batch, channels, height, width] ->[batch, num_frames, channels, height, width]
|
||||
image_latents = image_latents.unsqueeze(1).repeat(1, num_frames, 1, 1, 1)
|
||||
|
||||
# 5. Get Added Time IDs
|
||||
added_time_ids = self._get_add_time_ids(
|
||||
fps,
|
||||
motion_bucket_id,
|
||||
noise_aug_strength,
|
||||
image_embeddings.dtype,
|
||||
batch_size,
|
||||
num_videos_per_prompt,
|
||||
do_classifier_free_guidance,
|
||||
)
|
||||
added_time_ids = added_time_ids.to(device)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_frames,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
image_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 7. Prepare guidance scale
|
||||
guidance_scale = torch.linspace(min_guidance_scale, max_guidance_scale, num_frames).unsqueeze(0)
|
||||
guidance_scale = guidance_scale.to(device, latents.dtype)
|
||||
guidance_scale = guidance_scale.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
guidance_scale = _append_dims(guidance_scale, latents.ndim)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
|
||||
cache_features = None
|
||||
interval_seq = list(range(0, num_inference_steps, cache_interval))
|
||||
interval_seq = sorted(interval_seq)
|
||||
|
||||
# 8. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# Concatenate image_latents over channels dimention
|
||||
latent_model_input = torch.cat([latent_model_input, image_latents], dim=2)
|
||||
|
||||
if i in interval_seq:
|
||||
cache_features = None
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred, cache_features = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings,
|
||||
added_time_ids=added_time_ids,
|
||||
cache_features=cache_features,
|
||||
cache_branch=cache_branch,
|
||||
return_dict=False,
|
||||
)
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents).prev_sample
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if not output_type == "latent":
|
||||
# cast back to fp16 if needed
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=torch.float16)
|
||||
frames = self.decode_latents(latents, num_frames, decode_chunk_size)
|
||||
frames = tensor2vid(frames, self.image_processor, output_type=output_type)
|
||||
else:
|
||||
frames = latents
|
||||
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return frames
|
||||
|
||||
return StableVideoDiffusionPipelineOutput(frames=frames)
|
||||
|
||||
|
||||
# resizing utils
|
||||
# TODO: clean up later
|
||||
def _resize_with_antialiasing(input, size, interpolation="bicubic", align_corners=True):
|
||||
h, w = input.shape[-2:]
|
||||
factors = (h / size[0], w / size[1])
|
||||
|
||||
# First, we have to determine sigma
|
||||
# Taken from skimage: https://github.com/scikit-image/scikit-image/blob/v0.19.2/skimage/transform/_warps.py#L171
|
||||
sigmas = (
|
||||
max((factors[0] - 1.0) / 2.0, 0.001),
|
||||
max((factors[1] - 1.0) / 2.0, 0.001),
|
||||
)
|
||||
|
||||
# Now kernel size. Good results are for 3 sigma, but that is kind of slow. Pillow uses 1 sigma
|
||||
# https://github.com/python-pillow/Pillow/blob/master/src/libImaging/Resample.c#L206
|
||||
# But they do it in the 2 passes, which gives better results. Let's try 2 sigmas for now
|
||||
ks = int(max(2.0 * 2 * sigmas[0], 3)), int(max(2.0 * 2 * sigmas[1], 3))
|
||||
|
||||
# Make sure it is odd
|
||||
if (ks[0] % 2) == 0:
|
||||
ks = ks[0] + 1, ks[1]
|
||||
|
||||
if (ks[1] % 2) == 0:
|
||||
ks = ks[0], ks[1] + 1
|
||||
|
||||
input = _gaussian_blur2d(input, ks, sigmas)
|
||||
|
||||
output = torch.nn.functional.interpolate(input, size=size, mode=interpolation, align_corners=align_corners)
|
||||
return output
|
||||
|
||||
|
||||
def _compute_padding(kernel_size):
|
||||
"""Compute padding tuple."""
|
||||
# 4 or 6 ints: (padding_left, padding_right,padding_top,padding_bottom)
|
||||
# https://pytorch.org/docs/stable/nn.html#torch.nn.functional.pad
|
||||
if len(kernel_size) < 2:
|
||||
raise AssertionError(kernel_size)
|
||||
computed = [k - 1 for k in kernel_size]
|
||||
|
||||
# for even kernels we need to do asymmetric padding :(
|
||||
out_padding = 2 * len(kernel_size) * [0]
|
||||
|
||||
for i in range(len(kernel_size)):
|
||||
computed_tmp = computed[-(i + 1)]
|
||||
|
||||
pad_front = computed_tmp // 2
|
||||
pad_rear = computed_tmp - pad_front
|
||||
|
||||
out_padding[2 * i + 0] = pad_front
|
||||
out_padding[2 * i + 1] = pad_rear
|
||||
|
||||
return out_padding
|
||||
|
||||
|
||||
def _filter2d(input, kernel):
|
||||
# prepare kernel
|
||||
b, c, h, w = input.shape
|
||||
tmp_kernel = kernel[:, None, ...].to(device=input.device, dtype=input.dtype)
|
||||
|
||||
tmp_kernel = tmp_kernel.expand(-1, c, -1, -1)
|
||||
|
||||
height, width = tmp_kernel.shape[-2:]
|
||||
|
||||
padding_shape: list[int] = _compute_padding([height, width])
|
||||
input = torch.nn.functional.pad(input, padding_shape, mode="reflect")
|
||||
|
||||
# kernel and input tensor reshape to align element-wise or batch-wise params
|
||||
tmp_kernel = tmp_kernel.reshape(-1, 1, height, width)
|
||||
input = input.view(-1, tmp_kernel.size(0), input.size(-2), input.size(-1))
|
||||
|
||||
# convolve the tensor with the kernel.
|
||||
output = torch.nn.functional.conv2d(input, tmp_kernel, groups=tmp_kernel.size(0), padding=0, stride=1)
|
||||
|
||||
out = output.view(b, c, h, w)
|
||||
return out
|
||||
|
||||
|
||||
def _gaussian(window_size: int, sigma):
|
||||
if isinstance(sigma, float):
|
||||
sigma = torch.tensor([[sigma]])
|
||||
|
||||
batch_size = sigma.shape[0]
|
||||
|
||||
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1)
|
||||
|
||||
if window_size % 2 == 0:
|
||||
x = x + 0.5
|
||||
|
||||
gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0)))
|
||||
|
||||
return gauss / gauss.sum(-1, keepdim=True)
|
||||
|
||||
|
||||
def _gaussian_blur2d(input, kernel_size, sigma):
|
||||
if isinstance(sigma, tuple):
|
||||
sigma = torch.tensor([sigma], dtype=input.dtype)
|
||||
else:
|
||||
sigma = sigma.to(dtype=input.dtype)
|
||||
|
||||
ky, kx = int(kernel_size[0]), int(kernel_size[1])
|
||||
bs = sigma.shape[0]
|
||||
kernel_x = _gaussian(kx, sigma[:, 1].view(bs, 1))
|
||||
kernel_y = _gaussian(ky, sigma[:, 0].view(bs, 1))
|
||||
out_x = _filter2d(input, kernel_x[..., None, :])
|
||||
out = _filter2d(out_x, kernel_y[..., None])
|
||||
|
||||
return out
|
||||
2108
ixformer_sdk/contrib/DeepCache/svd/pipeline_utils.py
Normal file
2108
ixformer_sdk/contrib/DeepCache/svd/pipeline_utils.py
Normal file
File diff suppressed because it is too large
Load Diff
2412
ixformer_sdk/contrib/DeepCache/svd/unet_3d_blocks.py
Normal file
2412
ixformer_sdk/contrib/DeepCache/svd/unet_3d_blocks.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,566 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import UNet2DConditionLoadersMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.models.attention_processor import CROSS_ATTENTION_PROCESSORS, AttentionProcessor, AttnProcessor
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
from .unet_3d_blocks import UNetMidBlockSpatioTemporal, get_down_block, get_up_block
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class UNetSpatioTemporalConditionOutput(BaseOutput):
|
||||
"""
|
||||
The output of [`UNetSpatioTemporalConditionModel`].
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor` of shape `(batch_size, num_frames, num_channels, height, width)`):
|
||||
The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model.
|
||||
"""
|
||||
|
||||
sample: torch.FloatTensor = None
|
||||
|
||||
|
||||
class UNetSpatioTemporalConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin):
|
||||
r"""
|
||||
A conditional Spatio-Temporal UNet model that takes a noisy video frames, conditional state, and a timestep and returns a sample
|
||||
shaped output.
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
|
||||
Parameters:
|
||||
sample_size (`int` or `Tuple[int, int]`, *optional*, defaults to `None`):
|
||||
Height and width of input/output sample.
|
||||
in_channels (`int`, *optional*, defaults to 8): Number of channels in the input sample.
|
||||
out_channels (`int`, *optional*, defaults to 4): Number of channels in the output.
|
||||
down_block_types (`Tuple[str]`, *optional*, defaults to `("CrossAttnDownBlockSpatioTemporal", "CrossAttnDownBlockSpatioTemporal", "CrossAttnDownBlockSpatioTemporal", "DownBlockSpatioTemporal")`):
|
||||
The tuple of downsample blocks to use.
|
||||
up_block_types (`Tuple[str]`, *optional*, defaults to `("UpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal")`):
|
||||
The tuple of upsample blocks to use.
|
||||
block_out_channels (`Tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`):
|
||||
The tuple of output channels for each block.
|
||||
addition_time_embed_dim: (`int`, defaults to 256):
|
||||
Dimension to to encode the additional time ids.
|
||||
projection_class_embeddings_input_dim (`int`, defaults to 768):
|
||||
The dimension of the projection of encoded `added_time_ids`.
|
||||
layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block.
|
||||
cross_attention_dim (`int` or `Tuple[int]`, *optional*, defaults to 1280):
|
||||
The dimension of the cross attention features.
|
||||
transformer_layers_per_block (`int`, `Tuple[int]`, or `Tuple[Tuple]` , *optional*, defaults to 1):
|
||||
The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for
|
||||
[`~models.unet_3d_blocks.CrossAttnDownBlockSpatioTemporal`], [`~models.unet_3d_blocks.CrossAttnUpBlockSpatioTemporal`],
|
||||
[`~models.unet_3d_blocks.UNetMidBlockSpatioTemporal`].
|
||||
num_attention_heads (`int`, `Tuple[int]`, defaults to `(5, 10, 10, 20)`):
|
||||
The number of attention heads.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
sample_size: Optional[int] = None,
|
||||
in_channels: int = 8,
|
||||
out_channels: int = 4,
|
||||
down_block_types: Tuple[str] = (
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"DownBlockSpatioTemporal",
|
||||
),
|
||||
up_block_types: Tuple[str] = (
|
||||
"UpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
),
|
||||
block_out_channels: Tuple[int] = (320, 640, 1280, 1280),
|
||||
addition_time_embed_dim: int = 256,
|
||||
projection_class_embeddings_input_dim: int = 768,
|
||||
layers_per_block: Union[int, Tuple[int]] = 2,
|
||||
cross_attention_dim: Union[int, Tuple[int]] = 1024,
|
||||
transformer_layers_per_block: Union[int, Tuple[int], Tuple[Tuple]] = 1,
|
||||
num_attention_heads: Union[int, Tuple[int]] = (5, 10, 10, 20),
|
||||
num_frames: int = 25,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.sample_size = sample_size
|
||||
|
||||
# Check inputs
|
||||
if len(down_block_types) != len(up_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}."
|
||||
)
|
||||
|
||||
if len(block_out_channels) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `cross_attention_dim` as `down_block_types`. `cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `layers_per_block` as `down_block_types`. `layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
# input
|
||||
self.conv_in = nn.Conv2d(
|
||||
in_channels,
|
||||
block_out_channels[0],
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
# time
|
||||
time_embed_dim = block_out_channels[0] * 4
|
||||
|
||||
self.time_proj = Timesteps(block_out_channels[0], True, downscale_freq_shift=0)
|
||||
timestep_input_dim = block_out_channels[0]
|
||||
|
||||
self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||
|
||||
self.add_time_proj = Timesteps(addition_time_embed_dim, True, downscale_freq_shift=0)
|
||||
self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim)
|
||||
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
if isinstance(num_attention_heads, int):
|
||||
num_attention_heads = (num_attention_heads,) * len(down_block_types)
|
||||
|
||||
if isinstance(cross_attention_dim, int):
|
||||
cross_attention_dim = (cross_attention_dim,) * len(down_block_types)
|
||||
|
||||
if isinstance(layers_per_block, int):
|
||||
layers_per_block = [layers_per_block] * len(down_block_types)
|
||||
|
||||
if isinstance(transformer_layers_per_block, int):
|
||||
transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types)
|
||||
|
||||
blocks_time_embed_dim = time_embed_dim
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
down_block = get_down_block(
|
||||
down_block_type,
|
||||
num_layers=layers_per_block[i],
|
||||
transformer_layers_per_block=transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_downsample=not is_final_block,
|
||||
resnet_eps=1e-5,
|
||||
cross_attention_dim=cross_attention_dim[i],
|
||||
num_attention_heads=num_attention_heads[i],
|
||||
resnet_act_fn="silu",
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
# mid
|
||||
self.mid_block = UNetMidBlockSpatioTemporal(
|
||||
block_out_channels[-1],
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
transformer_layers_per_block=transformer_layers_per_block[-1],
|
||||
cross_attention_dim=cross_attention_dim[-1],
|
||||
num_attention_heads=num_attention_heads[-1],
|
||||
)
|
||||
|
||||
# count how many layers upsample the images
|
||||
self.num_upsamplers = 0
|
||||
|
||||
# up
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
reversed_num_attention_heads = list(reversed(num_attention_heads))
|
||||
reversed_layers_per_block = list(reversed(layers_per_block))
|
||||
reversed_cross_attention_dim = list(reversed(cross_attention_dim))
|
||||
reversed_transformer_layers_per_block = list(reversed(transformer_layers_per_block))
|
||||
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)]
|
||||
|
||||
# add upsample block for all BUT final layer
|
||||
if not is_final_block:
|
||||
add_upsample = True
|
||||
self.num_upsamplers += 1
|
||||
else:
|
||||
add_upsample = False
|
||||
|
||||
up_block = get_up_block(
|
||||
up_block_type,
|
||||
num_layers=reversed_layers_per_block[i] + 1,
|
||||
transformer_layers_per_block=reversed_transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
prev_output_channel=prev_output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_upsample=add_upsample,
|
||||
resnet_eps=1e-5,
|
||||
resolution_idx=i,
|
||||
cross_attention_dim=reversed_cross_attention_dim[i],
|
||||
num_attention_heads=reversed_num_attention_heads[i],
|
||||
resnet_act_fn="silu",
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
# out
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=32, eps=1e-5)
|
||||
self.conv_act = nn.SiLU()
|
||||
|
||||
self.conv_out = nn.Conv2d(
|
||||
block_out_channels[0],
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
@property
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(
|
||||
name: str,
|
||||
module: torch.nn.Module,
|
||||
processors: Dict[str, AttentionProcessor],
|
||||
):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
self.set_attn_processor(processor)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if hasattr(module, "gradient_checkpointing"):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
# Copied from diffusers.models.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking
|
||||
def enable_forward_chunking(self, chunk_size: Optional[int] = None, dim: int = 0) -> None:
|
||||
"""
|
||||
Sets the attention processor to use [feed forward
|
||||
chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers).
|
||||
|
||||
Parameters:
|
||||
chunk_size (`int`, *optional*):
|
||||
The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually
|
||||
over each tensor of dim=`dim`.
|
||||
dim (`int`, *optional*, defaults to `0`):
|
||||
The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch)
|
||||
or dim=1 (sequence length).
|
||||
"""
|
||||
if dim not in [0, 1]:
|
||||
raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}")
|
||||
|
||||
# By default chunk size is 1
|
||||
chunk_size = chunk_size or 1
|
||||
|
||||
def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int):
|
||||
if hasattr(module, "set_chunk_feed_forward"):
|
||||
module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim)
|
||||
|
||||
for child in module.children():
|
||||
fn_recursive_feed_forward(child, chunk_size, dim)
|
||||
|
||||
for module in self.children():
|
||||
fn_recursive_feed_forward(module, chunk_size, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[torch.Tensor, float, int],
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
added_time_ids: torch.Tensor,
|
||||
cache_features: Optional[torch.Tensor] = None,
|
||||
cache_branch: Optional[int] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[UNetSpatioTemporalConditionOutput, Tuple]:
|
||||
r"""
|
||||
The [`UNetSpatioTemporalConditionModel`] forward method.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The noisy input tensor with the following shape `(batch, num_frames, channel, height, width)`.
|
||||
timestep (`torch.FloatTensor` or `float` or `int`): The number of timesteps to denoise an input.
|
||||
encoder_hidden_states (`torch.FloatTensor`):
|
||||
The encoder hidden states with shape `(batch, sequence_length, cross_attention_dim)`.
|
||||
added_time_ids: (`torch.FloatTensor`):
|
||||
The additional time ids with shape `(batch, num_additional_ids)`. These are encoded with sinusoidal
|
||||
embeddings and added to the time embeddings.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] instead of a plain
|
||||
tuple.
|
||||
Returns:
|
||||
[`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] or `tuple`:
|
||||
If `return_dict` is True, an [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] is returned, otherwise
|
||||
a `tuple` is returned where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. time
|
||||
timesteps = timestep
|
||||
if not torch.is_tensor(timesteps):
|
||||
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||
# This would be a good case for the `match` statement (Python 3.10+)
|
||||
is_mps = sample.device.type == "mps"
|
||||
if isinstance(timestep, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
|
||||
elif len(timesteps.shape) == 0:
|
||||
timesteps = timesteps[None].to(sample.device)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
batch_size, num_frames = sample.shape[:2]
|
||||
timesteps = timesteps.expand(batch_size)
|
||||
|
||||
t_emb = self.time_proj(timesteps)
|
||||
|
||||
# `Timesteps` does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=sample.dtype)
|
||||
|
||||
emb = self.time_embedding(t_emb)
|
||||
|
||||
time_embeds = self.add_time_proj(added_time_ids.flatten())
|
||||
time_embeds = time_embeds.reshape((batch_size, -1))
|
||||
time_embeds = time_embeds.to(emb.dtype)
|
||||
aug_emb = self.add_embedding(time_embeds)
|
||||
emb = emb + aug_emb
|
||||
|
||||
# Flatten the batch and frames dimensions
|
||||
# sample: [batch, frames, channels, height, width] -> [batch * frames, channels, height, width]
|
||||
sample = sample.flatten(0, 1)
|
||||
# Repeat the embeddings num_video_frames times
|
||||
# emb: [batch, channels] -> [batch * frames, channels]
|
||||
emb = emb.repeat_interleave(num_frames, dim=0)
|
||||
# encoder_hidden_states: [batch, 1, channels] -> [batch * frames, 1, channels]
|
||||
encoder_hidden_states = encoder_hidden_states.repeat_interleave(num_frames, dim=0)
|
||||
|
||||
# 2. pre-process
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
image_only_indicator = torch.zeros(batch_size, num_frames, dtype=sample.dtype, device=sample.device)
|
||||
|
||||
# Branch: 4 down_blocks, each with 3 skip connections. Here we ignore the first skip branch, whose computations only has up_blocks but without down_blocks.
|
||||
if cache_branch is not None:
|
||||
each_module_num = len(self.down_blocks[0].resnets) + 1
|
||||
down_cache_block_idx = cache_branch // each_module_num
|
||||
down_cache_module_idx = cache_branch % each_module_num
|
||||
|
||||
up_cache_block_idx = len(self.up_blocks) - 1 - down_cache_block_idx
|
||||
up_cache_module_idx = 1 - down_cache_module_idx
|
||||
if down_cache_module_idx == each_module_num - 1:
|
||||
up_cache_block_idx -= 1
|
||||
up_cache_module_idx = 2
|
||||
|
||||
if cache_features is not None:
|
||||
# 3. down
|
||||
down_block_res_samples = (sample,)
|
||||
for block_id, downsample_block in enumerate(self.down_blocks):
|
||||
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
exist_module_idx=down_cache_module_idx if down_cache_block_idx == block_id else None
|
||||
)
|
||||
else:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
image_only_indicator=image_only_indicator,
|
||||
exist_module_idx=down_cache_module_idx if down_cache_block_idx == block_id else None
|
||||
)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
if down_cache_block_idx == block_id:
|
||||
break
|
||||
|
||||
# 4. no mid
|
||||
sample = cache_features
|
||||
|
||||
# 5. up
|
||||
for i, upsample_block in enumerate(self.up_blocks):
|
||||
if i < up_cache_block_idx:
|
||||
continue
|
||||
|
||||
if i == up_cache_block_idx:
|
||||
trunc_res_samples_len = len(upsample_block.resnets) - up_cache_module_idx
|
||||
else:
|
||||
trunc_res_samples_len = len(upsample_block.resnets)
|
||||
|
||||
res_samples = down_block_res_samples[-trunc_res_samples_len :]
|
||||
down_block_res_samples = down_block_res_samples[: -trunc_res_samples_len]
|
||||
|
||||
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||
sample, _ = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
enter_module_idx=up_cache_module_idx if i == up_cache_block_idx else None
|
||||
)
|
||||
else:
|
||||
sample, _ = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
image_only_indicator=image_only_indicator,
|
||||
enter_module_idx=up_cache_module_idx if i == up_cache_block_idx else None
|
||||
)
|
||||
else:
|
||||
# 3. down
|
||||
down_block_res_samples = (sample,)
|
||||
for downsample_block in self.down_blocks:
|
||||
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
else:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
|
||||
# 4. mid
|
||||
sample = self.mid_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
# 5. up
|
||||
for i, upsample_block in enumerate(self.up_blocks):
|
||||
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||
down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]
|
||||
|
||||
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||
sample, current_record_f = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
else:
|
||||
sample, current_record_f = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
if cache_branch is not None and i == up_cache_block_idx:
|
||||
cache_features = current_record_f[up_cache_module_idx]
|
||||
|
||||
# 6. post-process
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
# 7. Reshape back to original shape
|
||||
sample = sample.reshape(batch_size, num_frames, *sample.shape[1:])
|
||||
|
||||
if not return_dict:
|
||||
return (sample, cache_features)
|
||||
|
||||
return UNetSpatioTemporalConditionOutput(sample=sample)
|
||||
Reference in New Issue
Block a user