Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1393a46d67 | ||
|
|
27421363bf | ||
|
|
01400a541d | ||
|
|
c981354387 | ||
|
|
c92c8bfb1f | ||
|
|
9334025801 | ||
|
|
144893ecd2 | ||
|
|
bce8505ce2 | ||
|
|
0992657e54 | ||
|
|
2fa036992b | ||
|
|
9432a34c38 | ||
|
|
d4e8ee28fb | ||
|
|
34b8e83831 | ||
|
|
e29d480a1a |
@@ -26,6 +26,7 @@ jobs:
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install .[dev]
|
||||
pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
|
||||
- name: Run Linting
|
||||
run: |
|
||||
ruff check .
|
||||
|
||||
@@ -11,3 +11,5 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: comfy-org/node-diff@main
|
||||
with:
|
||||
base_ref: ${{ github.event.repository.default_branch }}
|
||||
|
||||
+80
-2
@@ -10,7 +10,85 @@ __author__ = """LanPaint"""
|
||||
__email__ = "czhengac@connect.ust.hk"
|
||||
__version__ = "0.0.1"
|
||||
|
||||
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
|
||||
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
def _install_lightweight_runtime_stubs() -> None:
|
||||
"""Install lightweight stubs so tooling can import this package without ComfyUI.
|
||||
|
||||
This is used by CI tooling (e.g., comfy-org/node-diff) that imports NODE_CLASS_MAPPINGS
|
||||
in an environment where ComfyUI isn't installed.
|
||||
"""
|
||||
import sys
|
||||
import types
|
||||
|
||||
# `src/LanPaint/nodes.py` uses `torch.Tensor` in type annotations.
|
||||
try:
|
||||
import torch # noqa: F401
|
||||
except ModuleNotFoundError:
|
||||
torch_mod = types.ModuleType("torch")
|
||||
|
||||
class Tensor: # noqa: N801 (match torch naming)
|
||||
pass
|
||||
|
||||
torch_mod.Tensor = Tensor
|
||||
torch_mod.nn = types.SimpleNamespace(functional=types.SimpleNamespace())
|
||||
sys.modules["torch"] = torch_mod
|
||||
|
||||
if "comfyui_version" not in sys.modules:
|
||||
comfyui_version_mod = types.ModuleType("comfyui_version")
|
||||
comfyui_version_mod.__version__ = "0.0.0"
|
||||
sys.modules["comfyui_version"] = comfyui_version_mod
|
||||
|
||||
sys.modules.setdefault("nodes", types.ModuleType("nodes"))
|
||||
sys.modules.setdefault("latent_preview", types.ModuleType("latent_preview"))
|
||||
|
||||
if "comfy" not in sys.modules:
|
||||
comfy_mod = types.ModuleType("comfy")
|
||||
comfy_mod.__path__ = []
|
||||
|
||||
comfy_utils_mod = types.ModuleType("comfy.utils")
|
||||
|
||||
def repeat_to_batch_size(tensor, batch_size): # type: ignore[no-untyped-def]
|
||||
if getattr(tensor, "shape", ())[0] == batch_size:
|
||||
return tensor
|
||||
return tensor
|
||||
|
||||
comfy_utils_mod.repeat_to_batch_size = repeat_to_batch_size
|
||||
|
||||
comfy_samplers_mod = types.ModuleType("comfy.samplers")
|
||||
|
||||
class DummyKSAMPLER: # noqa: N801 (match ComfyUI naming)
|
||||
pass
|
||||
|
||||
comfy_samplers_mod.KSAMPLER = DummyKSAMPLER
|
||||
|
||||
comfy_model_base_mod = types.ModuleType("comfy.model_base")
|
||||
|
||||
class ModelType: # noqa: N801 (match ComfyUI naming)
|
||||
FLUX = "FLUX"
|
||||
FLOW = "FLOW"
|
||||
|
||||
class WAN22: # noqa: N801 (match ComfyUI naming)
|
||||
pass
|
||||
|
||||
comfy_model_base_mod.ModelType = ModelType
|
||||
comfy_model_base_mod.WAN22 = WAN22
|
||||
|
||||
comfy_mod.utils = comfy_utils_mod
|
||||
comfy_mod.samplers = comfy_samplers_mod
|
||||
comfy_mod.model_base = comfy_model_base_mod
|
||||
|
||||
sys.modules["comfy"] = comfy_mod
|
||||
sys.modules["comfy.utils"] = comfy_utils_mod
|
||||
sys.modules["comfy.samplers"] = comfy_samplers_mod
|
||||
sys.modules["comfy.model_base"] = comfy_model_base_mod
|
||||
|
||||
|
||||
try:
|
||||
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
|
||||
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
|
||||
except ModuleNotFoundError:
|
||||
_install_lightweight_runtime_stubs()
|
||||
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
|
||||
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
+4
-1
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "LanPaint"
|
||||
version = "1.4.9"
|
||||
version = "1.4.10"
|
||||
description = "Achieve seamless inpainting results without needing a specialized inpainting model."
|
||||
authors = [
|
||||
{name = "LanPaint", email = "czhengac@connect.ust.hk"}
|
||||
@@ -75,5 +75,8 @@ select = [
|
||||
# See all rules here: https://docs.astral.sh/ruff/rules/#pyflakes-f
|
||||
]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"src/LanPaint/nodes.py" = ["F403", "F405"]
|
||||
|
||||
[tool.ruff.lint.flake8-quotes]
|
||||
inline-quotes = "double"
|
||||
|
||||
+64
-21
@@ -1,5 +1,5 @@
|
||||
import torch
|
||||
from .utils import *
|
||||
from .utils import StochasticHarmonicOscillator
|
||||
from functools import partial
|
||||
|
||||
class LanPaint():
|
||||
@@ -34,7 +34,6 @@ class LanPaint():
|
||||
def LanPaint(self, x, sigma, latent_mask, current_times, n_steps, model_options, seed, IS_FLUX, IS_FLOW):
|
||||
VE_Sigma, abt, Flow_t = current_times
|
||||
|
||||
|
||||
step_size = self.step_size * (1 - abt)
|
||||
step_size = self.add_none_dims(step_size)
|
||||
# self.inner_model.inner_model.scale_latent_inpaint returns variance exploding x_t values
|
||||
@@ -43,8 +42,6 @@ class LanPaint():
|
||||
return self.inner_model.inner_model.model_sampling.noise_scaling(sigma.reshape([sigma.shape[0]] + [1] * (len(noise.shape) - 1)), noise, latent_image)
|
||||
|
||||
x = x * (1 - latent_mask) + scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image)* latent_mask
|
||||
|
||||
|
||||
|
||||
if IS_FLUX or IS_FLOW:
|
||||
x_t = x * ( self.add_none_dims(abt)**0.5 + (1-self.add_none_dims(abt))**0.5 )
|
||||
@@ -68,7 +65,6 @@ class LanPaint():
|
||||
return out
|
||||
|
||||
def score_model(self, x_t, y, mask, abt, sigma, tflow, model_options, seed):
|
||||
|
||||
lamb = self.chara_lamb
|
||||
if self.IS_FLUX or self.IS_FLOW:
|
||||
# compute t for flow model, with a small epsilon compensating for numerical error.
|
||||
@@ -119,24 +115,71 @@ class LanPaint():
|
||||
x_t = x_t.to(dtype)
|
||||
v = v.to(dtype)
|
||||
return x_t, v
|
||||
if args is None:
|
||||
#v = torch.zeros_like(x_t)
|
||||
v = None
|
||||
C = Coef_C(x_t)
|
||||
#print(torch.squeeze(dtx), torch.squeeze(dty))
|
||||
x_t, v = advance_time(x_t, v, dt, Gamma, A, C, D)
|
||||
else:
|
||||
v, C = args
|
||||
|
||||
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
|
||||
def advance_time_overdamped(x_t, dt, A, C, D):
|
||||
"""
|
||||
Overdamped (Gamma -> infinity) limit:
|
||||
dx = -A x dt + C dt + D dW_t
|
||||
with C treated as constant over this substep.
|
||||
"""
|
||||
dtype = x_t.dtype
|
||||
with torch.autocast(device_type=x_t.device.type, dtype=torch.float32):
|
||||
A_dt = A * dt
|
||||
exp_neg = torch.exp(-A_dt)
|
||||
|
||||
C_new = Coef_C(x_t)
|
||||
v = v + Gamma**0.5 * ( C_new - C) *dt
|
||||
eps = 1e-8
|
||||
abs_A = torch.abs(A)
|
||||
# k = (1 - exp(-A dt)) / A -> dt when A -> 0
|
||||
k = torch.where(abs_A < eps, dt, (-torch.expm1(-A_dt)) / A)
|
||||
# k2 = (1 - exp(-2 A dt)) / (2 A) -> dt when A -> 0
|
||||
k2 = torch.where(abs_A < eps, dt, (-torch.expm1(-2 * A_dt)) / (2 * A))
|
||||
|
||||
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
|
||||
mean = exp_neg * x_t + k * C
|
||||
var = (D ** 2) * k2
|
||||
noise = torch.randn_like(x_t) * torch.sqrt(torch.clamp(var, min=0.0))
|
||||
x_t = mean + noise
|
||||
return x_t.to(dtype)
|
||||
|
||||
def run_damped(x_t, args):
|
||||
if args is None:
|
||||
v = None
|
||||
C = Coef_C(x_t)
|
||||
x_t, v = advance_time(x_t, v, dt, Gamma, A, C, D)
|
||||
else:
|
||||
v, C = args
|
||||
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
|
||||
C_new = Coef_C(x_t)
|
||||
v = v + Gamma**0.5 * ( C_new - C) *dt
|
||||
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
|
||||
C = C_new
|
||||
return x_t, (v, C)
|
||||
|
||||
def run_overdamped(x_t, args):
|
||||
if args is None:
|
||||
C = Coef_C(x_t)
|
||||
x_t = advance_time_overdamped(x_t, dt, A, C, D)
|
||||
else:
|
||||
_, C = args
|
||||
x_t = advance_time_overdamped(x_t, dt / 2, A, C, D)
|
||||
C_new = Coef_C(x_t)
|
||||
x_t = x_t + (C_new - C) * dt
|
||||
x_t = advance_time_overdamped(x_t, dt / 2, A, C, D)
|
||||
C = C_new
|
||||
return x_t, (None, C)
|
||||
|
||||
try:
|
||||
x_t_next, (v_next, C_next) = run_damped(x_t, args)
|
||||
|
||||
if torch.isnan(x_t_next).any() or torch.isnan(v_next).any():
|
||||
raise ValueError("NaN detected")
|
||||
|
||||
x_t = x_t_next
|
||||
v = v_next
|
||||
C = C_next
|
||||
|
||||
except Exception:
|
||||
x_t, (v, C) = run_overdamped(x_t, args)
|
||||
|
||||
C = C_new
|
||||
|
||||
return x_t, (v, C)
|
||||
|
||||
def prepare_step_size(self, current_times, step_size, sigma_x, sigma_y):
|
||||
@@ -148,7 +191,7 @@ class LanPaint():
|
||||
# Compute time step (dtx, dty) for x and y branches.
|
||||
dtx = 2 * step_size * sigma_x
|
||||
dty = 2 * step_size * sigma_y
|
||||
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
# Define friction parameter Gamma_hat for each branch.
|
||||
# Using dtx**0 provides a tensor of the proper device/dtype.
|
||||
@@ -178,4 +221,4 @@ class LanPaint():
|
||||
|
||||
def x0_evalutation(self, x_t, score, sigma, args):
|
||||
x0 = x_t + score(x_t)
|
||||
return x0
|
||||
return x0
|
||||
|
||||
+60
-216
@@ -1,112 +1,69 @@
|
||||
from contextlib import contextmanager
|
||||
from inspect import cleandoc
|
||||
import inspect
|
||||
import math
|
||||
# import nodes.py
|
||||
import comfy
|
||||
import nodes
|
||||
import latent_preview
|
||||
from functools import partial
|
||||
import torch
|
||||
from comfy.utils import repeat_to_batch_size
|
||||
from comfy.samplers import *
|
||||
from comfy.model_base import ModelType
|
||||
from .utils import *
|
||||
from .lanpaint import LanPaint
|
||||
from comfy.model_base import WAN22
|
||||
import comfyui_version
|
||||
import comfy.nested_tensor
|
||||
|
||||
def reshape_mask(input_mask, output_shape,video_inpainting=False):
|
||||
|
||||
import comfy.nested_tensor
|
||||
|
||||
# 修改这里的判断条件,不能只用 hasattr("unbind")
|
||||
if isinstance(input_mask, comfy.nested_tensor.NestedTensor):
|
||||
masks = input_mask.unbind()
|
||||
|
||||
# 如果 output_shape 也是嵌套的(通常 noise.shape 在 NestedTensor 下返回 tuple of shapes)
|
||||
if isinstance(output_shape, (list, tuple)) and len(output_shape) > 0 and not isinstance(output_shape[0], int):
|
||||
reshaped_parts = []
|
||||
for i in range(len(masks)):
|
||||
# 递归处理每一个子部分,并传入对应的子 shape
|
||||
reshaped_parts.append(reshape_mask(masks[i], output_shape[i], video_inpainting))
|
||||
return comfy.nested_tensor.NestedTensor(tuple(reshaped_parts))
|
||||
else:
|
||||
# 如果 output_shape 是单一形状(降级处理)
|
||||
return comfy.nested_tensor.NestedTensor(tuple(reshape_mask(m, output_shape, video_inpainting) for m in masks))
|
||||
|
||||
dims = len(output_shape) - 2
|
||||
print('output shape',output_shape)
|
||||
scale_mode = "nearest-exact"
|
||||
print('input mask',input_mask.shape,type(input_mask),torch.max(input_mask),torch.min(input_mask))
|
||||
print('target output_shape',output_shape)
|
||||
print('input_mask.ndim:', input_mask.ndim, 'output_shape len:', len(output_shape))
|
||||
|
||||
|
||||
# Handle input mask dimensions
|
||||
if input_mask.ndim == 2:
|
||||
input_mask = input_mask.unsqueeze(0).unsqueeze(0)
|
||||
elif input_mask.ndim == 3:
|
||||
input_mask = input_mask.unsqueeze(1)
|
||||
|
||||
# Handle 5D output shape (B, C, F, H, W) by ensuring input is 5D
|
||||
if len(output_shape) == 5 and input_mask.ndim == 4:
|
||||
if comfyui_version.__version__ >= "0.6.0":
|
||||
input_mask = input_mask.unsqueeze(2) # (B, C, 1, H, W)
|
||||
|
||||
# Handle video case with temporal dimension
|
||||
# if video_inpainting: # Video case: (batch, channels, frames, height, width)
|
||||
# target_frames = output_shape[2]
|
||||
# target_height, target_width = output_shape[-2:]
|
||||
|
||||
# print('Video case - input_mask initial shape:', input_mask.shape)
|
||||
|
||||
# # First reshape input_mask to have proper dimensions for video processing
|
||||
# # Assume input is (frames, channels, height, width) -> (1, channels, frames, height, width)
|
||||
# ## if comfy version < 0.6.0
|
||||
# if comfyui_version.__version__ < "0.6.0":
|
||||
# input_mask = input_mask.permute(1, 0, 2, 3).unsqueeze(0)
|
||||
# print('Video case - input_mask after reshaping:', input_mask.shape)
|
||||
# # Ensure we have the correct 5D shape: (batch, channels, frames, height, width)
|
||||
# batch_size, channels, frames, height, width = input_mask.shape
|
||||
# print('Video case - dimensions: batch_size={}, channels={}, frames={}, height={}, width={}'.format(batch_size, channels, frames, height, width))
|
||||
# print('Video case - target size:', (target_frames, target_height, target_width))
|
||||
|
||||
# # 3D nearest-exact interpolation: (batch, channels, frames, height, width) -> (batch, channels, target_frames, target_height, target_width)
|
||||
# temp_mask = torch.nn.functional.interpolate(
|
||||
# input_mask,
|
||||
# size=(target_frames, target_height, target_width),
|
||||
# mode=scale_mode,
|
||||
# )
|
||||
|
||||
# # temp_mask is already 5D: (batch, channels, target_frames, target_height, target_width)
|
||||
# mask = temp_mask
|
||||
# print('after mask',mask.shape)
|
||||
# # Handle channel dimension expansion if needed
|
||||
# if mask.shape[1] < output_shape[1]:
|
||||
# mask = mask.repeat(1, output_shape[1], 1, 1, 1)[:, :output_shape[1]]
|
||||
# # Handle batch dimension
|
||||
# mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
if video_inpainting:
|
||||
# 如果是 3D Token 序列 (LTXV 压平后的情况)
|
||||
if input_mask.ndim == 3 and len(output_shape) == 3:
|
||||
mask = torch.nn.functional.interpolate(
|
||||
input_mask,
|
||||
size=output_shape[2],
|
||||
mode=scale_mode
|
||||
)
|
||||
return mask
|
||||
|
||||
# 只有在确认为 5D 视频张量时才执行原有逻辑
|
||||
if input_mask.ndim == 5:
|
||||
target_frames = output_shape[2]
|
||||
target_height, target_width = output_shape[-2:]
|
||||
|
||||
# (这里保留你原有的 permute 和 unsqueeze 逻辑,但要确保它是针对非 5D 输入的补救)
|
||||
if input_mask.ndim < 5:
|
||||
# 假设输入是 (F, C, H, W) -> (1, C, F, H, W)
|
||||
if hasattr(comfyui_version, "__version__") and comfyui_version.__version__ < "0.6.0":
|
||||
input_mask = input_mask.permute(1, 0, 2, 3).unsqueeze(0)
|
||||
|
||||
# 现在可以安全地解包 5D 形状了
|
||||
batch_size, channels, frames, height, width = input_mask.shape
|
||||
mask = torch.nn.functional.interpolate(
|
||||
input_mask,
|
||||
size=(target_frames, target_height, target_width),
|
||||
mode=scale_mode,
|
||||
)
|
||||
|
||||
if mask.shape[1] < output_shape[1]:
|
||||
mask = mask.repeat(1, output_shape[1], 1, 1, 1)[:, :output_shape[1]]
|
||||
mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
return mask
|
||||
if video_inpainting: # Video case: (batch, channels, frames, height, width)
|
||||
target_frames = output_shape[2]
|
||||
target_height, target_width = output_shape[-2:]
|
||||
|
||||
print('Video case - input_mask initial shape:', input_mask.shape)
|
||||
|
||||
# First reshape input_mask to have proper dimensions for video processing
|
||||
# Assume input is (frames, channels, height, width) -> (1, channels, frames, height, width)
|
||||
## if comfy version < 0.6.0
|
||||
if comfyui_version.__version__ < "0.6.0":
|
||||
input_mask = input_mask.permute(1, 0, 2, 3).unsqueeze(0)
|
||||
print('Video case - input_mask after reshaping:', input_mask.shape)
|
||||
# Ensure we have the correct 5D shape: (batch, channels, frames, height, width)
|
||||
batch_size, channels, frames, height, width = input_mask.shape
|
||||
print('Video case - dimensions: batch_size={}, channels={}, frames={}, height={}, width={}'.format(batch_size, channels, frames, height, width))
|
||||
print('Video case - target size:', (target_frames, target_height, target_width))
|
||||
|
||||
# 3D nearest-exact interpolation: (batch, channels, frames, height, width) -> (batch, channels, target_frames, target_height, target_width)
|
||||
temp_mask = torch.nn.functional.interpolate(
|
||||
input_mask,
|
||||
size=(target_frames, target_height, target_width),
|
||||
mode=scale_mode,
|
||||
)
|
||||
|
||||
# temp_mask is already 5D: (batch, channels, target_frames, target_height, target_width)
|
||||
mask = temp_mask
|
||||
print('after mask',mask.shape)
|
||||
# Handle channel dimension expansion if needed
|
||||
if mask.shape[1] < output_shape[1]:
|
||||
mask = mask.repeat(1, output_shape[1], 1, 1, 1)[:, :output_shape[1]]
|
||||
# Handle batch dimension
|
||||
mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
else: # Original 2D image case
|
||||
if comfyui_version.__version__ < "0.6.0":
|
||||
mask = torch.nn.functional.interpolate(input_mask, size=output_shape[-2:], mode=scale_mode)
|
||||
@@ -115,7 +72,7 @@ def reshape_mask(input_mask, output_shape,video_inpainting=False):
|
||||
if mask.shape[1] < output_shape[1]:
|
||||
mask = mask.repeat((1, output_shape[1]) + (1,) * dims)[:,:output_shape[1]]
|
||||
mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
|
||||
|
||||
|
||||
return mask
|
||||
def prepare_mask(noise_mask, shape, device,video_inpainting=False):
|
||||
@@ -146,9 +103,9 @@ class CFGGuider_LanPaint:
|
||||
if isinstance(self.inner_model, WAN22):
|
||||
print("WAN22 detected")
|
||||
self.inner_model.extra_conds = super(WAN22, self.inner_model).extra_conds
|
||||
|
||||
if denoise_mask is not None:
|
||||
video_inpainting = self.model_options.get("video_inpainting", False)
|
||||
print('denoise_mask',denoise_mask.shape,type(denoise_mask))
|
||||
denoise_mask = prepare_mask(denoise_mask, noise.shape, device, video_inpainting)
|
||||
|
||||
noise = noise.to(device)
|
||||
@@ -196,6 +153,8 @@ class KSamplerX0Inpaint:
|
||||
abt = (1 - Flow_t)**2 / ((1 - Flow_t)**2 + Flow_t**2 )
|
||||
VE_Sigma = Flow_t / (1 - Flow_t)
|
||||
#print("t", torch.mean( sigma ).item(), "VE_Sigma", torch.mean( VE_Sigma ).item())
|
||||
|
||||
|
||||
else:
|
||||
VE_Sigma = sigma
|
||||
abt = 1/( 1+VE_Sigma**2 )
|
||||
@@ -205,31 +164,6 @@ class KSamplerX0Inpaint:
|
||||
if "denoise_mask_function" in model_options:
|
||||
denoise_mask = model_options["denoise_mask_function"](sigma, denoise_mask, extra_options={"model": self.inner_model, "sigmas": self.sigmas})
|
||||
|
||||
if isinstance(denoise_mask, comfy.nested_tensor.NestedTensor):
|
||||
masks = denoise_mask.unbind()
|
||||
xs = x.unbind()
|
||||
latent_imgs = self.latent_image.unbind()
|
||||
noises = self.noise.unbind()
|
||||
|
||||
outs = []
|
||||
# 针对 LTXV,通常 i=0 是视频,i=1 是音频
|
||||
for i in range(len(xs)):
|
||||
m = (masks[i] > 0.5).float()
|
||||
lm = 1 - m
|
||||
# 这里的 PaintMethod 通常只支持普通 Tensor,所以我们分块处理
|
||||
# 注意:如果音频部分不需要 Inpaint,可以增加判断
|
||||
current_times = (VE_Sigma, abt, Flow_t)
|
||||
|
||||
# 只有视频部分 (i=0) 应用 LanPaint 逻辑,音频部分通常直接 pass 或原样返回
|
||||
if i == 0:
|
||||
out_part = self.PaintMethod(xs[i], latent_imgs[i], noises[i], sigma, lm, current_times, model_options, seed)
|
||||
else:
|
||||
# 音频部分如果没有对应的 Inpaint 逻辑,通常直接调用 inner_model
|
||||
out_part, _ = self.inner_model(xs[i], sigma, model_options=model_options, seed=seed)
|
||||
outs.append(out_part)
|
||||
|
||||
return comfy.nested_tensor.NestedTensor(tuple(outs))
|
||||
|
||||
denoise_mask = (denoise_mask > 0.5).float()
|
||||
|
||||
latent_mask = 1 - denoise_mask
|
||||
@@ -244,7 +178,7 @@ class KSamplerX0Inpaint:
|
||||
out = self.PaintMethod(x, self.latent_image, self.noise, sigma, latent_mask, current_times, model_options, seed)
|
||||
else:
|
||||
out, _ = self.inner_model(x, sigma, model_options=model_options, seed=seed)
|
||||
|
||||
|
||||
# Add TAESD preview support - directly use the latent_preview module
|
||||
current_step = model_options.get("i", kwargs.get("i", 0))
|
||||
total_steps = model_options.get("total_steps", 0)
|
||||
@@ -255,7 +189,7 @@ class KSamplerX0Inpaint:
|
||||
callback = model_options.get("callback", None)
|
||||
if callback is not None:
|
||||
callback({"i": current_step, "denoised": out, "x": x})
|
||||
|
||||
|
||||
return out
|
||||
|
||||
# Custom sampler class extending ComfyUI's KSAMPLER for LanPaint
|
||||
@@ -264,7 +198,6 @@ class KSAMPLER(comfy.samplers.KSAMPLER):
|
||||
#noise here is a randn noise from comfy.sample.prepare_noise
|
||||
#latent_image is the latent image as input of the KSampler node. For inpainting, it is the masked latent image. Otherwise it is zero tensor.
|
||||
extra_args["denoise_mask"] = denoise_mask
|
||||
print("LanPaint KSampler start sampler_function",denoise_mask.shape if denoise_mask is not None else None)
|
||||
model_k = KSamplerX0Inpaint(model_wrap, sigmas)
|
||||
model_k.latent_image = latent_image
|
||||
if self.inpaint_options.get("random", False): #TODO: Should this be the default?
|
||||
@@ -395,13 +328,13 @@ class LanPaint_KSampler():
|
||||
model.LanPaint_cfg_BIG = cfg
|
||||
else:
|
||||
model.LanPaint_cfg_BIG = 0*cfg - 0.5
|
||||
|
||||
|
||||
# Convert inpainting_mode to boolean for video_inpainting
|
||||
video_inpainting = (Inpainting_mode == "🎬 Video Inpainting")
|
||||
if not hasattr(model, 'model_options') or model.model_options is None:
|
||||
model.model_options = {}
|
||||
model.model_options["video_inpainting"] = video_inpainting
|
||||
|
||||
|
||||
with override_sample_function():
|
||||
return nodes.common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=denoise)
|
||||
class LanPaint_KSamplerAdvanced:
|
||||
@@ -455,7 +388,7 @@ class LanPaint_KSamplerAdvanced:
|
||||
model.LanPaint_cfg_BIG = cfg
|
||||
else:
|
||||
model.LanPaint_cfg_BIG = 0*cfg - 0.5
|
||||
|
||||
|
||||
# Convert inpainting_mode to boolean for video_inpainting
|
||||
video_inpainting = (Inpainting_mode == "🎬 Video Inpainting")
|
||||
if not hasattr(model, 'model_options') or model.model_options is None:
|
||||
@@ -507,7 +440,7 @@ class MaskBlend:
|
||||
kernel = self.gaussian_kernel(blend_overlap)
|
||||
kernel = kernel.to(image1.device)
|
||||
kernel = kernel[None, None, ...]
|
||||
|
||||
|
||||
mask = torch.nn.functional.conv2d(mask[:,None,:,:], kernel, padding=blend_overlap//2)[:,0,:,:]
|
||||
|
||||
|
||||
@@ -529,77 +462,6 @@ class MaskBlend:
|
||||
|
||||
return kernel
|
||||
|
||||
class MaskBlendAlpha:
|
||||
"""
|
||||
Create an RGBA image by writing the mask into the PNG alpha channel.
|
||||
|
||||
Requirement:
|
||||
- inpaint region: alpha = 0 (transparent)
|
||||
- other region: alpha = 1 (opaque)
|
||||
|
||||
This node writes the mask into the PNG alpha channel.
|
||||
Current default behavior matches the previous `invert_mask=True` behavior:
|
||||
alpha = mask.
|
||||
"""
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"tooltip": "VAE-decoded image (RGB)."}),
|
||||
"mask": ("MASK", {"tooltip": "Mask used as alpha channel (alpha = mask)."}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "to_rgba"
|
||||
CATEGORY = "image/postprocessing"
|
||||
|
||||
def to_rgba(self, image: torch.Tensor, mask: torch.Tensor):
|
||||
"""
|
||||
image: [B,H,W,3] float in [0,1]
|
||||
mask: [B,H,W] (or [H,W]) float in [0,1] used as alpha
|
||||
returns RGBA image: [B,H,W,4] float in [0,1]
|
||||
"""
|
||||
if image.ndim != 4 or image.shape[-1] != 3:
|
||||
raise ValueError(f"Expected IMAGE tensor [B,H,W,3], got {tuple(image.shape)}")
|
||||
|
||||
# Normalize mask shape to [B,H,W]
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
elif mask.ndim == 3:
|
||||
pass
|
||||
else:
|
||||
# Some pipelines may carry mask as [B,1,H,W]
|
||||
if mask.ndim == 4 and mask.shape[1] == 1:
|
||||
mask = mask[:, 0, :, :]
|
||||
else:
|
||||
raise ValueError(f"Expected MASK tensor [B,H,W] or [H,W], got {tuple(mask.shape)}")
|
||||
|
||||
b, h, w, _ = image.shape
|
||||
|
||||
# Batch align
|
||||
if mask.shape[0] != b:
|
||||
if mask.shape[0] == 1:
|
||||
mask = mask.repeat(b, 1, 1)
|
||||
else:
|
||||
raise ValueError(f"Batch mismatch: image batch={b}, mask batch={mask.shape[0]}")
|
||||
|
||||
# Spatial align (resize mask to image resolution if needed)
|
||||
if mask.shape[1] != h or mask.shape[2] != w:
|
||||
mask_4d = mask.unsqueeze(1) # [B,1,H,W]
|
||||
mask_4d = torch.nn.functional.interpolate(mask_4d, size=(h, w), mode="nearest")
|
||||
mask = mask_4d[:, 0, :, :]
|
||||
|
||||
mask = mask.float().clamp(0.0, 1.0)
|
||||
|
||||
# Default behavior (matches previous invert_mask=True path):
|
||||
# alpha = mask
|
||||
rgba = torch.cat([image, mask.unsqueeze(-1)], dim=-1)
|
||||
return (rgba,)
|
||||
|
||||
class Noise_EmptyNoise:
|
||||
def generate_noise(self, latent):
|
||||
return torch.zeros_like(latent["samples"])
|
||||
@@ -628,7 +490,6 @@ class LanPaint_SamplerCustom:
|
||||
"LanPaint_NumSteps": ("INT", {"default": 5, "min": 0, "max": 100, "tooltip": "Number of steps for Langevin dynamics, representing turns of thinking per step."}),
|
||||
"LanPaint_PromptMode": (["Image First", "Prompt First"], {"tooltip": "Image First: prioritizes image quality; Prompt First: prioritizes prompt adherence."}),
|
||||
"LanPaint_Info": ("STRING", {"default": "LanPaint Custom Sampler. For more info, visit https://github.com/scraed/LanPaint. If you find it useful, please give a star ⭐️!", "multiline": True}),
|
||||
"Inpainting_mode": (["🖼️ Image Inpainting", "🎬 Video Inpainting"], {"default": "🖼️ Image Inpainting", "tooltip": "Choose Image mode for photos or Video mode for video frames with temporal consistency"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -637,7 +498,7 @@ class LanPaint_SamplerCustom:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, model, sampler, sigmas, add_noise, noise_seed, cfg, positive, negative, latent_image, LanPaint_NumSteps, LanPaint_PromptMode, LanPaint_Info="",Inpainting_mode="🖼️ Image Inpainting"):
|
||||
def sample(self, model, sampler, sigmas, add_noise, noise_seed, cfg, positive, negative, latent_image, LanPaint_NumSteps, LanPaint_PromptMode, LanPaint_Info=""):
|
||||
model.LanPaint_StepSize = 0.2
|
||||
model.LanPaint_Lambda = 16.0
|
||||
model.LanPaint_Beta = 1.
|
||||
@@ -648,10 +509,6 @@ class LanPaint_SamplerCustom:
|
||||
model.LanPaint_cfg_BIG = cfg
|
||||
else:
|
||||
model.LanPaint_cfg_BIG = 0 * cfg - 0.5
|
||||
video_inpainting = (Inpainting_mode == "🎬 Video Inpainting")
|
||||
if not hasattr(model, 'model_options') or model.model_options is None:
|
||||
model.model_options = {}
|
||||
model.model_options["video_inpainting"] = video_inpainting
|
||||
with override_sample_function():
|
||||
latent = latent_image.copy()
|
||||
latent_image = latent["samples"]
|
||||
@@ -681,7 +538,7 @@ class LanPaint_SamplerCustom:
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class LanPaint_SamplerCustomAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -699,7 +556,6 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
"LanPaint_PromptMode": (["Image First", "Prompt First"], {"tooltip": "Image First: prioritizes image quality; Prompt First: prioritizes prompt adherence."}),
|
||||
"LanPaint_EarlyStop": ("INT", {"default": 1, "min": 0, "max": 10000, "tooltip": "Steps to stop LanPaint early, preventing irregular patterns."}),
|
||||
"LanPaint_Info": ("STRING", {"default": "LanPaint Custom Sampler Adv. For more info, visit https://github.com/scraed/LanPaint. If you find it useful, please give a star ⭐️!", "multiline": True}),
|
||||
"Inpainting_mode": (["🖼️ Image Inpainting", "🎬 Video Inpainting"], {"default": "🖼️ Image Inpainting", "tooltip": "Choose Image mode for photos or Video mode for video frames with temporal consistency"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -710,7 +566,7 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, noise, guider, sampler, sigmas, latent_image, LanPaint_NumSteps, LanPaint_Lambda, LanPaint_StepSize, LanPaint_Beta, LanPaint_Friction, LanPaint_PromptMode, LanPaint_EarlyStop, LanPaint_Info="",Inpainting_mode="🖼️ Image Inpainting"):
|
||||
def sample(self, noise, guider, sampler, sigmas, latent_image, LanPaint_NumSteps, LanPaint_Lambda, LanPaint_StepSize, LanPaint_Beta, LanPaint_Friction, LanPaint_PromptMode, LanPaint_EarlyStop, LanPaint_Info=""):
|
||||
model = guider.model_patcher
|
||||
model.LanPaint_StepSize = LanPaint_StepSize
|
||||
model.LanPaint_Lambda = LanPaint_Lambda
|
||||
@@ -722,25 +578,16 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
model.LanPaint_cfg_BIG = guider.cfg
|
||||
else:
|
||||
model.LanPaint_cfg_BIG = 0 * guider.cfg - 0.5
|
||||
|
||||
video_inpainting = (Inpainting_mode == "🎬 Video Inpainting")
|
||||
if not hasattr(model, 'model_options') or model.model_options is None:
|
||||
model.model_options = {}
|
||||
model.model_options["video_inpainting"] = video_inpainting
|
||||
with override_sample_function():
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
print('before fix_empty_latent_channels latent_image shape',latent_image.shape)
|
||||
latent = latent.copy()
|
||||
latent_image = comfy.sample.fix_empty_latent_channels(guider.model_patcher, latent_image)
|
||||
latent["samples"] = latent_image
|
||||
print('latent_image shape',latent_image.shape)
|
||||
print('outside noise_mask',latent["noise_mask"].shape if "noise_mask" in latent else 'no noise_mask')
|
||||
print('latent keys',latent.keys())
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
print('inside noise_mask shape',noise_mask.shape)
|
||||
|
||||
x0_output = {}
|
||||
callback = latent_preview.prepare_callback(guider.model_patcher, sigmas.shape[-1] - 1, x0_output)
|
||||
@@ -756,7 +603,6 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
out_denoised["samples"] = guider.model_patcher.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
# print('output',out.keys(),out["samples"].shape,out['noise_mask'].shape)
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
@@ -768,7 +614,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"LanPaint_SamplerCustom" : LanPaint_SamplerCustom,
|
||||
"LanPaint_SamplerCustomAdvanced" : LanPaint_SamplerCustomAdvanced,
|
||||
"LanPaint_MaskBlend": MaskBlend,
|
||||
"LanPaint_MaskBlendAlpha": MaskBlendAlpha,
|
||||
# "LanPaint_UpSale_LatentNoiseMask": LanPaint_UpSale_LatentNoiseMask,
|
||||
}
|
||||
|
||||
@@ -779,6 +624,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LanPaint_SamplerCustom" : "LanPaint Sampler Custom",
|
||||
"LanPaint_SamplerCustomAdvanced" : "LanPaint Sampler Custom (Advanced)",
|
||||
"LanPaint_MaskBlend": "LanPaint Mask Blend",
|
||||
"LanPaint_MaskBlendAlpha": "MaskBlend (alpha)",
|
||||
# "LanPaint_UpSale_LatentNoiseMask": "LanPaint UpSale Latent Noise Mask"
|
||||
}
|
||||
|
||||
+20
-21
@@ -28,11 +28,11 @@ def expm1mxmhx2_x3(x):
|
||||
def exp_1mcosh_GD(gamma_t, delta):
|
||||
"""
|
||||
Compute e^(-Γt) * (1 - cosh(Γt√Δ))/ ( (Γt)**2 Δ )
|
||||
|
||||
|
||||
Parameters:
|
||||
gamma_t: Γ*t term (could be a scalar or tensor)
|
||||
delta: Δ term (could be a scalar or tensor)
|
||||
|
||||
|
||||
Returns:
|
||||
Result of the computation with numerical stability handling
|
||||
"""
|
||||
@@ -68,7 +68,6 @@ def exp_sinh_GsqrtD(gamma_t, delta):
|
||||
sqrt_abs_delta = torch.sqrt(torch.abs(delta))
|
||||
gamma_t_sqrt_delta = gamma_t * sqrt_abs_delta
|
||||
numerator_pos = (torch.exp(gamma_t * (sqrt_abs_delta - 1)) - torch.exp(gamma_t * (-sqrt_abs_delta - 1))) / 2
|
||||
denominator_pos = gamma_t_sqrt_delta
|
||||
result_pos = numerator_pos / gamma_t_sqrt_delta
|
||||
result_pos = torch.where(torch.isfinite(result_pos), result_pos, torch.zeros_like(result_pos))
|
||||
|
||||
@@ -117,15 +116,15 @@ def zeta1(gamma_t, delta):
|
||||
exp_cosh_term = exp_cosh(half_gamma_t, delta)
|
||||
exp_sinh_term = exp_sinh_sqrtD(half_gamma_t, delta)
|
||||
|
||||
|
||||
|
||||
# Main computation
|
||||
numerator = 1 - (exp_cosh_term + exp_sinh_term)
|
||||
denominator = gamma_t * (1 - delta) / 4
|
||||
result = 1 - numerator / denominator
|
||||
|
||||
|
||||
# Handle numerical instability
|
||||
result = torch.where(torch.isfinite(result), result, torch.zeros_like(result))
|
||||
|
||||
|
||||
# Taylor expansion for small x (similar to your epxm1Dx approach)
|
||||
mask = torch.abs(denominator) < 5e-3
|
||||
term1 = epxm1_x(-gamma_t)
|
||||
@@ -133,17 +132,17 @@ def zeta1(gamma_t, delta):
|
||||
term3 = expm1mxmhx2_x3(-gamma_t)
|
||||
taylor = term1 + (1/2.+ term1-3*term2)*denominator + (-1/6. + term1/2 - 4 * term2 + 10 * term3) * denominator**2
|
||||
result = torch.where(mask, taylor, result)
|
||||
|
||||
|
||||
return result
|
||||
|
||||
def exp_cosh_minus_terms(gamma_t, delta):
|
||||
"""
|
||||
Compute E^(-tΓ) * (Cosh[tΓ] - 1 - (Cosh[tΓ√Δ] - 1)/Δ) / (tΓ(1 - Δ))
|
||||
|
||||
|
||||
Parameters:
|
||||
gamma_t: Γ*t term (could be a scalar or tensor)
|
||||
delta: Δ term (could be a scalar or tensor)
|
||||
|
||||
|
||||
Returns:
|
||||
Result of the computation with numerical stability handling
|
||||
"""
|
||||
@@ -151,17 +150,17 @@ def exp_cosh_minus_terms(gamma_t, delta):
|
||||
# Compute individual terms
|
||||
exp_cosh_term = exp_cosh(gamma_t, gamma_t**0) - exp_term # E^(-tΓ) (Cosh[tΓ] - 1) term
|
||||
exp_cosh_delta_term = - gamma_t**2 * exp_1mcosh_GD(gamma_t, delta) # E^(-tΓ) (Cosh[tΓ√Δ] - 1)/Δ term
|
||||
|
||||
|
||||
#exp_1mcosh_GD e^(-Γt) * (1 - cosh(Γt√Δ))/ ( (Γt)**2 Δ )
|
||||
# Main computation
|
||||
numerator = exp_cosh_term - exp_cosh_delta_term
|
||||
denominator = gamma_t * (1 - delta)
|
||||
|
||||
|
||||
result = numerator / denominator
|
||||
|
||||
|
||||
# Handle numerical instability
|
||||
result = torch.where(torch.isfinite(result), result, torch.zeros_like(result))
|
||||
|
||||
|
||||
# Taylor expansion for small gamma_t and delta near 1
|
||||
mask = (torch.abs(denominator) < 1e-1)
|
||||
exp_1mcosh_GD_term = exp_1mcosh_GD(gamma_t, delta**0)
|
||||
@@ -170,7 +169,7 @@ def exp_cosh_minus_terms(gamma_t, delta):
|
||||
- denominator / 4 * ( 0.5 * exp_cosh(gamma_t, delta**0) - 4 * exp_1mcosh_GD_term - 5 /2 * exp_sinh_GsqrtD(gamma_t, delta**0) )
|
||||
)
|
||||
result = torch.where(mask, taylor, result)
|
||||
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@@ -185,7 +184,7 @@ def sig11(gamma_t, delta):
|
||||
def Zcoefs(gamma_t, delta):
|
||||
Zeta1 = zeta1(gamma_t, delta)
|
||||
Zeta2 = zeta2(gamma_t, delta)
|
||||
|
||||
|
||||
sq_total = 1 - Zeta1 + gamma_t * (delta - 1) * (Zeta1 - 1)**2 / 8
|
||||
amplitude = torch.sqrt(sq_total)
|
||||
Zcoef1 = ( gamma_t**0.5 * Zeta2 / 2 **0.5 ) / amplitude
|
||||
@@ -208,7 +207,7 @@ class StochasticHarmonicOscillator:
|
||||
dq(t) = -Γ A y(t) dt + Γ C dt + Γ D dw(t) - Γ q(t) dt
|
||||
|
||||
Also define v(t) = q(t) / √Γ, which is numerically more stable.
|
||||
|
||||
|
||||
Where:
|
||||
y(t) - Position variable
|
||||
q(t) - Velocity variable
|
||||
@@ -239,7 +238,7 @@ class StochasticHarmonicOscillator:
|
||||
Returns:
|
||||
tuple: (y(t), v(t))
|
||||
"""
|
||||
|
||||
|
||||
dummyzero = y0.new_zeros(1) # convert scalar to tensor with same device and dtype as y0
|
||||
Delta = self.Delta + dummyzero
|
||||
Gamma_hat = self.Gamma * t + dummyzero
|
||||
@@ -254,12 +253,12 @@ class StochasticHarmonicOscillator:
|
||||
if v0 is None:
|
||||
v0 = torch.randn_like(y0) * D / 2 ** 0.5
|
||||
#v0 = (C - A * y0)/Gamma**0.5
|
||||
|
||||
|
||||
# Calculate mean position and velocity
|
||||
term1 = (1 - zeta_1) * (C * t - A * t * y0) + zeta_2 * (Gamma ** 0.5) * v0 * t
|
||||
y_mean = term1 + y0
|
||||
v_mean = (1 - EE)*(C - A * y0) / (Gamma ** 0.5) + (EE - A * t * (1 - zeta_1)) * v0
|
||||
|
||||
|
||||
cov_yy = D**2 * t * self.sig22(Gamma_hat, Delta)
|
||||
cov_vv = D**2 * self.sig11(Gamma_hat, Delta) / 2
|
||||
cov_yv = (zeta2(Gamma_hat, Delta) * Gamma_hat * D ) **2 / 2 / (Gamma ** 0.5)
|
||||
@@ -274,7 +273,7 @@ class StochasticHarmonicOscillator:
|
||||
cov_matrix[..., 1, 1] = cov_vv
|
||||
|
||||
|
||||
|
||||
|
||||
# Compute the Cholesky decomposition to get scale_tril
|
||||
#scale_tril = torch.linalg.cholesky(cov_matrix)
|
||||
scale_tril = torch.zeros(*batch_shape, 2, 2, device=y0.device, dtype=y0.dtype)
|
||||
@@ -298,4 +297,4 @@ class StochasticHarmonicOscillator:
|
||||
scale_tril=scale_tril
|
||||
).sample()
|
||||
|
||||
return new_yv[...,0], new_yv[...,1]
|
||||
return new_yv[...,0], new_yv[...,1]
|
||||
|
||||
+4
-3
@@ -1,4 +1,5 @@
|
||||
[pytest]
|
||||
testpaths = . # Run tests in the current directory
|
||||
python_files = test_*.py # Run tests in files that start with "test_"
|
||||
norecursedirs = .. # Don't run tests in the parent directory
|
||||
# Keep settings value-only; pytest does not treat inline `# ...` as comments.
|
||||
testpaths = .
|
||||
python_files = test_*.py
|
||||
norecursedirs = ..
|
||||
|
||||
+9
-17
@@ -1,21 +1,13 @@
|
||||
#!/usr/bin/env python
|
||||
"""Basic import tests for LanPaint.
|
||||
|
||||
"""Tests for `LanPaint` package."""
|
||||
The ComfyUI runtime dependencies (e.g. `comfy`) are intentionally optional for unit tests.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from src.LanPaint.nodes import Example
|
||||
|
||||
@pytest.fixture
|
||||
def example_node():
|
||||
"""Fixture to create an Example node instance."""
|
||||
return Example()
|
||||
def test_package_imports_without_comfy() -> None:
|
||||
import LanPaint
|
||||
|
||||
def test_example_node_initialization(example_node):
|
||||
"""Test that the node can be instantiated."""
|
||||
assert isinstance(example_node, Example)
|
||||
|
||||
def test_return_types():
|
||||
"""Test the node's metadata."""
|
||||
assert Example.RETURN_TYPES == ("IMAGE",)
|
||||
assert Example.FUNCTION == "test"
|
||||
assert Example.CATEGORY == "Example"
|
||||
assert isinstance(LanPaint.NODE_CLASS_MAPPINGS, dict)
|
||||
assert isinstance(LanPaint.NODE_DISPLAY_NAME_MAPPINGS, dict)
|
||||
assert "LanPaint_KSampler" in LanPaint.NODE_CLASS_MAPPINGS
|
||||
assert LanPaint.WEB_DIRECTORY == "./web"
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
import importlib
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
def _repeat_to_batch_size(tensor: torch.Tensor, batch_size: int) -> torch.Tensor:
|
||||
if tensor.shape[0] == batch_size:
|
||||
return tensor
|
||||
if tensor.shape[0] == 1:
|
||||
return tensor.repeat((batch_size,) + (1,) * (tensor.ndim - 1))
|
||||
repeats = (batch_size + tensor.shape[0] - 1) // tensor.shape[0]
|
||||
return tensor.repeat((repeats,) + (1,) * (tensor.ndim - 1))[:batch_size]
|
||||
|
||||
|
||||
def _import_nodes(monkeypatch, comfyui_version: str):
|
||||
comfy_mod = types.ModuleType("comfy")
|
||||
comfy_mod.__path__ = []
|
||||
|
||||
comfy_utils_mod = types.ModuleType("comfy.utils")
|
||||
comfy_utils_mod.repeat_to_batch_size = _repeat_to_batch_size
|
||||
|
||||
comfy_samplers_mod = types.ModuleType("comfy.samplers")
|
||||
class DummyKSAMPLER: ...
|
||||
comfy_samplers_mod.KSAMPLER = DummyKSAMPLER
|
||||
|
||||
comfy_model_base_mod = types.ModuleType("comfy.model_base")
|
||||
class ModelType:
|
||||
FLUX = "FLUX"
|
||||
FLOW = "FLOW"
|
||||
|
||||
class WAN22: ...
|
||||
comfy_model_base_mod.ModelType = ModelType
|
||||
comfy_model_base_mod.WAN22 = WAN22
|
||||
|
||||
comfyui_version_mod = types.ModuleType("comfyui_version")
|
||||
comfyui_version_mod.__version__ = comfyui_version
|
||||
|
||||
comfy_mod.utils = comfy_utils_mod
|
||||
comfy_mod.samplers = comfy_samplers_mod
|
||||
comfy_mod.model_base = comfy_model_base_mod
|
||||
|
||||
monkeypatch.setitem(sys.modules, "comfy", comfy_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.utils", comfy_utils_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.samplers", comfy_samplers_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.model_base", comfy_model_base_mod)
|
||||
monkeypatch.setitem(sys.modules, "nodes", types.ModuleType("nodes"))
|
||||
monkeypatch.setitem(sys.modules, "latent_preview", types.ModuleType("latent_preview"))
|
||||
monkeypatch.setitem(sys.modules, "comfyui_version", comfyui_version_mod)
|
||||
|
||||
sys.modules.pop("src.LanPaint.nodes", None)
|
||||
return importlib.import_module("src.LanPaint.nodes")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("comfyui_version", ["0.5.0", "0.6.0"])
|
||||
def test_reshape_mask_accepts_bhw_and_5d_output_shape(monkeypatch, comfyui_version: str) -> None:
|
||||
lanpaint_nodes = _import_nodes(monkeypatch, comfyui_version)
|
||||
input_mask = torch.zeros((1, 4, 4))
|
||||
output_shape = (1, 16, 1, 8, 8)
|
||||
|
||||
out = lanpaint_nodes.reshape_mask(input_mask, output_shape, video_inpainting=False)
|
||||
assert tuple(out.shape) == output_shape
|
||||
|
||||
|
||||
def test_prepare_mask_accepts_hw_and_moves_device(monkeypatch) -> None:
|
||||
lanpaint_nodes = _import_nodes(monkeypatch, "0.5.0")
|
||||
input_mask = torch.zeros((4, 4))
|
||||
output_shape = (2, 3, 8, 8)
|
||||
|
||||
out = lanpaint_nodes.prepare_mask(input_mask, output_shape, device=torch.device("cpu"), video_inpainting=False)
|
||||
assert tuple(out.shape) == output_shape
|
||||
assert out.device.type == "cpu"
|
||||
@@ -0,0 +1,51 @@
|
||||
import torch
|
||||
from unittest.mock import MagicMock, patch
|
||||
from src.LanPaint.lanpaint import LanPaint
|
||||
def test_langevin_dynamics_fallback_on_nan() -> None:
|
||||
"""Test that langevin_dynamics falls back to overdamped dynamics if damped dynamics produces NaNs."""
|
||||
torch.manual_seed(0)
|
||||
# Setup minimal LanPaint instance
|
||||
lp = LanPaint(Model=MagicMock(), NSteps=10, Friction=1.0, Lambda=1.0, Beta=1.0, StepSize=0.1)
|
||||
# Dummy inputs
|
||||
# Shape: (Batch, Channel, Height, Width)
|
||||
x_t = torch.randn(1, 4, 8, 8)
|
||||
lp.img_dim_size = 4
|
||||
mask = torch.zeros_like(x_t)
|
||||
# Simple score function
|
||||
def score(x):
|
||||
return torch.zeros_like(x)
|
||||
step_size = torch.tensor([0.1])
|
||||
# (sigma, abt, flow_t)
|
||||
current_times = (torch.tensor([0.5]), torch.tensor([0.5]), torch.tensor([0.5]))
|
||||
# Mock StochasticHarmonicOscillator to return NaNs
|
||||
# We patch it where it is used (imported) in lanpaint.py
|
||||
with patch("src.LanPaint.lanpaint.StochasticHarmonicOscillator") as MockSHO:
|
||||
mock_instance = MockSHO.return_value
|
||||
# Configure dynamics to return NaNs
|
||||
nan_tensor = torch.full_like(x_t, float('nan'))
|
||||
mock_instance.dynamics.return_value = (nan_tensor, nan_tensor)
|
||||
# Execute langevin_dynamics
|
||||
# This should try run_damped -> get NaNs -> raise ValueError -> catch -> run_overdamped
|
||||
x_out, (v_out, C_out) = lp.langevin_dynamics(x_t, score, mask, step_size, current_times, sigma_y=1.0)
|
||||
# Verify that SHO was initialized and dynamics called
|
||||
MockSHO.assert_called()
|
||||
mock_instance.dynamics.assert_called()
|
||||
# Verify result is finite (indicating fallback to overdamped logic was successful)
|
||||
assert torch.isfinite(x_out).all(), "Output contains NaNs, fallback failed"
|
||||
assert v_out is None or torch.isfinite(v_out).all()
|
||||
def test_langevin_dynamics_fallback_on_exception() -> None:
|
||||
"""Test that langevin_dynamics falls back to overdamped dynamics if damped dynamics raises Exception."""
|
||||
torch.manual_seed(0)
|
||||
lp = LanPaint(Model=MagicMock(), NSteps=10, Friction=1.0, Lambda=1.0, Beta=1.0, StepSize=0.1)
|
||||
x_t = torch.randn(1, 4, 8, 8)
|
||||
lp.img_dim_size = 4
|
||||
mask = torch.zeros_like(x_t)
|
||||
score = lambda x: torch.zeros_like(x)
|
||||
step_size = torch.tensor([0.1])
|
||||
current_times = (torch.tensor([0.5]), torch.tensor([0.5]), torch.tensor([0.5]))
|
||||
with patch("src.LanPaint.lanpaint.StochasticHarmonicOscillator") as MockSHO:
|
||||
mock_instance = MockSHO.return_value
|
||||
# Configure dynamics to raise Exception
|
||||
mock_instance.dynamics.side_effect = RuntimeError("Simulation exploded")
|
||||
x_out, (v_out, C_out) = lp.langevin_dynamics(x_t, score, mask, step_size, current_times, sigma_y=1.0)
|
||||
assert torch.isfinite(x_out).all()
|
||||
Reference in New Issue
Block a user