commit 8a8861305194a958061ea3f236c19dcffbc4f20e Author: db Date: Fri Nov 10 21:15:46 2023 +0800 init diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..8e2efbc --- /dev/null +++ b/.gitignore @@ -0,0 +1,198 @@ +### Example user template template +### Example user template + +# IntelliJ project files +.idea +*.iml +out +gen +### macOS template +# General +.DS_Store +.AppleDouble +.LSOverride + +# Icon must end with two \r +Icon + +# Thumbnails +._* + +# Files that might appear in the root of a volume +.DocumentRevisions-V100 +.fseventsd +.Spotlight-V100 +.TemporaryItems +.Trashes +.VolumeIcon.icns +.com.apple.timemachine.donotpresent + +# Directories potentially created on remote AFP share +.AppleDB +.AppleDesktop +Network Trash Folder +Temporary Items +.apdisk + +### Python template +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/#use-with-ide +.pdm.toml + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ + diff --git a/README.md b/README.md new file mode 100644 index 0000000..02ae77a --- /dev/null +++ b/README.md @@ -0,0 +1,3 @@ +# ComfyUI-Keyframe + +为关键帧设置重绘强度 \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..6c1633b --- /dev/null +++ b/__init__.py @@ -0,0 +1,7 @@ +from .launch_util import is_installed, run_pip +from .keyframe.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +if not is_installed("rich"): + run_pip("install rich", desc="Install rich", live=True) + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/keyframe/interface.py b/keyframe/interface.py new file mode 100644 index 0000000..7631498 --- /dev/null +++ b/keyframe/interface.py @@ -0,0 +1,57 @@ +from dataclasses import dataclass, field +from typing import Union + +import torch + + +class KeyframePart: + def __init__(self, batch_index: int, image: torch.Tensor, denoise: float) -> None: + self.batch_index = batch_index + self.denoise = denoise + self.image = image + + +class KeyframePartGroup: + def __init__(self) -> None: + self.keyframes: list[KeyframePart] = [] + + def add(self, keyframe: KeyframePart) -> None: + added = False + for i in range(len(self.keyframes)): + if self.keyframes[i].batch_index == keyframe.batch_index: + self.keyframes[i] = keyframe + added = True + break + if not added: + self.keyframes.append(keyframe) + self.keyframes.sort(key=lambda k: k.batch_index) + + def get_index(self, index: int) -> Union[KeyframePart, None]: + try: + return self.keyframes[index] + except IndexError: + return None + + def __getitem__(self, index) -> KeyframePart: + return self.keyframes[index] + + def is_empty(self) -> bool: + return len(self.keyframes) == 0 + + +@dataclass +class ModelInjectParam: + keyframe_part_group: KeyframePartGroup + latent: dict = field(default_factory=dict, repr=False) + seed: int = 0 + steps: int = 0 + scheduler: str = 'normal' + denoise: float = 0 + noise: torch.Tensor = field(default=None, repr=False) + + def reset(self): + self.seed: int = 0 + self.steps: int = 0 + self.scheduler: str = 'normal' + self.denoise: float = 0 + self.noise: torch.Tensor = None diff --git a/keyframe/nodes.py b/keyframe/nodes.py new file mode 100644 index 0000000..7eb14a3 --- /dev/null +++ b/keyframe/nodes.py @@ -0,0 +1,166 @@ +import numpy as np + +from .interface import KeyframePartGroup, KeyframePart, ModelInjectParam +from .sampling import keyframe_sample_factory + +from .util import inject_model, print +from .samples import inject_samples +import comfy.sample as comfy_sample + +inject_samples() +comfy_sample.sample = keyframe_sample_factory(comfy_sample.sample) + + +class KeyframePartNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "batch_index": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1}), + "denoise": ("FLOAT", {"default": 0.1, "min": 0.001, "max": 1.0, "step": 0.01}), + }, + "optional": { + "part": ("LATENT_KEYFRAME_PART",), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME_PART",) + RETURN_NAMES = ("part",) + FUNCTION = "load_keyframe_part" + + CATEGORY = "KeyframePart" + + def load_keyframe_part(self, image, batch_index, denoise, part=None): + if not part: + part = KeyframePartGroup() + keyframe = KeyframePart(batch_index, image, denoise) + part.add(keyframe) + return (part,) + + +class KeyframeInterpolationPartNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "batch_index_from": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1}), + "batch_index_to": ("INT", {"default": 4, "min": 1, "max": 9999, "step": 1}), + "denoise_from": ("FLOAT", {"default": 0.1, "min": 0.001, "max": 1.0, "step": 0.01}), + "denoise_to": ("FLOAT", {"default": 1.0, "min": 0.001, "max": 1.0, "step": 0.01}), + "interpolation": (["linear", "ease-in", "ease-out", "ease-in-out"],), + }, + "optional": { + "part": ("LATENT_KEYFRAME_PART",), + } + } + + RETURN_TYPES = ("LATENT_KEYFRAME_PART",) + RETURN_NAMES = ("part",) + FUNCTION = "load_keyframe_part" + + CATEGORY = "KeyframeInterpolationPartNode" + + def load_keyframe_part(self, + image, + batch_index_from, + batch_index_to, + denoise_from, + denoise_to, + interpolation, + part=None): + if batch_index_from >= batch_index_to: + raise ValueError("batch_index_from must be less than batch_index_to") + + if not part: + part = KeyframePartGroup() + current_group = KeyframePartGroup() + + steps = batch_index_to - batch_index_from + diff = denoise_to - denoise_from + if interpolation == "linear": + weights = np.linspace(denoise_from, denoise_to, steps) + elif interpolation == "ease-in": + index = np.linspace(0, 1, steps) + weights = diff * np.power(index, 2) + denoise_from + elif interpolation == "ease-out": + index = np.linspace(0, 1, steps) + weights = diff * (1 - np.power(1 - index, 2)) + denoise_from + elif interpolation == "ease-in-out": + index = np.linspace(0, 1, steps) + weights = diff * ((1 - np.cos(index * np.pi)) / 2) + denoise_from + + for i in range(steps): + keyframe = KeyframePart(batch_index_from + i, image, float(weights[i])) + current_group.add(keyframe) + + # replace values with prev_latent_keyframes + for latent_keyframe in part.keyframes: + current_group.add(latent_keyframe) + + return (current_group,) + + +class KeyframeApplyNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "latent": ("LATENT",), + "part_group": ("LATENT_KEYFRAME_PART",), + "vae": ("VAE",), + } + } + + RETURN_TYPES = ("MODEL", "LATENT",) + RETURN_NAMES = ("model", "latent",) + FUNCTION = "apply_latent_keyframe" + + CATEGORY = "LatentKeyframeApply" + + @staticmethod + def vae_encode_crop_pixels(pixels): + x = (pixels.shape[1] // 8) * 8 + y = (pixels.shape[2] // 8) * 8 + if pixels.shape[1] != x or pixels.shape[2] != y: + x_offset = (pixels.shape[1] % 8) // 2 + y_offset = (pixels.shape[2] % 8) // 2 + pixels = pixels[:, x_offset:x + x_offset, y_offset:y + y_offset, :] + return pixels + + def encode(self, vae, pixels): + pixels = self.vae_encode_crop_pixels(pixels) + t = vae.encode(pixels[:, :, :, :3]) + return t + + def apply_latent_keyframe(self, model, latent, part_group: KeyframePartGroup, vae): + # 预处理latent,把图片替换进去 + for part in part_group.keyframes: + latent['samples'][part.batch_index] = self.encode(vae, part.image) + print(f"apply keyframe {part.batch_index}:{part.denoise}") + + # 注入参数,后续处理 + + inject_param = ModelInjectParam(part_group, latent) + inject_model(model.model, inject_param) + + return (model, latent,) + + +# NODE MAPPING +NODE_CLASS_MAPPINGS = { + # Keyframes + "KeyframePart": KeyframePartNode, + "KeyframeInterpolationPart": KeyframeInterpolationPartNode, + "KeyframeApply": KeyframeApplyNode, + +} + +NODE_DISPLAY_NAME_MAPPINGS = { + # Keyframes + "KeyframePart": "Keyframe Part", + "KeyframeInterpolationPart": "Keyframe Interpolation Part", + "KeyframeApply": "Keyframe Apply", +} diff --git a/keyframe/samples.py b/keyframe/samples.py new file mode 100644 index 0000000..993ee2c --- /dev/null +++ b/keyframe/samples.py @@ -0,0 +1,104 @@ +import torch +from tqdm.auto import trange + +import comfy.samplers +from comfy.k_diffusion import sampling as k_diffusion_sampling +from comfy.k_diffusion.sampling import to_d, default_noise_sampler +from .util import print, is_injected_model, get_injected_model, generate_sigmas, generate_noise, get_ancestral_step + +CUSTOM_SAMPLERS = [ + 'k_euler', 'k_euler_a' +] + + +def inject_samples(): + comfy.samplers.SAMPLER_NAMES.extend(CUSTOM_SAMPLERS) + k_diffusion_sampling.sample_k_euler = sample_k_euler + k_diffusion_sampling.sample_k_euler_a = sample_k_euler_a + print(f'Injected samplers: {CUSTOM_SAMPLERS}') + + +def get_sigmas_noise(model_wrap, x, noise, latent_image, sigmas, scheduler, steps, part_group): + sigmas = generate_sigmas(model_wrap.inner_model, x, sigmas, scheduler, steps, part_group, sigmas.device) + noise = noise.to(x.device) + latent_image = latent_image.to(x.device) + for i in range(noise.shape[0]): + noise[i] = generate_noise(model_wrap, sigmas[i], noise[i]) + + if latent_image is not None: + latent_image = model_wrap.inner_model.process_latent_in(latent_image) + noise += latent_image + return sigmas, noise + + +@torch.no_grad() +def sample_k_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., + s_tmax=float('inf'), s_noise=1.): + model_wrap = model.inner_model + real_model = model_wrap.inner_model + + if not is_injected_model(real_model): + raise Exception("model is not injected,please use LatentKeyframeApply node to inject model") + + inject_param = get_injected_model(real_model) + latent_image = inject_param.latent['samples'] + + sigmas, noise = get_sigmas_noise(model_wrap, x, inject_param.noise, latent_image, sigmas, inject_param.scheduler, + inject_param.steps, inject_param.keyframe_part_group) + + extra_args = {} if extra_args is None else extra_args + s_tmin = s_tmin * x.new_ones([sigmas.shape[0]]) + s_tmax = s_tmax * x.new_ones([sigmas.shape[0]]) + gammas = s_tmax * x.new_ones([sigmas.shape[0]]) + for i in trange(sigmas.shape[1] - 1, disable=disable): + for j, sigma in enumerate(sigmas): + gammas[j] = min(s_churn / (sigmas.shape[0] - 1), 2 ** 0.5 - 1) if s_tmin[j] <= sigma[i] <= s_tmax[j] else 0. + sigma_hat = sigmas[:, i] * (gammas + 1) + + for j, gamma in enumerate(gammas): + if gamma > 0: + eps = torch.randn_like(noise) * s_noise + noise[j] = noise[j] + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5 + + denoised = model(noise, sigma_hat, **extra_args) + d = to_d(noise, sigma_hat, denoised) + if callback is not None: + callback({'x': noise, 'i': i, 'sigma': sigma_hat, 'denoised': denoised}) + dt = sigmas[:, i + 1] - sigma_hat + noise += d * dt.view(d.shape[0], 1, 1, 1) + return noise + + +@torch.no_grad() +def sample_k_euler_a(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., + noise_sampler=None): + + model_wrap = model.inner_model + real_model = model_wrap.inner_model + + if not is_injected_model(real_model): + raise Exception("model is not injected,please use LatentKeyframeApply node to inject model") + + inject_param = get_injected_model(real_model) + latent_image = inject_param.latent['samples'] + + sigmas, noise = get_sigmas_noise(model_wrap, x, inject_param.noise, latent_image, sigmas, inject_param.scheduler, + inject_param.steps, inject_param.keyframe_part_group) + + extra_args = {} if extra_args is None else extra_args + noise_sampler = default_noise_sampler(noise) if noise_sampler is None else noise_sampler + + for i in trange(sigmas.shape[1] - 1, disable=disable): + s_in = sigmas[:, i] + denoised = model(noise, s_in, **extra_args) + sigma_down, sigma_up = get_ancestral_step(sigmas[:, i], sigmas[:, i + 1], eta=eta) + if callback is not None: + callback({'x': noise, 'i': i, 'sigma': s_in, 'denoised': denoised}) + d = to_d(noise, s_in, denoised) + # Euler method + dt = sigma_down - s_in + noise += d * dt.view(d.shape[0], 1, 1, 1) + for j, sigma in enumerate(sigmas): + if sigma[i + 1] > 0: + noise[j] = noise[j] + noise_sampler(sigma[i], sigma[i + 1])[j] * s_noise * sigma_up[j] + return noise diff --git a/keyframe/sampling.py b/keyframe/sampling.py new file mode 100644 index 0000000..88297e4 --- /dev/null +++ b/keyframe/sampling.py @@ -0,0 +1,24 @@ +from typing import Callable + +from comfy.model_patcher import ModelPatcher +from .util import is_injected_model, print, get_injected_model + + +def keyframe_sample_factory(orig_comfy_sample: Callable) -> Callable: + def keyframe_sample(model: ModelPatcher, *args, **kwargs): + if not is_injected_model(model.model): + return orig_comfy_sample(model, *args, **kwargs) + inject_param = get_injected_model(model.model) + try: + inject_param.reset() + + inject_param.noise = args[0] + inject_param.steps = args[1] + inject_param.scheduler = args[4] + inject_param.denoise = kwargs.get('denoise', None) + inject_param.seed = kwargs.get('seed', None) + return orig_comfy_sample(model, *args, **kwargs) + finally: + inject_param.reset() + + return keyframe_sample diff --git a/keyframe/util.py b/keyframe/util.py new file mode 100644 index 0000000..8a69d30 --- /dev/null +++ b/keyframe/util.py @@ -0,0 +1,72 @@ +import math + +import torch + +from comfy.samplers import calculate_sigmas_scheduler +from rich.console import Console + +KEYFRAME_INJECTED_ATTR = "keyframe_injected" +console = Console(color_system="truecolor", force_terminal=True) + + +def inject_model(model, inject_param): + # 注入模型参数 + setattr(model, KEYFRAME_INJECTED_ATTR, inject_param) + return model + + +def is_injected_model(model): + return hasattr(model, KEYFRAME_INJECTED_ATTR) + + +def get_injected_model(model): + return getattr(model, KEYFRAME_INJECTED_ATTR) + + +def clear_injected_model(model): + if is_injected_model(model): + delattr(model, KEYFRAME_INJECTED_ATTR) + + +def print(msg, *args, **kwargs): + msg = f'[bold red]Keyframe[/bold red] [green]{msg}[/green]' + console.print(msg) + + +def max_denoise(model_wrap, sigmas): + max_sigma = float(model_wrap.inner_model.model_sampling.sigma_max) + sigma = float(sigmas[0]) + return math.isclose(max_sigma, sigma, rel_tol=1e-05) or sigma > max_sigma + + +def generate_sigmas(real_model, x, origin_sigmas, scheduler, steps, part_group, device): + batch_size = x.shape[0] + new_sigmas = origin_sigmas.unsqueeze(0).repeat(batch_size, 1) + + for part in part_group: + if part.denoise is None or part.denoise > 0.9999: + new_sigmas[part.batch_index] = calculate_sigmas_scheduler(real_model, scheduler, steps).to(device) + else: + new_steps = int(steps / part.denoise) + sigmas = calculate_sigmas_scheduler(real_model, scheduler, new_steps).to(device) + new_sigmas[part.batch_index] = sigmas[-(steps + 1):] + return new_sigmas + + +def generate_noise(model_wrap, sigmas, noise): + if max_denoise(model_wrap, sigmas): + n = noise * torch.sqrt(1.0 + sigmas[0] ** 2.0) + else: + n = noise * sigmas[0] + return n + + +def get_ancestral_step(sigma_from: torch.Tensor, sigma_to: torch.Tensor, eta: float = 1.) -> ( + torch.Tensor, torch.Tensor): + if not eta: + return sigma_to, torch.zeros_like(sigma_to) + sigma_up = torch.min(sigma_to, + eta * (sigma_to ** 2 * (sigma_from ** 2 - sigma_to ** 2) / sigma_from ** 2) ** 0.5) + sigma_down = (sigma_to ** 2 - sigma_up ** 2) ** 0.5 + + return sigma_down, sigma_up diff --git a/launch_util.py b/launch_util.py new file mode 100644 index 0000000..797ad95 --- /dev/null +++ b/launch_util.py @@ -0,0 +1,52 @@ +import importlib.util +import subprocess +import sys +import os + +python = sys.executable + + +def is_installed(package): + try: + spec = importlib.util.find_spec(package) + except ModuleNotFoundError: + return False + + return spec is not None + + +def run_pip(command, desc=None, live=False): + return run(f'"{python}" -m pip {command}', desc=f"Installing {desc}", + errdesc=f"Couldn't install {desc}", live=live) + + +def run(command, desc=None, errdesc=None, custom_env=None, live: bool = False) -> str: + if desc is not None: + print(desc) + + run_kwargs = { + "args": command, + "shell": True, + "env": os.environ if custom_env is None else custom_env, + "encoding": 'utf8', + "errors": 'ignore', + } + + if not live: + run_kwargs["stdout"] = run_kwargs["stderr"] = subprocess.PIPE + + result = subprocess.run(**run_kwargs) + + if result.returncode != 0: + error_bits = [ + f"{errdesc or 'Error running command'}.", + f"Command: {command}", + f"Error code: {result.returncode}", + ] + if result.stdout: + error_bits.append(f"stdout: {result.stdout}") + if result.stderr: + error_bits.append(f"stderr: {result.stderr}") + raise RuntimeError("\n".join(error_bits)) + + return (result.stdout or "")