Compare commits

...
21 Commits
Author SHA1 Message Date
scraed aedb908a3f Add version tuple check for ComfyUI version handling
Introduced a helper function to compare ComfyUI versions as tuples instead of strings, improving reliability of version-dependent logic in reshape_mask. Updated all version checks to use the new tuple-based comparison.
2026-01-22 16:22:41 +08:00
scraed f9718ea3e1 Update Flux 2 klein image example link
Corrected the anchor link for the 'Flux 2 klein' image example to reference '2 steps of thinking' instead of '5 steps of thinking'.
2026-01-18 00:55:57 +08:00
scraed 0bafe2117a Merge branch 'master' of https://github.com/scraed/LanPaint 2026-01-18 00:46:59 +08:00
scraed 2c1a23777d Update example steps in README
Changed the number of steps in the Flux 2 klein InPaint example from 5 to 2 to reflect the correct workflow.
2026-01-18 00:46:56 +08:00
scraed dde4c82463 Bump version from 1.4.10 to 1.4.11 2026-01-18 00:46:14 +08:00
scraed 72cb484c40 Add Flux 2 Klein inpainting example and workflow
Added a new image inpainting example for Flux 2 Klein, including workflow JSON, sample images (original, masked, inpainted), and updated README with documentation and links. This provides users with a reference workflow and visual results for the Flux 2 Klein model.
2026-01-18 00:39:22 +08:00
scraed 0d15d15cf3 Update Discord link in README.md 2026-01-13 23:34:02 +08:00
scraed 1393a46d67 Bump version from 1.4.9 to 1.4.10 2026-01-13 22:17:42 +08:00
scraed 27421363bf Merge pull request #71 from godnight10061/fix/issue69-nan-multivariatenormal
Avoid MultivariateNormal crash on non-finite dynamics
2026-01-13 22:16:34 +08:00
scraed 01400a541d Update test_sho_regression.py 2026-01-13 22:13:30 +08:00
scraed c981354387 Update test_sho_regression.py 2026-01-13 21:46:56 +08:00
scraed c92c8bfb1f Update test_sho_regression.py 2026-01-13 21:38:49 +08:00
scraed 9334025801 Update test_sho_regression.py 2026-01-13 21:31:15 +08:00
scraed 144893ecd2 fall back to overdamped update if nan appears 2026-01-13 18:41:27 +08:00
scraed bce8505ce2 Revert "Avoid MultivariateNormal crash on non-finite dynamics"
This reverts commit 9432a34c38.
2026-01-13 18:29:06 +08:00
godnight10061 0992657e54 Fix CI workflows for torch and default branch 2026-01-13 09:25:41 +08:00
godnight10061 2fa036992b Add unit regression for NaN oscillator mean 2026-01-13 08:56:46 +08:00
godnight10061 9432a34c38 Avoid MultivariateNormal crash on non-finite dynamics 2026-01-13 01:39:48 +08:00
scraed d4e8ee28fb Merge pull request #70 from godnight10061/fix/reshape-mask-3d
Thanks for the contribution! Tested and merged.
2026-01-11 18:22:14 +08:00
godnight10061 34b8e83831 Make CI tooling importable 2026-01-11 11:07:41 +08:00
godnight10061 e29d480a1a Fix reshape_mask for 3D masks 2026-01-10 14:46:14 +08:00
17 changed files with 3025 additions and 86 deletions
+1
View File
@@ -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 .
+2
View File
@@ -11,3 +11,5 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: comfy-org/node-diff@main
with:
base_ref: ${{ github.event.repository.default_branch }}
+18 -1
View File
@@ -7,7 +7,7 @@
[![Hugging Face](https://img.shields.io/badge/Hugging%20Face-yellow?logo=huggingface&logoColor=white)](https://huggingface.co/charrywhite/LanPaint)
[![Blog](https://img.shields.io/badge/📝-Blog-9cf)](https://scraed.github.io/scraedBlog/)
[![GitHub stars](https://img.shields.io/github/stars/scraed/LanPaint)](https://github.com/scraed/LanPaint/stargazers)
[![Discord](https://img.shields.io/badge/Discord-5865F2?style=for-the-badge&logo=discord&logoColor=white)](https://discord.gg/aCGZutBV)
[![Discord](https://img.shields.io/badge/Discord-5865F2?style=for-the-badge&logo=discord&logoColor=white)](https://discord.gg/yN5wYDE6W4)
</div>
@@ -67,6 +67,7 @@ Check our latest [Wan 2.2 Video Examples](#video-examples-beta), [Wan 2.2 Image
- [Resource Consumption](#resource-consumption)
- [Image Examples](#image-examples)
- [Flux.2.Dev](#example-flux2dev-inpaintlanpaint-k-sampler-5-steps-of-thinking)
- [Flux 2 klein](#example-flux-2-klein-inpaintlanpaint-k-sampler-2-steps-of-thinking)
- [Z-image](#example-z-image-inpaintlanpaint-k-sampler-5-steps-of-thinking)
- [Hunyuan T2I](#example-hunyuan-t2i-inpaintlanpaint-k-sampler-5-steps-of-thinking)
- [Wan 2.2 T2I](#example-wan22-inpaintlanpaint-k-sampler-5-steps-of-thinking)
@@ -342,6 +343,22 @@ You need to follow the ComfyUI version of [SD 3.5 workflow](https://comfyui-wiki
(Note: Prompt First mode is disabled on Flux.2.Dev. As it does not use CFG guidance.)
### Example Flux 2 klein: InPaint(LanPaint K Sampler, 2 steps of thinking)
<details open>
<summary>View Original / Masked / Inpainted Comparison</summary>
| Original | Masked | Inpainted |
|:--------:|:------:|:---------:|
| ![Original Flux 2 klein](https://github.com/scraed/LanPaint/blob/master/examples/Example_24/Original_No_Mask.png) | ![Masked Flux 2 klein](https://github.com/scraed/LanPaint/blob/master/examples/Example_24/Masked_Load_Me_in_Loader.png) | ![Inpainted Flux 2 klein](https://github.com/scraed/LanPaint/blob/master/examples/Example_24/InPainted_Drag_Me_to_ComfyUI.png) |
</details>
[View Workflow & Masks](https://github.com/scraed/LanPaint/tree/master/examples/Example_24)
[Model Used in This Example](https://docs.comfy.org/zh-CN/tutorials/flux/flux-2-klein)
### Example Flux: InPaint(LanPaint K Sampler, 5 steps of thinking)
![Inpainting Result 7](https://github.com/scraed/LanPaint/blob/master/examples/InpaintChara_10.jpg)
[View Workflow & Masks](https://github.com/scraed/LanPaint/tree/master/examples/Example_7)
+80 -2
View File
@@ -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"
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.6 MiB

File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.8 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

+4 -1
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "LanPaint"
version = "1.4.9"
version = "1.4.11"
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
View File
@@ -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
+34 -20
View File
@@ -1,19 +1,22 @@
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
def _version_tuple(value):
return tuple(int(part) if part.isdigit() else 0 for part in value.split("."))
COMFYUI_VERSION_060_OR_NEWER = _version_tuple(comfyui_version.__version__) >= (0, 6, 0)
def reshape_mask(input_mask, output_shape,video_inpainting=False):
dims = len(output_shape) - 2
print('output shape',output_shape)
@@ -21,32 +24,43 @@ def reshape_mask(input_mask, output_shape,video_inpainting=False):
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_060_OR_NEWER:
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":
if not COMFYUI_VERSION_060_OR_NEWER:
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)
@@ -56,14 +70,14 @@ def reshape_mask(input_mask, output_shape,video_inpainting=False):
# Handle batch dimension
mask = repeat_to_batch_size(mask, output_shape[0])
else: # Original 2D image case
if comfyui_version.__version__ < "0.6.0":
if not COMFYUI_VERSION_060_OR_NEWER:
mask = torch.nn.functional.interpolate(input_mask, size=output_shape[-2:], mode=scale_mode)
else:
mask = torch.nn.functional.interpolate(input_mask, size=output_shape[2:], mode=scale_mode)
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):
@@ -144,7 +158,7 @@ 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
@@ -169,7 +183,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)
@@ -180,7 +194,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
@@ -319,13 +333,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:
@@ -379,7 +393,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:
@@ -431,7 +445,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,7 +543,7 @@ class LanPaint_SamplerCustom:
else:
out_denoised = out
return (out, out_denoised)
class LanPaint_SamplerCustomAdvanced:
@classmethod
def INPUT_TYPES(s):
+20 -21
View File
@@ -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
View File
@@ -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
View File
@@ -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"
+74
View File
@@ -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"
+51
View File
@@ -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()