Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
24051fb151 | ||
|
|
a76a5da0c5 | ||
|
|
21a6b155e8 | ||
|
|
8bc20d5e00 | ||
|
|
84b63adb11 | ||
|
|
725fb2ba37 | ||
|
|
4f19674299 | ||
|
|
558084416d | ||
|
|
a984dcb118 | ||
|
|
3f9193997a | ||
|
|
7372428114 | ||
|
|
ac6eece747 | ||
|
|
95b65ed1b1 | ||
|
|
30cfd2f2e4 | ||
|
|
2e3a566bb4 | ||
|
|
8fe1a9d312 | ||
|
|
4fcbfc7b5b | ||
|
|
651ddc2225 | ||
|
|
bd4606f27a | ||
|
|
8a72294f27 | ||
|
|
9ddd5f613e | ||
|
|
39c90bbf62 | ||
|
|
6ad4a18537 | ||
|
|
5e7d472738 | ||
|
|
d73f3db265 | ||
|
|
cd06069b52 | ||
|
|
114d091965 | ||
|
|
06da285a09 | ||
|
|
8f52a1d378 | ||
|
|
c7962bf784 | ||
|
|
a895d2a758 | ||
|
|
34ded92d8d | ||
|
|
1b87239553 | ||
|
|
cbc9f3abfd | ||
|
|
7d46147887 | ||
|
|
c5fe6351ea | ||
|
|
ac7a5a1f2f | ||
|
|
73ab537e86 | ||
|
|
6fd5f60383 | ||
|
|
9c9b5555fc | ||
|
|
450a25e5e8 | ||
|
|
e7080f9574 | ||
|
|
2086f00603 | ||
|
|
896729ccd4 | ||
|
|
207860edc3 | ||
|
|
4074473b90 | ||
|
|
11c1490140 | ||
|
|
7c79712fc8 | ||
|
|
fb35a1d032 | ||
|
|
060fcc4475 | ||
|
|
aedb908a3f | ||
|
|
f9718ea3e1 | ||
|
|
0bafe2117a | ||
|
|
2c1a23777d | ||
|
|
dde4c82463 | ||
|
|
72cb484c40 | ||
|
|
0d15d15cf3 | ||
|
|
1393a46d67 | ||
|
|
27421363bf | ||
|
|
01400a541d | ||
|
|
c981354387 | ||
|
|
c92c8bfb1f | ||
|
|
9334025801 | ||
|
|
144893ecd2 | ||
|
|
bce8505ce2 | ||
|
|
0992657e54 | ||
|
|
2fa036992b | ||
|
|
9432a34c38 | ||
|
|
d4e8ee28fb | ||
|
|
34b8e83831 | ||
|
|
e29d480a1a |
@@ -26,6 +26,7 @@ jobs:
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install .[dev]
|
||||
pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
|
||||
- name: Run Linting
|
||||
run: |
|
||||
ruff check .
|
||||
|
||||
@@ -11,3 +11,5 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: comfy-org/node-diff@main
|
||||
with:
|
||||
base_ref: ${{ github.event.repository.default_branch }}
|
||||
|
||||
@@ -7,13 +7,19 @@
|
||||
[](https://huggingface.co/charrywhite/LanPaint)
|
||||
[](https://scraed.github.io/scraedBlog/)
|
||||
[](https://github.com/scraed/LanPaint/stargazers)
|
||||
[](https://discord.gg/aCGZutBV)
|
||||
[](https://discord.gg/yN5wYDE6W4)
|
||||
</div>
|
||||
|
||||
|
||||
Universally applicable inpainting ability for every model. LanPaint sampler lets the model "think" through multiple iterations before denoising, enabling you to invest more computation time for superior inpainting quality.
|
||||
|
||||
This is the official implementation of ["LanPaint: Training-Free Diffusion Inpainting with Asymptotically Exact and Fast Conditional Sampling"](https://arxiv.org/abs/2502.03491), accepted by TMLR. The repository is for ComfyUI extension. Local Python benchmark code is published here: [LanPaintBench](https://github.com/scraed/LanPaintBench).
|
||||
This is the official implementation of ["LanPaint: Training-Free Diffusion Inpainting with Asymptotically Exact and Fast Conditional Sampling"](https://arxiv.org/abs/2502.03491), accepted by TMLR.
|
||||
|
||||
The repository is for ComfyUI extension.
|
||||
|
||||
Diffusers Support: [LanPaint-Diffusers](https://github.com/charrywhite/LanPaint-diffusers) by [@charrywhite](https://github.com/charrywhite/)
|
||||
|
||||
Benchmark code for paper reproduce: [LanPaintBench](https://github.com/scraed/LanPaintBench).
|
||||
|
||||
## Citation
|
||||
|
||||
@@ -31,14 +37,22 @@ note={}
|
||||
```
|
||||
**🎉 NEW 2026: Join our discord!**
|
||||
|
||||
[Join our Discord](https://discord.gg/aCGZutBV) to share experiences, discuss features, and explore future development.
|
||||
[Join our Discord](https://discord.gg/yN5wYDE6W4) to share experiences, discuss features, and explore future development.
|
||||
|
||||
**🎬 NEW: LanPaint now supports inpainting and outpainting based on Z-Image!**
|
||||
|
||||
`v1.5.0` fixes an important hidden bug that reduced performance and could blur images (especially with `z-image-base`) and also boosts overall LanPaint performance across other models.
|
||||
|
||||
| Original | Masked | Inpainted |
|
||||
|:--------:|:------:|:---------:|
|
||||
|  |  |  |
|
||||
|
||||
**🎬 NEW: LanPaint now supports Z-Image-Base too!**
|
||||
|
||||
| Original | Masked | Inpainted |
|
||||
|:--------:|:------:|:---------:|
|
||||
|  |  |  |
|
||||
|
||||
|
||||
**🎬 NEW: LanPaint now supports video inpainting and outpainting based on Wan 2.2!**
|
||||
|
||||
@@ -67,7 +81,9 @@ 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)
|
||||
@@ -90,7 +106,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, 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, 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.
|
||||

|
||||
- **No Training Needed** – Works out of the box with your existing model.
|
||||
- **Easy to Use** – Same workflow as standard ComfyUI KSampler.
|
||||
@@ -273,6 +289,24 @@ 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 |
|
||||
|:--------:|:------:|:---------:|
|
||||
|  |  |  |
|
||||
|
||||
</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.
|
||||
|
||||
@@ -342,6 +376,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 |
|
||||
|:--------:|:------:|:---------:|
|
||||
|  |  |  |
|
||||
|
||||
</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). If you have quality problem on Comfy 0.11 and 0.12, check [this issue](https://github.com/scraed/LanPaint/issues/80).
|
||||
|
||||
|
||||
### Example Flux: InPaint(LanPaint K Sampler, 5 steps of thinking)
|
||||

|
||||
[View Workflow & Masks](https://github.com/scraed/LanPaint/tree/master/examples/Example_7)
|
||||
@@ -462,6 +512,10 @@ 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/03/02
|
||||
- `v1.5.0`: Fixed a hidden bug that hurt performance and caused image blur (especially on `z-image-base`), and improved overall LanPaint performance on other models too.
|
||||
- 2026/01/30
|
||||
- Add Z-image-base documentation and Example_25 workflow images.
|
||||
- 2025/08/08
|
||||
- Add Qwen image support
|
||||
- 2025/06/21
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
After Width: | Height: | Size: 1.6 MiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 1.8 MiB |
|
After Width: | Height: | Size: 1.5 MiB |
|
After Width: | Height: | Size: 1.7 MiB |
|
After Width: | Height: | Size: 2.2 MiB |
|
After Width: | Height: | Size: 1.8 MiB |
|
After Width: | Height: | Size: 2.8 MiB |
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "LanPaint"
|
||||
version = "1.4.9"
|
||||
version = "1.5.0"
|
||||
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"
|
||||
|
||||
@@ -0,0 +1,337 @@
|
||||
"""
|
||||
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)
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import torch
|
||||
from .utils import *
|
||||
from .utils import StochasticHarmonicOscillator
|
||||
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):
|
||||
def __init__(self, Model, NSteps, Friction, Lambda, Beta, StepSize, IS_FLUX = False, IS_FLOW = False, EarlyStopThreshold = 0.0, EarlyStopPatience = 1, EarlyStopHook = None):
|
||||
self.n_steps = NSteps
|
||||
self.chara_lamb = Lambda
|
||||
self.IS_FLUX = IS_FLUX
|
||||
@@ -13,6 +15,9 @@ 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
|
||||
@@ -32,9 +37,9 @@ class LanPaint():
|
||||
n_steps = self.n_steps
|
||||
return self.LanPaint(x, sigma, latent_mask, current_times, n_steps, model_options, seed, self.IS_FLUX, self.IS_FLOW)
|
||||
def LanPaint(self, x, sigma, latent_mask, current_times, n_steps, model_options, seed, IS_FLUX, IS_FLOW):
|
||||
input_x = x
|
||||
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 +48,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 )
|
||||
@@ -54,21 +57,60 @@ 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:
|
||||
x = x_t * ( 1+self.add_none_dims(VE_Sigma)**2 )**0.5 # switch to variance perserving x_t values
|
||||
############ LanPaint Iterations End ###############
|
||||
# out is x_0
|
||||
|
||||
out, _ = self.inner_model(x, sigma, model_options=model_options, seed=seed)
|
||||
out = out * (1-latent_mask) + self.latent_image * latent_mask
|
||||
|
||||
input_x.copy_(x)
|
||||
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.
|
||||
@@ -89,6 +131,13 @@ 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)
|
||||
@@ -106,11 +155,10 @@ class LanPaint():
|
||||
dt = dtx * (1-mask) + dty * mask
|
||||
Gamma = Gamma_x * (1-mask) + Gamma_y * mask
|
||||
|
||||
|
||||
def Coef_C(x_t):
|
||||
x0 = self.x0_evalutation(x_t, score, sigma, args)
|
||||
x0 = x_t + score(x_t)
|
||||
C = (abt**0.5 * x0 - x_t )/ (1-abt) + A * x_t
|
||||
return C
|
||||
return C, x0
|
||||
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):
|
||||
@@ -119,25 +167,74 @@ 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)
|
||||
|
||||
C = C_new
|
||||
|
||||
return x_t, (v, C)
|
||||
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
|
||||
|
||||
def prepare_step_size(self, current_times, step_size, sigma_x, sigma_y):
|
||||
# -------------------------------------------------------------------------
|
||||
@@ -148,7 +245,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.
|
||||
@@ -173,9 +270,3 @@ 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
|
||||
@@ -1,121 +1,83 @@
|
||||
from contextlib import contextmanager
|
||||
from inspect import cleandoc
|
||||
import inspect
|
||||
import math
|
||||
# import nodes.py
|
||||
import comfy
|
||||
import nodes
|
||||
import latent_preview
|
||||
from functools import partial
|
||||
import torch
|
||||
from comfy.utils import repeat_to_batch_size
|
||||
from comfy.samplers import *
|
||||
from comfy.model_base import ModelType
|
||||
from .utils import *
|
||||
from .lanpaint import LanPaint
|
||||
from comfy.model_base import WAN22
|
||||
import comfyui_version
|
||||
import comfy.nested_tensor
|
||||
|
||||
def _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):
|
||||
|
||||
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 comfyui_version.__version__ < "0.6.0":
|
||||
# input_mask = input_mask.permute(1, 0, 2, 3).unsqueeze(0)
|
||||
# print('Video case - input_mask after reshaping:', input_mask.shape)
|
||||
# # Ensure we have the correct 5D shape: (batch, channels, frames, height, width)
|
||||
# batch_size, channels, frames, height, width = input_mask.shape
|
||||
# print('Video case - dimensions: batch_size={}, channels={}, frames={}, height={}, width={}'.format(batch_size, channels, frames, height, width))
|
||||
# print('Video case - target size:', (target_frames, target_height, target_width))
|
||||
|
||||
# # 3D nearest-exact interpolation: (batch, channels, frames, height, width) -> (batch, channels, target_frames, target_height, target_width)
|
||||
# temp_mask = torch.nn.functional.interpolate(
|
||||
# input_mask,
|
||||
# size=(target_frames, target_height, target_width),
|
||||
# mode=scale_mode,
|
||||
# )
|
||||
|
||||
# # temp_mask is already 5D: (batch, channels, target_frames, target_height, target_width)
|
||||
# mask = temp_mask
|
||||
# print('after mask',mask.shape)
|
||||
# # Handle channel dimension expansion if needed
|
||||
# if mask.shape[1] < output_shape[1]:
|
||||
# mask = mask.repeat(1, output_shape[1], 1, 1, 1)[:, :output_shape[1]]
|
||||
# # Handle batch dimension
|
||||
# mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
if video_inpainting:
|
||||
# 如果是 3D Token 序列 (LTXV 压平后的情况)
|
||||
if input_mask.ndim == 3 and len(output_shape) == 3:
|
||||
mask = torch.nn.functional.interpolate(
|
||||
input_mask,
|
||||
size=output_shape[2],
|
||||
mode=scale_mode
|
||||
)
|
||||
return mask
|
||||
|
||||
# 只有在确认为 5D 视频张量时才执行原有逻辑
|
||||
if input_mask.ndim == 5:
|
||||
target_frames = output_shape[2]
|
||||
target_height, target_width = output_shape[-2:]
|
||||
|
||||
# (这里保留你原有的 permute 和 unsqueeze 逻辑,但要确保它是针对非 5D 输入的补救)
|
||||
if input_mask.ndim < 5:
|
||||
# 假设输入是 (F, C, H, W) -> (1, C, F, H, W)
|
||||
if hasattr(comfyui_version, "__version__") and comfyui_version.__version__ < "0.6.0":
|
||||
input_mask = input_mask.permute(1, 0, 2, 3).unsqueeze(0)
|
||||
|
||||
# 现在可以安全地解包 5D 形状了
|
||||
batch_size, channels, frames, height, width = input_mask.shape
|
||||
mask = torch.nn.functional.interpolate(
|
||||
input_mask,
|
||||
size=(target_frames, target_height, target_width),
|
||||
mode=scale_mode,
|
||||
)
|
||||
|
||||
if mask.shape[1] < output_shape[1]:
|
||||
mask = mask.repeat(1, output_shape[1], 1, 1, 1)[:, :output_shape[1]]
|
||||
mask = repeat_to_batch_size(mask, output_shape[0])
|
||||
return mask
|
||||
if video_inpainting: # Video case: (batch, channels, frames, height, width)
|
||||
target_frames = output_shape[2]
|
||||
target_height, target_width = output_shape[-2:]
|
||||
|
||||
print('Video case - input_mask initial shape:', input_mask.shape)
|
||||
|
||||
# First reshape input_mask to have proper dimensions for video processing
|
||||
# Assume input is (frames, channels, height, width) -> (1, channels, frames, height, width)
|
||||
## if comfy version < 0.6.0
|
||||
if 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])
|
||||
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):
|
||||
@@ -146,9 +108,9 @@ class CFGGuider_LanPaint:
|
||||
if isinstance(self.inner_model, WAN22):
|
||||
print("WAN22 detected")
|
||||
self.inner_model.extra_conds = super(WAN22, self.inner_model).extra_conds
|
||||
|
||||
if denoise_mask is not None:
|
||||
video_inpainting = self.model_options.get("video_inpainting", False)
|
||||
print('denoise_mask',denoise_mask.shape,type(denoise_mask))
|
||||
denoise_mask = prepare_mask(denoise_mask, noise.shape, device, video_inpainting)
|
||||
|
||||
noise = noise.to(device)
|
||||
@@ -196,6 +158,8 @@ class KSamplerX0Inpaint:
|
||||
abt = (1 - Flow_t)**2 / ((1 - Flow_t)**2 + Flow_t**2 )
|
||||
VE_Sigma = Flow_t / (1 - Flow_t)
|
||||
#print("t", torch.mean( sigma ).item(), "VE_Sigma", torch.mean( VE_Sigma ).item())
|
||||
|
||||
|
||||
else:
|
||||
VE_Sigma = sigma
|
||||
abt = 1/( 1+VE_Sigma**2 )
|
||||
@@ -205,31 +169,6 @@ class KSamplerX0Inpaint:
|
||||
if "denoise_mask_function" in model_options:
|
||||
denoise_mask = model_options["denoise_mask_function"](sigma, denoise_mask, extra_options={"model": self.inner_model, "sigmas": self.sigmas})
|
||||
|
||||
if isinstance(denoise_mask, comfy.nested_tensor.NestedTensor):
|
||||
masks = denoise_mask.unbind()
|
||||
xs = x.unbind()
|
||||
latent_imgs = self.latent_image.unbind()
|
||||
noises = self.noise.unbind()
|
||||
|
||||
outs = []
|
||||
# 针对 LTXV,通常 i=0 是视频,i=1 是音频
|
||||
for i in range(len(xs)):
|
||||
m = (masks[i] > 0.5).float()
|
||||
lm = 1 - m
|
||||
# 这里的 PaintMethod 通常只支持普通 Tensor,所以我们分块处理
|
||||
# 注意:如果音频部分不需要 Inpaint,可以增加判断
|
||||
current_times = (VE_Sigma, abt, Flow_t)
|
||||
|
||||
# 只有视频部分 (i=0) 应用 LanPaint 逻辑,音频部分通常直接 pass 或原样返回
|
||||
if i == 0:
|
||||
out_part = self.PaintMethod(xs[i], latent_imgs[i], noises[i], sigma, lm, current_times, model_options, seed)
|
||||
else:
|
||||
# 音频部分如果没有对应的 Inpaint 逻辑,通常直接调用 inner_model
|
||||
out_part, _ = self.inner_model(xs[i], sigma, model_options=model_options, seed=seed)
|
||||
outs.append(out_part)
|
||||
|
||||
return comfy.nested_tensor.NestedTensor(tuple(outs))
|
||||
|
||||
denoise_mask = (denoise_mask > 0.5).float()
|
||||
|
||||
latent_mask = 1 - denoise_mask
|
||||
@@ -244,7 +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)
|
||||
@@ -255,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
|
||||
@@ -264,7 +203,6 @@ class KSAMPLER(comfy.samplers.KSAMPLER):
|
||||
#noise here is a randn noise from comfy.sample.prepare_noise
|
||||
#latent_image is the latent image as input of the KSampler node. For inpainting, it is the masked latent image. Otherwise it is zero tensor.
|
||||
extra_args["denoise_mask"] = denoise_mask
|
||||
print("LanPaint KSampler start sampler_function",denoise_mask.shape if denoise_mask is not None else None)
|
||||
model_k = KSamplerX0Inpaint(model_wrap, sigmas)
|
||||
model_k.latent_image = latent_image
|
||||
if self.inpaint_options.get("random", False): #TODO: Should this be the default?
|
||||
@@ -289,7 +227,10 @@ 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)
|
||||
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))
|
||||
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.
|
||||
@@ -391,17 +332,19 @@ 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:
|
||||
@@ -430,6 +373,8 @@ 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)"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -438,7 +383,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"):
|
||||
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):
|
||||
force_full_denoise = True
|
||||
if return_with_leftover_noise == "enable":
|
||||
force_full_denoise = False
|
||||
@@ -451,11 +396,13 @@ 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:
|
||||
@@ -507,7 +454,7 @@ class MaskBlend:
|
||||
kernel = self.gaussian_kernel(blend_overlap)
|
||||
kernel = kernel.to(image1.device)
|
||||
kernel = kernel[None, None, ...]
|
||||
|
||||
|
||||
mask = torch.nn.functional.conv2d(mask[:,None,:,:], kernel, padding=blend_overlap//2)[:,0,:,:]
|
||||
|
||||
|
||||
@@ -529,77 +476,6 @@ class MaskBlend:
|
||||
|
||||
return kernel
|
||||
|
||||
class MaskBlendAlpha:
|
||||
"""
|
||||
Create an RGBA image by writing the mask into the PNG alpha channel.
|
||||
|
||||
Requirement:
|
||||
- inpaint region: alpha = 0 (transparent)
|
||||
- other region: alpha = 1 (opaque)
|
||||
|
||||
This node writes the mask into the PNG alpha channel.
|
||||
Current default behavior matches the previous `invert_mask=True` behavior:
|
||||
alpha = mask.
|
||||
"""
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"tooltip": "VAE-decoded image (RGB)."}),
|
||||
"mask": ("MASK", {"tooltip": "Mask used as alpha channel (alpha = mask)."}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "to_rgba"
|
||||
CATEGORY = "image/postprocessing"
|
||||
|
||||
def to_rgba(self, image: torch.Tensor, mask: torch.Tensor):
|
||||
"""
|
||||
image: [B,H,W,3] float in [0,1]
|
||||
mask: [B,H,W] (or [H,W]) float in [0,1] used as alpha
|
||||
returns RGBA image: [B,H,W,4] float in [0,1]
|
||||
"""
|
||||
if image.ndim != 4 or image.shape[-1] != 3:
|
||||
raise ValueError(f"Expected IMAGE tensor [B,H,W,3], got {tuple(image.shape)}")
|
||||
|
||||
# Normalize mask shape to [B,H,W]
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
elif mask.ndim == 3:
|
||||
pass
|
||||
else:
|
||||
# Some pipelines may carry mask as [B,1,H,W]
|
||||
if mask.ndim == 4 and mask.shape[1] == 1:
|
||||
mask = mask[:, 0, :, :]
|
||||
else:
|
||||
raise ValueError(f"Expected MASK tensor [B,H,W] or [H,W], got {tuple(mask.shape)}")
|
||||
|
||||
b, h, w, _ = image.shape
|
||||
|
||||
# Batch align
|
||||
if mask.shape[0] != b:
|
||||
if mask.shape[0] == 1:
|
||||
mask = mask.repeat(b, 1, 1)
|
||||
else:
|
||||
raise ValueError(f"Batch mismatch: image batch={b}, mask batch={mask.shape[0]}")
|
||||
|
||||
# Spatial align (resize mask to image resolution if needed)
|
||||
if mask.shape[1] != h or mask.shape[2] != w:
|
||||
mask_4d = mask.unsqueeze(1) # [B,1,H,W]
|
||||
mask_4d = torch.nn.functional.interpolate(mask_4d, size=(h, w), mode="nearest")
|
||||
mask = mask_4d[:, 0, :, :]
|
||||
|
||||
mask = mask.float().clamp(0.0, 1.0)
|
||||
|
||||
# Default behavior (matches previous invert_mask=True path):
|
||||
# alpha = mask
|
||||
rgba = torch.cat([image, mask.unsqueeze(-1)], dim=-1)
|
||||
return (rgba,)
|
||||
|
||||
class Noise_EmptyNoise:
|
||||
def generate_noise(self, latent):
|
||||
return torch.zeros_like(latent["samples"])
|
||||
@@ -628,7 +504,6 @@ class LanPaint_SamplerCustom:
|
||||
"LanPaint_NumSteps": ("INT", {"default": 5, "min": 0, "max": 100, "tooltip": "Number of steps for Langevin dynamics, representing turns of thinking per step."}),
|
||||
"LanPaint_PromptMode": (["Image First", "Prompt First"], {"tooltip": "Image First: prioritizes image quality; Prompt First: prioritizes prompt adherence."}),
|
||||
"LanPaint_Info": ("STRING", {"default": "LanPaint Custom Sampler. For more info, visit https://github.com/scraed/LanPaint. If you find it useful, please give a star ⭐️!", "multiline": True}),
|
||||
"Inpainting_mode": (["🖼️ Image Inpainting", "🎬 Video Inpainting"], {"default": "🖼️ Image Inpainting", "tooltip": "Choose Image mode for photos or Video mode for video frames with temporal consistency"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -637,21 +512,19 @@ class LanPaint_SamplerCustom:
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, model, sampler, sigmas, add_noise, noise_seed, cfg, positive, negative, latent_image, LanPaint_NumSteps, LanPaint_PromptMode, LanPaint_Info="",Inpainting_mode="🖼️ Image Inpainting"):
|
||||
def sample(self, model, sampler, sigmas, add_noise, noise_seed, cfg, positive, negative, latent_image, LanPaint_NumSteps, LanPaint_PromptMode, LanPaint_Info=""):
|
||||
model.LanPaint_StepSize = 0.2
|
||||
model.LanPaint_Lambda = 16.0
|
||||
model.LanPaint_Beta = 1.
|
||||
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"]
|
||||
@@ -681,7 +554,7 @@ class LanPaint_SamplerCustom:
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class LanPaint_SamplerCustomAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -699,7 +572,8 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
"LanPaint_PromptMode": (["Image First", "Prompt First"], {"tooltip": "Image First: prioritizes image quality; Prompt First: prioritizes prompt adherence."}),
|
||||
"LanPaint_EarlyStop": ("INT", {"default": 1, "min": 0, "max": 10000, "tooltip": "Steps to stop LanPaint early, preventing irregular patterns."}),
|
||||
"LanPaint_Info": ("STRING", {"default": "LanPaint Custom Sampler Adv. For more info, visit https://github.com/scraed/LanPaint. If you find it useful, please give a star ⭐️!", "multiline": True}),
|
||||
"Inpainting_mode": (["🖼️ Image Inpainting", "🎬 Video Inpainting"], {"default": "🖼️ Image Inpainting", "tooltip": "Choose Image mode for photos or Video mode for video frames with temporal consistency"}),
|
||||
"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)"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -710,7 +584,7 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
|
||||
CATEGORY = "sampling/custom_sampling"
|
||||
|
||||
def sample(self, noise, guider, sampler, sigmas, latent_image, LanPaint_NumSteps, LanPaint_Lambda, LanPaint_StepSize, LanPaint_Beta, LanPaint_Friction, LanPaint_PromptMode, LanPaint_EarlyStop, LanPaint_Info="",Inpainting_mode="🖼️ Image Inpainting"):
|
||||
def sample(self, noise, guider, sampler, sigmas, latent_image, LanPaint_NumSteps, LanPaint_Lambda, LanPaint_StepSize, LanPaint_Beta, LanPaint_Friction, LanPaint_PromptMode, LanPaint_EarlyStop, LanPaint_Info="", LanPaint_InnerThreshold=0.0, LanPaint_InnerPatience=1):
|
||||
model = guider.model_patcher
|
||||
model.LanPaint_StepSize = LanPaint_StepSize
|
||||
model.LanPaint_Lambda = LanPaint_Lambda
|
||||
@@ -718,29 +592,22 @@ 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)
|
||||
@@ -756,7 +623,6 @@ class LanPaint_SamplerCustomAdvanced:
|
||||
out_denoised["samples"] = guider.model_patcher.model.process_latent_out(x0_output["x0"].cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
# print('output',out.keys(),out["samples"].shape,out['noise_mask'].shape)
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
@@ -768,7 +634,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"LanPaint_SamplerCustom" : LanPaint_SamplerCustom,
|
||||
"LanPaint_SamplerCustomAdvanced" : LanPaint_SamplerCustomAdvanced,
|
||||
"LanPaint_MaskBlend": MaskBlend,
|
||||
"LanPaint_MaskBlendAlpha": MaskBlendAlpha,
|
||||
# "LanPaint_UpSale_LatentNoiseMask": LanPaint_UpSale_LatentNoiseMask,
|
||||
}
|
||||
|
||||
@@ -779,6 +644,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LanPaint_SamplerCustom" : "LanPaint Sampler Custom",
|
||||
"LanPaint_SamplerCustomAdvanced" : "LanPaint Sampler Custom (Advanced)",
|
||||
"LanPaint_MaskBlend": "LanPaint Mask Blend",
|
||||
"LanPaint_MaskBlendAlpha": "MaskBlend (alpha)",
|
||||
# "LanPaint_UpSale_LatentNoiseMask": "LanPaint UpSale Latent Noise Mask"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
from typing import NamedTuple, Optional
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class LangevinState(NamedTuple):
|
||||
v: Optional[torch.Tensor]
|
||||
C: Optional[torch.Tensor]
|
||||
x0: Optional[torch.Tensor]
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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 = ..
|
||||
|
||||
@@ -1,21 +1,13 @@
|
||||
#!/usr/bin/env python
|
||||
"""Basic import tests for LanPaint.
|
||||
|
||||
"""Tests for `LanPaint` package."""
|
||||
The ComfyUI runtime dependencies (e.g. `comfy`) are intentionally optional for unit tests.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from src.LanPaint.nodes import Example
|
||||
|
||||
@pytest.fixture
|
||||
def example_node():
|
||||
"""Fixture to create an Example node instance."""
|
||||
return Example()
|
||||
def test_package_imports_without_comfy() -> None:
|
||||
import LanPaint
|
||||
|
||||
def test_example_node_initialization(example_node):
|
||||
"""Test that the node can be instantiated."""
|
||||
assert isinstance(example_node, Example)
|
||||
|
||||
def test_return_types():
|
||||
"""Test the node's metadata."""
|
||||
assert Example.RETURN_TYPES == ("IMAGE",)
|
||||
assert Example.FUNCTION == "test"
|
||||
assert Example.CATEGORY == "Example"
|
||||
assert isinstance(LanPaint.NODE_CLASS_MAPPINGS, dict)
|
||||
assert isinstance(LanPaint.NODE_DISPLAY_NAME_MAPPINGS, dict)
|
||||
assert "LanPaint_KSampler" in LanPaint.NODE_CLASS_MAPPINGS
|
||||
assert LanPaint.WEB_DIRECTORY == "./web"
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
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
|
||||
@@ -0,0 +1,74 @@
|
||||
import importlib
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
|
||||
def _repeat_to_batch_size(tensor: torch.Tensor, batch_size: int) -> torch.Tensor:
|
||||
if tensor.shape[0] == batch_size:
|
||||
return tensor
|
||||
if tensor.shape[0] == 1:
|
||||
return tensor.repeat((batch_size,) + (1,) * (tensor.ndim - 1))
|
||||
repeats = (batch_size + tensor.shape[0] - 1) // tensor.shape[0]
|
||||
return tensor.repeat((repeats,) + (1,) * (tensor.ndim - 1))[:batch_size]
|
||||
|
||||
|
||||
def _import_nodes(monkeypatch, comfyui_version: str):
|
||||
comfy_mod = types.ModuleType("comfy")
|
||||
comfy_mod.__path__ = []
|
||||
|
||||
comfy_utils_mod = types.ModuleType("comfy.utils")
|
||||
comfy_utils_mod.repeat_to_batch_size = _repeat_to_batch_size
|
||||
|
||||
comfy_samplers_mod = types.ModuleType("comfy.samplers")
|
||||
class DummyKSAMPLER: ...
|
||||
comfy_samplers_mod.KSAMPLER = DummyKSAMPLER
|
||||
|
||||
comfy_model_base_mod = types.ModuleType("comfy.model_base")
|
||||
class ModelType:
|
||||
FLUX = "FLUX"
|
||||
FLOW = "FLOW"
|
||||
|
||||
class WAN22: ...
|
||||
comfy_model_base_mod.ModelType = ModelType
|
||||
comfy_model_base_mod.WAN22 = WAN22
|
||||
|
||||
comfyui_version_mod = types.ModuleType("comfyui_version")
|
||||
comfyui_version_mod.__version__ = comfyui_version
|
||||
|
||||
comfy_mod.utils = comfy_utils_mod
|
||||
comfy_mod.samplers = comfy_samplers_mod
|
||||
comfy_mod.model_base = comfy_model_base_mod
|
||||
|
||||
monkeypatch.setitem(sys.modules, "comfy", comfy_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.utils", comfy_utils_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.samplers", comfy_samplers_mod)
|
||||
monkeypatch.setitem(sys.modules, "comfy.model_base", comfy_model_base_mod)
|
||||
monkeypatch.setitem(sys.modules, "nodes", types.ModuleType("nodes"))
|
||||
monkeypatch.setitem(sys.modules, "latent_preview", types.ModuleType("latent_preview"))
|
||||
monkeypatch.setitem(sys.modules, "comfyui_version", comfyui_version_mod)
|
||||
|
||||
sys.modules.pop("src.LanPaint.nodes", None)
|
||||
return importlib.import_module("src.LanPaint.nodes")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("comfyui_version", ["0.5.0", "0.6.0"])
|
||||
def test_reshape_mask_accepts_bhw_and_5d_output_shape(monkeypatch, comfyui_version: str) -> None:
|
||||
lanpaint_nodes = _import_nodes(monkeypatch, comfyui_version)
|
||||
input_mask = torch.zeros((1, 4, 4))
|
||||
output_shape = (1, 16, 1, 8, 8)
|
||||
|
||||
out = lanpaint_nodes.reshape_mask(input_mask, output_shape, video_inpainting=False)
|
||||
assert tuple(out.shape) == output_shape
|
||||
|
||||
|
||||
def test_prepare_mask_accepts_hw_and_moves_device(monkeypatch) -> None:
|
||||
lanpaint_nodes = _import_nodes(monkeypatch, "0.5.0")
|
||||
input_mask = torch.zeros((4, 4))
|
||||
output_shape = (2, 3, 8, 8)
|
||||
|
||||
out = lanpaint_nodes.prepare_mask(input_mask, output_shape, device=torch.device("cpu"), video_inpainting=False)
|
||||
assert tuple(out.shape) == output_shape
|
||||
assert out.device.type == "cpu"
|
||||
@@ -0,0 +1,45 @@
|
||||
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()
|
||||
|
||||