Compare commits

..
1 Commits
Author SHA1 Message Date
charrywhite 111f9fdffb 666 version code 2026-02-02 21:36:01 +08:00
25 changed files with 297 additions and 5151 deletions
-1
View File
@@ -26,7 +26,6 @@ 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,5 +11,3 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: comfy-org/node-diff@main
with:
base_ref: ${{ github.event.repository.default_branch }}
+3 -47
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/yN5wYDE6W4)
[![Discord](https://img.shields.io/badge/Discord-5865F2?style=for-the-badge&logo=discord&logoColor=white)](https://discord.gg/aCGZutBV)
</div>
@@ -31,7 +31,7 @@ note={}
```
**🎉 NEW 2026: Join our discord!**
[Join our Discord](https://discord.gg/yN5wYDE6W4) to share experiences, discuss features, and explore future development.
[Join our Discord](https://discord.gg/aCGZutBV) to share experiences, discuss features, and explore future development.
**🎬 NEW: LanPaint now supports inpainting and outpainting based on Z-Image!**
@@ -39,12 +39,6 @@ note={}
|:--------:|:------:|:---------:|
| ![Original Z-image](https://github.com/scraed/LanPaint/blob/master/examples/Example_21/Original_No_Mask.png) | ![Masked Z-image](https://github.com/scraed/LanPaint/blob/master/examples/Example_21/Masked_Load_Me_in_Loader.png) | ![Inpainted Z-image](https://github.com/scraed/LanPaint/blob/master/examples/Example_21/InPainted_Drag_Me_to_ComfyUI.png) |
**🎬 NEW: LanPaint now supports Z-Image-Base too!**
| Original | Masked | Inpainted |
|:--------:|:------:|:---------:|
| ![Original Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/Original_No_Mask.png) | ![Masked Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/Masked_Load_Me_in_Loader.png) | ![Inpainted Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/InPainted_Drag_Me_to_ComfyUI.png) |
**🎬 NEW: LanPaint now supports video inpainting and outpainting based on Wan 2.2!**
@@ -73,9 +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)
- [Z-image-base](#example-z-image-base-inpaintlanpaint-k-sampler-3-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)
- [Wan 2.2 T2I with reference](#example-wan22-partial-inpaintlanpaint-k-sampler-5-steps-of-thinking)
@@ -98,7 +90,7 @@ Check our latest [Wan 2.2 Video Examples](#video-examples-beta), [Wan 2.2 Image
## Features
- **Universal Compatibility** – Works instantly with almost any model (**Z-image, Z-image-base, Hunyuan, Wan 2.2, Qwen Image/Edit, HiDream, SD 3.5, Flux-series, SDXL, SD 1.5 or custom LoRAs**) and ControlNet.
- **Universal Compatibility** – Works instantly with almost any model (**Z-image, Hunyuan, Wan 2.2, Qwen Image/Edit, HiDream, SD 3.5, Flux-series, SDXL, SD 1.5 or custom LoRAs**) and ControlNet.
![Inpainting Result 13](https://github.com/scraed/LanPaint/blob/master/examples/InpaintChara_13.jpg)
- **No Training Needed** – Works out of the box with your existing model.
- **Easy to Use** – Same workflow as standard ComfyUI KSampler.
@@ -281,24 +273,6 @@ LanPaint also supports inpainting with the Z-image text-to-image model.
You can download the Z-image model for ComfyUI from [Z-image](https://docs.comfy.org/zh-CN/tutorials/image/z-image/z-image-turbo).
### Example Z-image-base: InPaint(LanPaint K Sampler, 3 steps of thinking)
LanPaint also supports inpainting with the Z-image-base model.
**Warning (stability)**: Z-image-base can easily diverge with LanPaint. Start with **small `LanPaint_StepSize`** and **fewer thinking iterations** (lower `LanPaint_NumSteps`) and increase gradually only if stable.
<details open>
<summary>View Original / Masked / Inpainted Comparison</summary>
| Original | Masked | Inpainted |
|:--------:|:------:|:---------:|
| ![Original Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/Original_No_Mask.png) | ![Masked Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/Masked_Load_Me_in_Loader.png) | ![Inpainted Z-image-base](https://github.com/scraed/LanPaint/blob/master/examples/Example_25/InPainted_Drag_Me_to_ComfyUI.png) |
</details>
[View Workflow & Masks](https://github.com/scraed/LanPaint/tree/master/examples/Example_25)
Workflow template (JSON): [Z_image_base_Inpaint.json](https://github.com/scraed/LanPaint/blob/master/example_workflows/Z_image_base_Inpaint.json)
### Example Wan2.2: Partial InPaint(LanPaint K Sampler, 5 steps of thinking)
Sometimes we don't want to inpaint completely new content, but rather let the inpainted image reference the original image. One option to achieve this is to inpaint with an edit model like Qwen Image Edit. Another option is to perform a partial inpaint: allowing the diffusion process to start at some middle steps rather than from 0.
@@ -368,22 +342,6 @@ 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)
@@ -504,8 +462,6 @@ Submit a PR to add your tutorial/video here, or open an [Issue](https://github.c
[Working togather with crop&stitch](https://github.com/scraed/LanPaint/issues/46)
## Updates
- 2026/01/30
- Add Z-image-base documentation and Example_25 workflow images.
- 2025/08/08
- Add Qwen image support
- 2025/06/21
+2 -80
View File
@@ -10,85 +10,7 @@ __author__ = """LanPaint"""
__email__ = "czhengac@connect.ust.hk"
__version__ = "0.0.1"
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
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
WEB_DIRECTORY = "./web"
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.6 MiB

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

Before

Width:  |  Height:  |  Size: 1.1 MiB

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

Before

Width:  |  Height:  |  Size: 1.8 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.8 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.8 MiB

+1 -4
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "LanPaint"
version = "1.4.13"
version = "1.4.9"
description = "Achieve seamless inpainting results without needing a specialized inpainting model."
authors = [
{name = "LanPaint", email = "czhengac@connect.ust.hk"}
@@ -75,8 +75,5 @@ 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"
-337
View File
@@ -1,337 +0,0 @@
"""
Early Stop Logic Contributed by `https://github.com/godnight10061`.
"""
import inspect
from typing import Any, Callable, Optional
import torch
from .types import LangevinState
def _clamp01(val: float) -> float:
if val <= 0.0:
return 0.0
if val >= 1.0:
return 1.0
return val
def _abt_scale(abt_val: float) -> float:
"""
Smooth, parameter-free scale based on outer-step noise level.
- 0 at abt=0/1 (disable at extreme noise / extreme tail)
- 1 at abt=0.5 (mid-schedule)
"""
abt_val = _clamp01(abt_val)
return _clamp01(4.0 * abt_val * (1.0 - abt_val))
def _boundary_weight(latent_mask: torch.Tensor, inpaint_weight: torch.Tensor) -> Optional[torch.Tensor]:
"""
Return a 4-neighbor boundary weight: unknown pixels adjacent to known pixels.
This replaces the previous dilation-based "ring" (kernel/padding) and has no tunable hyperparameters.
"""
if latent_mask.dim() != 4:
return None
known = latent_mask > 0.5
neighbor_known = torch.zeros_like(known)
neighbor_known[:, :, 1:, :] |= known[:, :, :-1, :]
neighbor_known[:, :, :-1, :] |= known[:, :, 1:, :]
neighbor_known[:, :, :, 1:] |= known[:, :, :, :-1]
neighbor_known[:, :, :, :-1] |= known[:, :, :, 1:]
boundary = (~known) & neighbor_known
return boundary.to(dtype=torch.float32) * inpaint_weight
def _weighted_mse(t1: torch.Tensor, t2: torch.Tensor, weight: torch.Tensor) -> float:
diff_sq = (t1.to(dtype=torch.float32) - t2.to(dtype=torch.float32)) ** 2
denom = torch.sum(weight) + 1e-12
return float((torch.sum(diff_sq * weight) / denom).item())
class LanPaintEarlyStopper:
"""
Per-step early-stop logic for LanPaint inner (Langevin) iterations.
"""
@classmethod
def from_options(
cls,
*,
model_options: Optional[dict],
latent_mask: torch.Tensor,
abt: torch.Tensor,
default_threshold: float,
default_patience: int,
default_distance_fn: Optional[Callable[..., Any]],
) -> Optional["LanPaintEarlyStopper"]:
semantic_stop = model_options.get("lanpaint_semantic_stop") if isinstance(model_options, dict) else None
threshold = float(default_threshold)
patience = int(default_patience)
distance_fn = default_distance_fn
# distance_fn contract: return None (use default metric) or a scalar (Python number / 0-d (1-element) torch.Tensor)
if isinstance(semantic_stop, dict):
threshold = float(semantic_stop.get("threshold", threshold))
patience = int(semantic_stop.get("patience", patience))
distance_fn = semantic_stop.get("distance_fn", distance_fn)
# Backward compatibility: map legacy 'min_steps' to a patience floor so it is not an independent knob.
if patience > 0:
min_steps = semantic_stop.get("min_steps")
if min_steps is not None:
try:
min_steps_int = int(min_steps)
except (TypeError, ValueError):
min_steps_int = 0
if min_steps_int > 1:
patience = max(patience, min_steps_int - 1)
enabled_early_stop = (threshold > 0.0) and (patience > 0)
# Require N+1 consecutive stable checks:
# - the first stable step sets patience_counter to 1
# - `patience=1` therefore stops after 2 stable steps
patience_eff = max(1, patience) + 1
threshold_eff = threshold
inpaint_weight = ring_weight = trace = abt_val = None
if enabled_early_stop:
try:
abt_val = float(torch.mean(abt).item())
except (TypeError, ValueError):
abt_val = 0.0
threshold_eff = threshold * _abt_scale(abt_val)
if threshold_eff <= 0.0:
enabled_early_stop = False
else:
inpaint_weight = (1 - latent_mask).to(dtype=torch.float32)
if float(torch.sum(inpaint_weight).item()) < 1e-6:
enabled_early_stop = False
else:
ring_weight = _boundary_weight(latent_mask, inpaint_weight)
if isinstance(model_options, dict):
trace = model_options.get("lanpaint_semantic_trace")
if not enabled_early_stop:
return None
# Pre-fetch trace keys to avoid repeated dict lookups
bench_case_id = bench_outer_step = bench_timestep = None
if isinstance(trace, list) and isinstance(model_options, dict):
bench_case_id = model_options.get("bench_case_id")
bench_outer_step = model_options.get("bench_outer_step")
bench_timestep = model_options.get("bench_timestep")
return cls(
enabled=enabled_early_stop,
threshold=threshold,
threshold_eff=threshold_eff,
patience_eff=patience_eff,
inpaint_weight=inpaint_weight,
ring_weight=ring_weight,
distance_fn=distance_fn,
trace=trace,
bench_case_id=bench_case_id,
bench_outer_step=bench_outer_step,
bench_timestep=bench_timestep,
abt_val=abt_val,
)
def __init__(
self,
*,
enabled: bool,
threshold: float,
threshold_eff: float,
patience_eff: int,
inpaint_weight: Optional[torch.Tensor],
ring_weight: Optional[torch.Tensor],
distance_fn: Optional[Callable[..., Any]] = None,
trace: Optional[list] = None,
bench_case_id: Any = None,
bench_outer_step: Any = None,
bench_timestep: Any = None,
abt_val: Optional[float] = None,
) -> None:
self.enabled = bool(enabled)
self.threshold = float(threshold)
self.threshold_eff = float(threshold_eff)
self.patience_eff = int(patience_eff)
self.inpaint_weight = inpaint_weight
self.ring_weight = ring_weight
self.trace = trace
self.bench_case_id = bench_case_id
self.bench_outer_step = bench_outer_step
self.bench_timestep = bench_timestep
self.abt_val = abt_val
self.patience_counter = 0
self.x0_anchor = None
self._dist_wrapper = self._wrap_distance_fn(distance_fn) if self.enabled else None
@property
def has_custom_distance_fn(self) -> bool:
return self._dist_wrapper is not None
@staticmethod
def _wrap_distance_fn(distance_fn: Optional[Callable[..., Any]]):
"""
Wrap a user-provided `distance_fn` into a normalized callable: fn(prev, cur, ctx) -> dist|None.
Supported signatures:
- 3+ positional (or *args): `distance_fn(prev, cur, ctx)`
- explicit / **kwargs ctx: `distance_fn(prev, cur, ctx=ctx)`
- default 2-arg: `distance_fn(cur, prev)`
Return contract: None (use default metric) or a scalar (Python number / 0-d (1-element) torch.Tensor).
"""
if not callable(distance_fn):
return None
try:
sig = inspect.signature(distance_fn)
params = list(sig.parameters.values())
has_ctx_param = "ctx" in sig.parameters
has_var_kw = any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params)
has_var_pos = any(p.kind == inspect.Parameter.VAR_POSITIONAL for p in params)
pos_params = [
p
for p in params
if p.kind in (inspect.Parameter.POSITIONAL_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD)
]
if len(pos_params) >= 3 or has_var_pos:
# 3-arg positional: fn(prev, cur, ctx)
return lambda p, c, ctx: distance_fn(p, c, ctx)
if has_ctx_param or has_var_kw:
# keyword ctx: fn(prev, cur, ctx=ctx)
return lambda p, c, ctx: distance_fn(p, c, ctx=ctx)
# Default 2-arg: fn(cur, prev)
return lambda p, c, ctx: distance_fn(c, p)
except (ValueError, TypeError):
# Fallback for built-ins or complex callables.
def fallback_wrapper(p, c, ctx):
try:
return distance_fn(p, c, ctx)
except TypeError as e:
tb = e.__traceback__
if tb is not None and tb.tb_frame.f_code is not fallback_wrapper.__code__:
raise
return distance_fn(c, p)
return fallback_wrapper
def step(
self,
*,
i: int,
n_steps: int,
x_t_before: torch.Tensor,
x_t_after: torch.Tensor,
x_t_prev_for_custom: Optional[torch.Tensor],
prev_args: Any,
args: Any,
ctx: dict,
) -> bool:
if not self.enabled:
return False
# 'inpaint_weight' is guaranteed to be set when enabled is True in the caller.
inpaint = self.inpaint_weight
if inpaint is None:
return False
dist = None
custom_dist = False
dist_inpaint = dist_ring = dist_drift = x0_prev = x0_cur = None
if self._dist_wrapper is not None:
dist = self._dist_wrapper(x_t_prev_for_custom, x_t_after, ctx)
if dist is not None:
if isinstance(dist, torch.Tensor):
if dist.numel() != 1:
raise TypeError("distance_fn must return None or a scalar / 0-d (1-element) tensor")
dist = float(dist.item())
else:
dist = float(dist)
custom_dist = dist is not None
if dist is None:
def _get_x0(arg: Any) -> Optional[torch.Tensor]:
if isinstance(arg, LangevinState):
return arg.x0
if isinstance(arg, tuple) and len(arg) >= 3:
return arg[2]
return None
x0_prev = _get_x0(prev_args)
x0_cur = _get_x0(args)
if x0_prev is not None and x0_cur is not None:
dist_inpaint = _weighted_mse(x0_cur, x0_prev, inpaint)
dist_ring = _weighted_mse(x0_cur, x0_prev, self.ring_weight) if self.ring_weight is not None else None
dist = dist_inpaint if dist_ring is None else max(dist_inpaint, dist_ring)
else:
dist_inpaint = _weighted_mse(x_t_after, x_t_before, inpaint)
dist = dist_inpaint
threshold_used = self.threshold if custom_dist else self.threshold_eff
# Drift guard (only for default metric with x0_cur).
if x0_cur is not None and not custom_dist:
if dist <= threshold_used:
if self.x0_anchor is None:
self.x0_anchor = x0_cur.detach()
else:
drift_inpaint = _weighted_mse(x0_cur, self.x0_anchor, inpaint)
drift_ring = _weighted_mse(x0_cur, self.x0_anchor, self.ring_weight) if self.ring_weight is not None else None
dist_drift = drift_inpaint if drift_ring is None else max(drift_inpaint, drift_ring)
dist = max(dist, dist_drift)
else:
self.x0_anchor = None
if dist <= threshold_used:
self.patience_counter += 1
else:
self.patience_counter = 0
self.x0_anchor = None
should_stop = self.patience_counter >= self.patience_eff
if isinstance(self.trace, list):
self.trace.append(
{
"case_id": self.bench_case_id,
"outer_step": self.bench_outer_step,
"bench_timestep": self.bench_timestep,
"inner_step": i + 1,
"dist": dist,
"dist_inpaint": None if dist_inpaint is None else float(dist_inpaint),
"dist_ring": None if dist_ring is None else float(dist_ring),
"dist_drift": None if dist_drift is None else float(dist_drift),
"threshold": float(threshold_used),
"threshold_eff": float(self.threshold_eff),
"patience_counter": int(self.patience_counter),
"patience_eff": int(self.patience_eff),
"abt": None if self.abt_val is None else float(self.abt_val),
"custom_dist": bool(custom_dist),
"stopped": bool(should_stop),
}
)
return bool(should_stop)
+31 -118
View File
@@ -1,11 +1,9 @@
import torch
from .utils import StochasticHarmonicOscillator
from .utils import *
from functools import partial
from .earlystop import LanPaintEarlyStopper
from .types import LangevinState
class LanPaint():
def __init__(self, Model, NSteps, Friction, Lambda, Beta, StepSize, IS_FLUX = False, IS_FLOW = False, EarlyStopThreshold = 0.0, EarlyStopPatience = 1, EarlyStopHook = None):
def __init__(self, Model, NSteps, Friction, Lambda, Beta, StepSize, IS_FLUX = False, IS_FLOW = False):
self.n_steps = NSteps
self.chara_lamb = Lambda
self.IS_FLUX = IS_FLUX
@@ -15,9 +13,6 @@ class LanPaint():
self.friction = Friction
self.chara_beta = Beta
self.img_dim_size = None
self.early_stop_threshold = EarlyStopThreshold
self.early_stop_patience = EarlyStopPatience
self.early_stop_hook = EarlyStopHook
def add_none_dims(self, array):
# Create a tuple with ':' for the first dimension and 'None' repeated num_nones times
@@ -39,6 +34,7 @@ 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
@@ -47,6 +43,8 @@ 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 )
@@ -56,46 +54,9 @@ class LanPaint():
############ LanPaint Iterations Start ###############
# after noise_scaling, noise = latent_image + noise * sigma, which is x_t in the variance exploding diffusion model notation for the known region.
args = None
stopper = LanPaintEarlyStopper.from_options(
model_options=model_options if isinstance(model_options, dict) else None,
latent_mask=latent_mask,
abt=abt,
default_threshold=self.early_stop_threshold,
default_patience=self.early_stop_patience,
default_distance_fn=self.early_stop_hook,
)
for i in range(n_steps):
score_func = partial( self.score_model, y = self.latent_image, mask = latent_mask, abt = self.add_none_dims(abt), sigma = self.add_none_dims(VE_Sigma), tflow = self.add_none_dims(Flow_t), model_options = model_options, seed = seed )
prev_args = args
x_t_prev = x_t.detach() if (stopper is not None and stopper.has_custom_distance_fn) else None
x_t_before = x_t if (stopper is not None and stopper.enabled) else None
x_t, args = self.langevin_dynamics(x_t, score_func , latent_mask, step_size , current_times, sigma_x = self.add_none_dims(self.sigma_x(abt)), sigma_y = self.add_none_dims(self.sigma_y(abt)), args = args)
if stopper is not None:
ctx = {
"step": i,
"steps_done": i + 1,
"n_steps": n_steps,
"mask": latent_mask,
"latent_image": self.latent_image,
"current_times": current_times,
"seed": seed,
}
if stopper.step(
i=i,
n_steps=n_steps,
x_t_before=x_t_before,
x_t_after=x_t,
x_t_prev_for_custom=x_t_prev,
prev_args=prev_args,
args=args,
ctx=ctx,
):
break
if IS_FLUX or IS_FLOW:
x = x_t / ( self.add_none_dims(abt)**0.5 + (1-self.add_none_dims(abt))**0.5 )
else:
@@ -107,6 +68,7 @@ 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.
@@ -127,13 +89,6 @@ class LanPaint():
return beta
def langevin_dynamics(self, x_t, score, mask, step_size, current_times, sigma_x=1, sigma_y=0, args=None):
if args is not None and not isinstance(args, LangevinState):
if isinstance(args, tuple):
if len(args) == 2:
# Backwards compat: older state was (v, C) without x0.
args = LangevinState(args[0], args[1], None)
elif len(args) >= 3:
args = LangevinState(args[0], args[1], args[2])
# prepare the step size and time parameters
with torch.autocast(device_type=x_t.device.type, dtype=torch.float32):
step_sizes = self.prepare_step_size(current_times, step_size, sigma_x, sigma_y)
@@ -151,10 +106,11 @@ class LanPaint():
dt = dtx * (1-mask) + dty * mask
Gamma = Gamma_x * (1-mask) + Gamma_y * mask
def Coef_C(x_t):
x0 = x_t + score(x_t)
x0 = self.x0_evalutation(x_t, score, sigma, args)
C = (abt**0.5 * x0 - x_t )/ (1-abt) + A * x_t
return C, x0
return C
def advance_time(x_t, v, dt, Gamma, A, C, D):
dtype = x_t.dtype
with torch.autocast(device_type=x_t.device.type, dtype=torch.float32):
@@ -163,74 +119,25 @@ 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
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)
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
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))
C_new = Coef_C(x_t)
v = v + Gamma**0.5 * ( C_new - C) *dt
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)
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
def run_damped(x_t, args):
if args is None:
v = None
C, x0 = Coef_C(x_t)
x_t, v = advance_time(x_t, v, dt, Gamma, A, C, D)
else:
v = args.v
C = args.C
x_t, v = advance_time(x_t, v, dt/2, Gamma, A, C, D)
C_new, x0 = 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
# args is (v, C, x0) for the next inner step.
return x_t, LangevinState(v, C, x0)
def run_overdamped(x_t, args):
if args is None:
C, x0 = Coef_C(x_t)
x_t = advance_time_overdamped(x_t, dt, A, C, D)
else:
C = args.C
x_t = advance_time_overdamped(x_t, dt / 2, A, C, D)
C_new, x0 = 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
# args is (v, C, x0); v is None in the overdamped fallback.
return x_t, LangevinState(None, C, x0)
try:
x_t_next, state = run_damped(x_t, args)
v_next = state.v
if torch.isnan(x_t_next).any() or (v_next is not None and torch.isnan(v_next).any()):
raise ValueError("NaN detected")
x_t = x_t_next
except Exception:
x_t, state = run_overdamped(x_t, args)
# args is (v, C, x0); v can be None if we fell back to the overdamped update.
return x_t, state
C = C_new
return x_t, (v, C)
def prepare_step_size(self, current_times, step_size, sigma_x, sigma_y):
# -------------------------------------------------------------------------
@@ -241,7 +148,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.
@@ -266,3 +173,9 @@ class LanPaint():
D_x = (2 * abt**0 )**0.5
D_y = (2 * abt**0 )**0.5
return sigma, abt, dtx/2, dty/2, Gamma_x, Gamma_y, A_x, A_y, D_x, D_y
def x0_evalutation(self, x_t, score, sigma, args):
x0 = x_t + score(x_t)
return x0
+219 -83
View File
@@ -1,83 +1,121 @@
from contextlib import contextmanager
import math
from inspect import cleandoc
import inspect
# import nodes.py
import comfy
import nodes
import latent_preview
import torch
from functools import partial
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)
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_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 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)
# 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: # 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
else: # Original 2D image case
if not COMFYUI_VERSION_060_OR_NEWER:
if comfyui_version.__version__ < "0.6.0":
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):
@@ -108,9 +146,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)
@@ -158,8 +196,6 @@ 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 )
@@ -169,6 +205,31 @@ 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
@@ -183,7 +244,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)
@@ -194,7 +255,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
@@ -203,6 +264,7 @@ 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?
@@ -227,10 +289,7 @@ class KSAMPLER(comfy.samplers.KSAMPLER):
model_wrap.model_patcher.LanPaint_Beta,
model_wrap.model_patcher.LanPaint_StepSize,
IS_FLUX = IS_FLUX,
IS_FLOW = IS_FLOW,
EarlyStopThreshold = getattr(model_wrap.model_patcher, "LanPaint_InnerThreshold", 0.0),
EarlyStopPatience = getattr(model_wrap.model_patcher, "LanPaint_InnerPatience", 1),
EarlyStopHook = extra_args.get("model_options", {}).get("lanpaint_semantic_hook", None))
IS_FLOW = IS_FLOW)
model_k.LanPaint_early_stop = model_wrap.model_patcher.LanPaint_EarlyStop
#if not inpainting, after noise_scaling, noise = noise * sigma, which is the noise added to the clean latent image in the variance exploding diffusion model notation.
#if inpainting, after noise_scaling, noise = latent_image + noise * sigma, which is x_t in the variance exploding diffusion model notation for the known region.
@@ -332,19 +391,17 @@ class LanPaint_KSampler():
model.LanPaint_NumSteps = LanPaint_NumSteps
model.LanPaint_Friction = 15.
model.LanPaint_EarlyStop = 1
model.LanPaint_InnerThreshold = 0.0
model.LanPaint_InnerPatience = 1
if LanPaint_PromptMode == "Image First":
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:
@@ -373,8 +430,6 @@ class LanPaint_KSamplerAdvanced:
"LanPaint_EarlyStop": ("INT", {"default": 1, "min": 0, "max": 10000, "tooltip": "The number of steps to stop the LanPaint early, useful for preventing the image from irregular patterns."}),
"LanPaint_Info": ("STRING", {"default": "LanPaint KSampler 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"}),
"LanPaint_InnerThreshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.0001, "round": 0.0001, "tooltip": "Early stop threshold for Langevin iterations based on semantic distance. 0.0 to disable. (Contributed by godnight10061)"}),
"LanPaint_InnerPatience": ("INT", {"default": 1, "min": 1, "max": 100, "tooltip": "Number of consecutive steps below threshold required to stop. (Contributed by godnight10061)"}),
},
}
@@ -383,7 +438,7 @@ class LanPaint_KSamplerAdvanced:
CATEGORY = "sampling"
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, LanPaint_NumSteps=5, LanPaint_Lambda=16.0, LanPaint_StepSize=0.2, LanPaint_Beta=1.0, LanPaint_Friction=15.0, LanPaint_PromptMode="Image First", LanPaint_EarlyStop=1, LanPaint_Info="", Inpainting_mode="🖼️ Image Inpainting", LanPaint_InnerThreshold=0.0, LanPaint_InnerPatience=1):
def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, LanPaint_NumSteps=5, LanPaint_Lambda=16.0, LanPaint_StepSize=0.2, LanPaint_Beta=1.0, LanPaint_Friction=15.0, LanPaint_PromptMode="Image First", LanPaint_EarlyStop=1, LanPaint_Info="", Inpainting_mode="🖼️ Image Inpainting"):
force_full_denoise = True
if return_with_leftover_noise == "enable":
force_full_denoise = False
@@ -396,13 +451,11 @@ class LanPaint_KSamplerAdvanced:
model.LanPaint_NumSteps = LanPaint_NumSteps
model.LanPaint_Friction = LanPaint_Friction
model.LanPaint_EarlyStop = LanPaint_EarlyStop
model.LanPaint_InnerThreshold = LanPaint_InnerThreshold
model.LanPaint_InnerPatience = LanPaint_InnerPatience
if LanPaint_PromptMode == "Image First":
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:
@@ -454,7 +507,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,:,:]
@@ -476,6 +529,77 @@ 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"])
@@ -504,6 +628,7 @@ 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"}),
}
}
@@ -512,19 +637,21 @@ 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=""):
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"):
model.LanPaint_StepSize = 0.2
model.LanPaint_Lambda = 16.0
model.LanPaint_Beta = 1.
model.LanPaint_NumSteps = LanPaint_NumSteps
model.LanPaint_Friction = 15.
model.LanPaint_EarlyStop = 1
model.LanPaint_InnerThreshold = 0.0
model.LanPaint_InnerPatience = 1
if LanPaint_PromptMode == "Image First":
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"]
@@ -554,7 +681,7 @@ class LanPaint_SamplerCustom:
else:
out_denoised = out
return (out, out_denoised)
class LanPaint_SamplerCustomAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -572,8 +699,7 @@ 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}),
"LanPaint_InnerThreshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.0001, "round": 0.0001, "tooltip": "Early stop threshold for Langevin iterations based on semantic distance. 0.0 to disable. (Contributed by godnight10061)"}),
"LanPaint_InnerPatience": ("INT", {"default": 1, "min": 1, "max": 100, "tooltip": "Number of consecutive steps below threshold required to stop. (Contributed by godnight10061)"}),
"Inpainting_mode": (["🖼️ Image Inpainting", "🎬 Video Inpainting"], {"default": "🖼️ Image Inpainting", "tooltip": "Choose Image mode for photos or Video mode for video frames with temporal consistency"}),
}
}
@@ -584,7 +710,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="", LanPaint_InnerThreshold=0.0, LanPaint_InnerPatience=1):
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"):
model = guider.model_patcher
model.LanPaint_StepSize = LanPaint_StepSize
model.LanPaint_Lambda = LanPaint_Lambda
@@ -592,22 +718,29 @@ class LanPaint_SamplerCustomAdvanced:
model.LanPaint_NumSteps = LanPaint_NumSteps
model.LanPaint_Friction = LanPaint_Friction
model.LanPaint_EarlyStop = LanPaint_EarlyStop
model.LanPaint_InnerThreshold = LanPaint_InnerThreshold
model.LanPaint_InnerPatience = LanPaint_InnerPatience
if LanPaint_PromptMode == "Image First":
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)
@@ -623,6 +756,7 @@ 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)
@@ -634,6 +768,7 @@ NODE_CLASS_MAPPINGS = {
"LanPaint_SamplerCustom" : LanPaint_SamplerCustom,
"LanPaint_SamplerCustomAdvanced" : LanPaint_SamplerCustomAdvanced,
"LanPaint_MaskBlend": MaskBlend,
"LanPaint_MaskBlendAlpha": MaskBlendAlpha,
# "LanPaint_UpSale_LatentNoiseMask": LanPaint_UpSale_LatentNoiseMask,
}
@@ -644,5 +779,6 @@ 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"
}
-10
View File
@@ -1,10 +0,0 @@
from typing import NamedTuple, Optional
import torch
class LangevinState(NamedTuple):
v: Optional[torch.Tensor]
C: Optional[torch.Tensor]
x0: Optional[torch.Tensor]
+21 -20
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,6 +68,7 @@ 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))
@@ -116,15 +117,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)
@@ -132,17 +133,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
"""
@@ -150,17 +151,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)
@@ -169,7 +170,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
@@ -184,7 +185,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
@@ -207,7 +208,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
@@ -238,7 +239,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
@@ -253,12 +254,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)
@@ -273,7 +274,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)
@@ -297,4 +298,4 @@ class StochasticHarmonicOscillator:
scale_tril=scale_tril
).sample()
return new_yv[...,0], new_yv[...,1]
return new_yv[...,0], new_yv[...,1]
+3 -4
View File
@@ -1,5 +1,4 @@
[pytest]
# Keep settings value-only; pytest does not treat inline `# ...` as comments.
testpaths = .
python_files = test_*.py
norecursedirs = ..
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
+17 -9
View File
@@ -1,13 +1,21 @@
"""Basic import tests for LanPaint.
#!/usr/bin/env python
The ComfyUI runtime dependencies (e.g. `comfy`) are intentionally optional for unit tests.
"""
"""Tests for `LanPaint` package."""
import pytest
from src.LanPaint.nodes import Example
def test_package_imports_without_comfy() -> None:
import LanPaint
@pytest.fixture
def example_node():
"""Fixture to create an Example node instance."""
return 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"
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"
-104
View File
@@ -1,104 +0,0 @@
import torch
from src.LanPaint.lanpaint import LanPaint as LanPaintEngine
class _DummySampling:
def noise_scaling(self, sigma, noise, latent_image): # type: ignore[no-untyped-def]
return latent_image + noise * sigma
class _DummyModel:
def __init__(self) -> None:
self.inner_model = self
self.model_sampling = _DummySampling()
def __call__(self, x, sigma, model_options=None, seed=None): # type: ignore[no-untyped-def]
return x, x
def _inputs(): # type: ignore[no-untyped-def]
x = torch.zeros((1, 4, 8, 8))
latent_image = torch.zeros_like(x)
noise = torch.ones_like(x)
sigma = torch.tensor([1.0])
latent_mask = torch.zeros_like(x)
current_times = (sigma, torch.tensor([0.5]), torch.tensor([0.0]))
return x, latent_image, noise, sigma, latent_mask, current_times
def test_default_semantic_stop_triggers_at_patience_without_custom_distance_fn() -> None:
engine = LanPaintEngine(
_DummyModel(),
NSteps=10,
Friction=15.0,
Lambda=1.0,
Beta=1.0,
StepSize=0.2,
)
calls = {"langevin": 0, "with_score": 0, "without_score": 0}
def fake_langevin(x_t, score, mask, step_size, current_times, sigma_x=1, sigma_y=0, args=None): # type: ignore[no-untyped-def]
calls["langevin"] += 1
if score is None:
calls["without_score"] += 1
else:
calls["with_score"] += 1
return x_t, args
engine.langevin_dynamics = fake_langevin # type: ignore[method-assign]
model_options = {
"lanpaint_semantic_stop": {
"threshold": 1e-6,
"patience": 2,
}
}
x, latent_image, noise, sigma, latent_mask, current_times = _inputs()
engine(x, latent_image, noise, sigma, latent_mask, current_times, model_options=model_options, seed=0, n_steps=10)
assert calls["langevin"] == 3
assert calls["with_score"] == 3
assert calls["without_score"] == 0
def test_semantic_stop_is_disabled_when_no_inpaint_region() -> None:
engine = LanPaintEngine(
_DummyModel(),
NSteps=10,
Friction=15.0,
Lambda=1.0,
Beta=1.0,
StepSize=0.2,
)
calls = {"langevin": 0, "with_score": 0, "without_score": 0}
def fake_langevin(x_t, score, mask, step_size, current_times, sigma_x=1, sigma_y=0, args=None): # type: ignore[no-untyped-def]
calls["langevin"] += 1
if score is None:
calls["without_score"] += 1
else:
calls["with_score"] += 1
return x_t, args
engine.langevin_dynamics = fake_langevin # type: ignore[method-assign]
model_options = {
"lanpaint_semantic_stop": {
"threshold": 1e-6,
"patience": 1,
}
}
x, latent_image, noise, sigma, latent_mask, _ = _inputs()
current_times = (sigma, torch.tensor([0.5]), torch.tensor([0.0]))
no_inpaint_mask = torch.ones_like(latent_mask)
engine(x, latent_image, noise, sigma, no_inpaint_mask, current_times, model_options=model_options, seed=0, n_steps=10)
assert calls["langevin"] == 10
assert calls["with_score"] == 10
assert calls["without_score"] == 0
-74
View File
@@ -1,74 +0,0 @@
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"
-45
View File
@@ -1,45 +0,0 @@
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, args_out = lp.langevin_dynamics(x_t, score, mask, step_size, current_times, sigma_y=1.0)
assert hasattr(args_out, "v")
assert hasattr(args_out, "C")
assert hasattr(args_out, "x0")
assert args_out[0] is args_out.v
assert args_out[1] is args_out.C
assert args_out[2] is args_out.x0
v_out = args_out[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()