update
@@ -2,8 +2,47 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
import warnings
|
||||||
|
import platform
|
||||||
|
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||||
|
backends = []
|
||||||
|
|
||||||
from torch.backends.cuda import sdp_kernel
|
def get_sdpa_settings():
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
old_gpu = torch.cuda.get_device_properties(0).major < 7
|
||||||
|
# only use Flash Attention on Ampere (8.0) or newer GPUs
|
||||||
|
use_flash_attn = torch.cuda.get_device_properties(0).major >= 8 and platform.system() == 'Linux'
|
||||||
|
if not use_flash_attn:
|
||||||
|
warnings.warn(
|
||||||
|
"Flash Attention is disabled as it requires a GPU with Ampere (8.0) CUDA capability.",
|
||||||
|
category=UserWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
|
# keep math kernel for PyTorch versions before 2.2 (Flash Attention v2 is only
|
||||||
|
# available on PyTorch 2.2+, while Flash Attention v1 cannot handle all cases)
|
||||||
|
pytorch_version = tuple(int(v) for v in torch.__version__.split(".")[:2])
|
||||||
|
if pytorch_version < (2, 2):
|
||||||
|
warnings.warn(
|
||||||
|
f"You are using PyTorch {torch.__version__} without Flash Attention v2 support. "
|
||||||
|
"Consider upgrading to PyTorch 2.2+ for Flash Attention v2 (which could be faster).",
|
||||||
|
category=UserWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
|
math_kernel_on = pytorch_version < (2, 2) or not use_flash_attn
|
||||||
|
else:
|
||||||
|
old_gpu = True
|
||||||
|
use_flash_attn = False
|
||||||
|
math_kernel_on = True
|
||||||
|
|
||||||
|
return old_gpu, use_flash_attn, math_kernel_on
|
||||||
|
|
||||||
|
OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()
|
||||||
|
if USE_FLASH_ATTN:
|
||||||
|
backends.append(SDPBackend.FLASH_ATTENTION)
|
||||||
|
if MATH_KERNEL_ON:
|
||||||
|
backends.append(SDPBackend.MATH)
|
||||||
|
if OLD_GPU or not USE_FLASH_ATTN:
|
||||||
|
backends.append(SDPBackend.EFFICIENT_ATTENTION)
|
||||||
|
|
||||||
def remove_all_hooks(model: torch.nn.Module) -> None:
|
def remove_all_hooks(model: torch.nn.Module) -> None:
|
||||||
for child in model.children():
|
for child in model.children():
|
||||||
@@ -167,7 +206,9 @@ class Reference(Operator):
|
|||||||
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||||
# Attention
|
# Attention
|
||||||
N = q.shape[-2]
|
N = q.shape[-2]
|
||||||
with sdp_kernel(**{"enable_math": True, "enable_flash": True, "enable_mem_efficient": True}):
|
if not MATH_KERNEL_ON and OLD_GPU:
|
||||||
|
backends.append(SDPBackend.MATH)
|
||||||
|
with sdpa_kernel(backends):
|
||||||
if layer_ind > 12 or self.mode == 'normal':
|
if layer_ind > 12 or self.mode == 'normal':
|
||||||
attn_bias = None
|
attn_bias = None
|
||||||
else:
|
else:
|
||||||
@@ -177,9 +218,7 @@ class Reference(Operator):
|
|||||||
amplify = amplify.log()
|
amplify = amplify.log()
|
||||||
attn_bias[:, :, :, :N] = amplify[:, :, :, [0]]
|
attn_bias[:, :, :, :N] = amplify[:, :, :, [0]]
|
||||||
attn_bias[:, :, :, N:] = amplify[:, :, :, [1]]
|
attn_bias[:, :, :, N:] = amplify[:, :, :, [1]]
|
||||||
out = F.scaled_dot_product_attention(
|
out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_bias)
|
||||||
q, k, v, attn_mask=attn_bias,
|
|
||||||
)
|
|
||||||
del q, k, v
|
del q, k, v
|
||||||
|
|
||||||
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
out = rearrange(out, "b h n d -> b n (h d)", h=h)
|
||||||
|
|||||||
@@ -1,11 +1,8 @@
|
|||||||
# @title Sampling function
|
# @title Sampling function
|
||||||
import math
|
import math
|
||||||
import os
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
import copy
|
import copy
|
||||||
|
|
||||||
import cv2
|
|
||||||
import numpy as np
|
|
||||||
import torch
|
import torch
|
||||||
from einops import rearrange, repeat
|
from einops import rearrange, repeat
|
||||||
|
|
||||||
@@ -115,7 +112,7 @@ def decode_video(model, device, latents, arg):
|
|||||||
end = 0
|
end = 0
|
||||||
|
|
||||||
i = 0
|
i = 0
|
||||||
|
comfy_pbar = ProgressBar(N)
|
||||||
with torch.autocast('cuda'):
|
with torch.autocast('cuda'):
|
||||||
while end < N:
|
while end < N:
|
||||||
start = i * (B - f - olap) + f
|
start = i * (B - f - olap) + f
|
||||||
@@ -131,6 +128,7 @@ def decode_video(model, device, latents, arg):
|
|||||||
else:
|
else:
|
||||||
outputs = torch.cat([ outputs, out[f+olap:] ])
|
outputs = torch.cat([ outputs, out[f+olap:] ])
|
||||||
i += 1
|
i += 1
|
||||||
|
comfy_pbar.update(1)
|
||||||
|
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
@@ -201,6 +199,7 @@ def sample(
|
|||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
with torch.autocast('cuda'):
|
with torch.autocast('cuda'):
|
||||||
# Prepare conditions
|
# Prepare conditions
|
||||||
|
print("Preparing conditions...")
|
||||||
c, uc, additional_model_inputs = get_conditioning(
|
c, uc, additional_model_inputs = get_conditioning(
|
||||||
model,
|
model,
|
||||||
get_unique_embedder_keys_from_conditioner(model.conditioner),
|
get_unique_embedder_keys_from_conditioner(model.conditioner),
|
||||||
@@ -212,6 +211,7 @@ def sample(
|
|||||||
controls=controls, palette=palette, anchor=anchor, first_control=first_control,
|
controls=controls, palette=palette, anchor=anchor, first_control=first_control,
|
||||||
additional_conditions=additional_conditions,
|
additional_conditions=additional_conditions,
|
||||||
)
|
)
|
||||||
|
print("conditions prepared")
|
||||||
# Initial noise
|
# Initial noise
|
||||||
if x_T is None:
|
if x_T is None:
|
||||||
randn = torch.randn(shape, dtype=torch.float32, device="cpu").to(device)
|
randn = torch.randn(shape, dtype=torch.float32, device="cpu").to(device)
|
||||||
@@ -435,29 +435,6 @@ def guidance(denoised, scale, num_frames):
|
|||||||
|
|
||||||
return denoised
|
return denoised
|
||||||
|
|
||||||
|
|
||||||
def write_video(output_folder, fps_id, samples):
|
|
||||||
os.makedirs(output_folder, exist_ok=True)
|
|
||||||
video_path = os.path.join(output_folder, f".mp4")
|
|
||||||
writer = cv2.VideoWriter(
|
|
||||||
video_path,
|
|
||||||
cv2.VideoWriter_fourcc(*"MP4V"),
|
|
||||||
fps_id + 1,
|
|
||||||
(samples.shape[-1], samples.shape[-2]),
|
|
||||||
)
|
|
||||||
|
|
||||||
vid = (
|
|
||||||
(rearrange(samples, "t c h w -> t h w c") * 255)
|
|
||||||
.cpu()
|
|
||||||
.numpy()
|
|
||||||
.astype(np.uint8)
|
|
||||||
)
|
|
||||||
for frame in vid:
|
|
||||||
frame = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
|
||||||
writer.write(frame)
|
|
||||||
writer.release()
|
|
||||||
|
|
||||||
|
|
||||||
def seed_everything(seed: int):
|
def seed_everything(seed: int):
|
||||||
import random, os
|
import random, os
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|||||||
|
Before Width: | Height: | Size: 103 KiB |
|
Before Width: | Height: | Size: 103 KiB |
|
Before Width: | Height: | Size: 100 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 101 KiB |
|
Before Width: | Height: | Size: 101 KiB |
|
Before Width: | Height: | Size: 102 KiB |
|
Before Width: | Height: | Size: 101 KiB |
|
Before Width: | Height: | Size: 103 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 100 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 100 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 102 KiB |
|
Before Width: | Height: | Size: 102 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 108 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 108 KiB |
|
Before Width: | Height: | Size: 109 KiB |
|
Before Width: | Height: | Size: 111 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 108 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 105 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 106 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |
|
Before Width: | Height: | Size: 107 KiB |