Compare commits

..
71 Commits
Author SHA1 Message Date
scraed 24051fb151 update README 2026-03-02 22:28:29 +08:00
scraed a76a5da0c5 bump version 2026-03-02 22:25:37 +08:00
scraed 21a6b155e8 fix the langevin output x has not been used issue 2026-03-02 14:16:32 +08:00
scraed 8bc20d5e00 Update README with additional resources and structure
Reorganized README content for clarity and added links to Diffusers
2026-02-16 01:03:30 +08:00
scraed 84b63adb11 replace flux example mask for bug fixing 2026-02-07 12:15:29 +08:00
scraed 725fb2ba37 Replace Klein example image 2026-02-07 11:56:03 +08:00
scraed 4f19674299 Update README.md 2026-02-04 14:40:08 +08:00
scraed 558084416d version 1.4.13 2026-02-02 23:49:35 +08:00
scraed a984dcb118 Fix flux klein example workflow 2026-02-02 23:48:33 +08:00
scraed 3f9193997a Update README.md 2026-02-02 20:36:55 +08:00
scraed 7372428114 Update README.md 2026-02-02 20:18:57 +08:00
scraed ac6eece747 Update Discord link in README.md 2026-02-02 20:15:45 +08:00
scraed 95b65ed1b1 Merge pull request #74 from godnight10061/fix/semantic-stop-tail
Add semantic inner-loop early stop mechanism
2026-01-30 11:44:51 +08:00
scraed 30cfd2f2e4 add contribution info 2026-01-30 11:40:02 +08:00
scraed 2e3a566bb4 Keep a minimal set of test 2026-01-30 11:23:51 +08:00
scraed 8fe1a9d312 Move parameters to the last for backward compatibility 2026-01-30 10:45:40 +08:00
scraed 4fcbfc7b5b Merge branch 'master' into pr/74 2026-01-30 10:20:59 +08:00
scraed 651ddc2225 fix typo 2026-01-30 00:48:22 +08:00
scraed bd4606f27a bump version to 1.4.12 2026-01-30 00:46:18 +08:00
scraed 8a72294f27 Z Image Base support 2026-01-30 00:45:20 +08:00
godnight10061 9ddd5f613e Map legacy min_steps into patience 2026-01-26 21:40:36 +08:00
godnight10061 39c90bbf62 Revert benchmark workflow 2026-01-26 21:16:05 +08:00
godnight10061 6ad4a18537 Refactor earlystop to reduce duplication 2026-01-26 18:09:21 +08:00
godnight10061 5e7d472738 Use LangevinState in semantic stop test mocks 2026-01-26 14:39:47 +08:00
godnight10061 d73f3db265 Document distance_fn hook 2026-01-26 14:02:07 +08:00
godnight10061 cd06069b52 Avoid early-stop when no inpaint region 2026-01-26 10:50:39 +08:00
godnight10061 114d091965 Simplify LanPaint inner-loop glue 2026-01-26 02:50:36 +08:00
godnight10061 06da285a09 Fix LangevinState backward compat 2026-01-25 23:42:04 +08:00
godnight10061 8f52a1d378 Use NamedTuple for Langevin state 2026-01-25 20:32:51 +08:00
godnight10061 c7962bf784 Refactor semantic early-stop metric 2026-01-25 10:57:29 +08:00
godnight10061 a895d2a758 Tighten semantic early-stop robustness 2026-01-24 12:43:23 +08:00
godnight10061 34ded92d8d Simplify semantic early-stop knobs 2026-01-24 10:19:37 +08:00
scraed 1b87239553 Move Coef_C definition above advance_time
Reordered the Coef_C function definition to appear before advance_time for improved code organization and readability.
2026-01-24 10:19:37 +08:00
scraed cbc9f3abfd Refactor early-stop logic into LanPaintEarlyStopper class
Moved the early-stop logic from lanpaint.py into a new src/LanPaint/earlystop.py module as the LanPaintEarlyStopper class. This improves code organization and maintainability by encapsulating early-stop behavior, reducing complexity in the main LanPaint class.
2026-01-24 10:19:37 +08:00
godnight10061 7d46147887 Guard torch.nn.functional import 2026-01-24 10:19:37 +08:00
godnight10061 c5fe6351ea Harden semantic-stop distance hook 2026-01-24 10:19:37 +08:00
godnight10061 ac7a5a1f2f Fix semantic-stop distance_fn dispatch 2026-01-24 10:19:37 +08:00
godnight10061 73ab537e86 Make CI-only imports lighter 2026-01-24 10:19:37 +08:00
godnight10061 6fd5f60383 Adjust SHO regression test for x0 return 2026-01-24 10:19:36 +08:00
godnight10061 9c9b5555fc Tighten semantic stop in mid-noise steps 2026-01-24 10:19:36 +08:00
godnight10061 450a25e5e8 Refine semantic inner stop 2026-01-24 10:19:36 +08:00
godnight10061 e7080f9574 Test ring changes prevent premature freeze 2026-01-24 10:19:36 +08:00
godnight10061 2086f00603 Skip score model after semantic convergence
# Conflicts:
#	src/LanPaint/lanpaint.py
2026-01-24 10:19:36 +08:00
godnight10061 896729ccd4 Fix semantic early-stop distance 2026-01-24 10:19:36 +08:00
godnight10061 207860edc3 Avoid semantic early-stop at low sigma 2026-01-24 10:19:36 +08:00
godnight10061 4074473b90 Add per-case quality gates to benchmark 2026-01-24 10:19:36 +08:00
godnight10061 11c1490140 Make benchmark workflow test semantic early stop 2026-01-24 10:19:36 +08:00
godnight10061 7c79712fc8 Make LPIPS benchmark long-running e2e 2026-01-24 10:19:36 +08:00
godnight10061 fb35a1d032 Add ImageNet LPIPS benchmark workflow 2026-01-24 10:19:36 +08:00
godnight10061 060fcc4475 Add semantic early-stop for Langevin iterations 2026-01-24 10:19:36 +08:00
scraed aedb908a3f Add version tuple check for ComfyUI version handling
Introduced a helper function to compare ComfyUI versions as tuples instead of strings, improving reliability of version-dependent logic in reshape_mask. Updated all version checks to use the new tuple-based comparison.
2026-01-22 16:22:41 +08:00
scraed f9718ea3e1 Update Flux 2 klein image example link
Corrected the anchor link for the 'Flux 2 klein' image example to reference '2 steps of thinking' instead of '5 steps of thinking'.
2026-01-18 00:55:57 +08:00
scraed 0bafe2117a Merge branch 'master' of https://github.com/scraed/LanPaint 2026-01-18 00:46:59 +08:00
scraed 2c1a23777d Update example steps in README
Changed the number of steps in the Flux 2 klein InPaint example from 5 to 2 to reflect the correct workflow.
2026-01-18 00:46:56 +08:00
scraed dde4c82463 Bump version from 1.4.10 to 1.4.11 2026-01-18 00:46:14 +08:00
scraed 72cb484c40 Add Flux 2 Klein inpainting example and workflow
Added a new image inpainting example for Flux 2 Klein, including workflow JSON, sample images (original, masked, inpainted), and updated README with documentation and links. This provides users with a reference workflow and visual results for the Flux 2 Klein model.
2026-01-18 00:39:22 +08:00
scraed 0d15d15cf3 Update Discord link in README.md 2026-01-13 23:34:02 +08:00
scraed 1393a46d67 Bump version from 1.4.9 to 1.4.10 2026-01-13 22:17:42 +08:00
scraed 27421363bf Merge pull request #71 from godnight10061/fix/issue69-nan-multivariatenormal
Avoid MultivariateNormal crash on non-finite dynamics
2026-01-13 22:16:34 +08:00
scraed 01400a541d Update test_sho_regression.py 2026-01-13 22:13:30 +08:00
scraed c981354387 Update test_sho_regression.py 2026-01-13 21:46:56 +08:00
scraed c92c8bfb1f Update test_sho_regression.py 2026-01-13 21:38:49 +08:00
scraed 9334025801 Update test_sho_regression.py 2026-01-13 21:31:15 +08:00
scraed 144893ecd2 fall back to overdamped update if nan appears 2026-01-13 18:41:27 +08:00
scraed bce8505ce2 Revert "Avoid MultivariateNormal crash on non-finite dynamics"
This reverts commit 9432a34c38.
2026-01-13 18:29:06 +08:00
godnight10061 0992657e54 Fix CI workflows for torch and default branch 2026-01-13 09:25:41 +08:00
godnight10061 2fa036992b Add unit regression for NaN oscillator mean 2026-01-13 08:56:46 +08:00
godnight10061 9432a34c38 Avoid MultivariateNormal crash on non-finite dynamics 2026-01-13 01:39:48 +08:00
scraed d4e8ee28fb Merge pull request #70 from godnight10061/fix/reshape-mask-3d
Thanks for the contribution! Tested and merged.
2026-01-11 18:22:14 +08:00
godnight10061 34b8e83831 Make CI tooling importable 2026-01-11 11:07:41 +08:00
godnight10061 e29d480a1a Fix reshape_mask for 3D masks 2026-01-10 14:46:14 +08:00
25 changed files with 5166 additions and 298 deletions
+1
View File
@@ -26,6 +26,7 @@ jobs:
run: |
python -m pip install --upgrade pip
pip install .[dev]
pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
- name: Run Linting
run: |
ruff check .
+2
View File
@@ -11,3 +11,5 @@ jobs:
runs-on: ubuntu-latest
steps:
- uses: comfy-org/node-diff@main
with:
base_ref: ${{ github.event.repository.default_branch }}
+58 -4
View File
@@ -7,13 +7,19 @@
[![Hugging Face](https://img.shields.io/badge/Hugging%20Face-yellow?logo=huggingface&logoColor=white)](https://huggingface.co/charrywhite/LanPaint)
[![Blog](https://img.shields.io/badge/📝-Blog-9cf)](https://scraed.github.io/scraedBlog/)
[![GitHub stars](https://img.shields.io/github/stars/scraed/LanPaint)](https://github.com/scraed/LanPaint/stargazers)
[![Discord](https://img.shields.io/badge/Discord-5865F2?style=for-the-badge&logo=discord&logoColor=white)](https://discord.gg/aCGZutBV)
[![Discord](https://img.shields.io/badge/Discord-5865F2?style=for-the-badge&logo=discord&logoColor=white)](https://discord.gg/yN5wYDE6W4)
</div>
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 |
|:--------:|:------:|:---------:|
| ![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!**
@@ -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.
![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.
@@ -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 |
|:--------:|:------:|:---------:|
| ![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.
@@ -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 |
|:--------:|:------:|:---------:|
| ![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). 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)
![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)
@@ -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
+80 -2
View File
@@ -10,7 +10,85 @@ __author__ = """LanPaint"""
__email__ = "czhengac@connect.ust.hk"
__version__ = "0.0.1"
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
def _install_lightweight_runtime_stubs() -> None:
"""Install lightweight stubs so tooling can import this package without ComfyUI.
This is used by CI tooling (e.g., comfy-org/node-diff) that imports NODE_CLASS_MAPPINGS
in an environment where ComfyUI isn't installed.
"""
import sys
import types
# `src/LanPaint/nodes.py` uses `torch.Tensor` in type annotations.
try:
import torch # noqa: F401
except ModuleNotFoundError:
torch_mod = types.ModuleType("torch")
class Tensor: # noqa: N801 (match torch naming)
pass
torch_mod.Tensor = Tensor
torch_mod.nn = types.SimpleNamespace(functional=types.SimpleNamespace())
sys.modules["torch"] = torch_mod
if "comfyui_version" not in sys.modules:
comfyui_version_mod = types.ModuleType("comfyui_version")
comfyui_version_mod.__version__ = "0.0.0"
sys.modules["comfyui_version"] = comfyui_version_mod
sys.modules.setdefault("nodes", types.ModuleType("nodes"))
sys.modules.setdefault("latent_preview", types.ModuleType("latent_preview"))
if "comfy" not in sys.modules:
comfy_mod = types.ModuleType("comfy")
comfy_mod.__path__ = []
comfy_utils_mod = types.ModuleType("comfy.utils")
def repeat_to_batch_size(tensor, batch_size): # type: ignore[no-untyped-def]
if getattr(tensor, "shape", ())[0] == batch_size:
return tensor
return tensor
comfy_utils_mod.repeat_to_batch_size = repeat_to_batch_size
comfy_samplers_mod = types.ModuleType("comfy.samplers")
class DummyKSAMPLER: # noqa: N801 (match ComfyUI naming)
pass
comfy_samplers_mod.KSAMPLER = DummyKSAMPLER
comfy_model_base_mod = types.ModuleType("comfy.model_base")
class ModelType: # noqa: N801 (match ComfyUI naming)
FLUX = "FLUX"
FLOW = "FLOW"
class WAN22: # noqa: N801 (match ComfyUI naming)
pass
comfy_model_base_mod.ModelType = ModelType
comfy_model_base_mod.WAN22 = WAN22
comfy_mod.utils = comfy_utils_mod
comfy_mod.samplers = comfy_samplers_mod
comfy_mod.model_base = comfy_model_base_mod
sys.modules["comfy"] = comfy_mod
sys.modules["comfy.utils"] = comfy_utils_mod
sys.modules["comfy.samplers"] = comfy_samplers_mod
sys.modules["comfy.model_base"] = comfy_model_base_mod
try:
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
except ModuleNotFoundError:
_install_lightweight_runtime_stubs()
from .src.LanPaint.nodes import NODE_CLASS_MAPPINGS
from .src.LanPaint.nodes import NODE_DISPLAY_NAME_MAPPINGS
WEB_DIRECTORY = "./web"
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.6 MiB

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

After

Width:  |  Height:  |  Size: 1.1 MiB

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

After

Width:  |  Height:  |  Size: 1.8 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.5 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.8 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.8 MiB

+4 -1
View File
@@ -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"
+337
View File
@@ -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)
+122 -31
View File
@@ -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
+83 -219
View File
@@ -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"
}
+10
View File
@@ -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]
+20 -21
View File
@@ -28,11 +28,11 @@ def expm1mxmhx2_x3(x):
def exp_1mcosh_GD(gamma_t, delta):
"""
Compute e^(-Γt) * (1 - cosh(Γt√Δ))/ ( (Γt)**2 Δ )
Parameters:
gamma_t: Γ*t term (could be a scalar or tensor)
delta: Δ term (could be a scalar or tensor)
Returns:
Result of the computation with numerical stability handling
"""
@@ -68,7 +68,6 @@ def exp_sinh_GsqrtD(gamma_t, delta):
sqrt_abs_delta = torch.sqrt(torch.abs(delta))
gamma_t_sqrt_delta = gamma_t * sqrt_abs_delta
numerator_pos = (torch.exp(gamma_t * (sqrt_abs_delta - 1)) - torch.exp(gamma_t * (-sqrt_abs_delta - 1))) / 2
denominator_pos = gamma_t_sqrt_delta
result_pos = numerator_pos / gamma_t_sqrt_delta
result_pos = torch.where(torch.isfinite(result_pos), result_pos, torch.zeros_like(result_pos))
@@ -117,15 +116,15 @@ def zeta1(gamma_t, delta):
exp_cosh_term = exp_cosh(half_gamma_t, delta)
exp_sinh_term = exp_sinh_sqrtD(half_gamma_t, delta)
# Main computation
numerator = 1 - (exp_cosh_term + exp_sinh_term)
denominator = gamma_t * (1 - delta) / 4
result = 1 - numerator / denominator
# Handle numerical instability
result = torch.where(torch.isfinite(result), result, torch.zeros_like(result))
# Taylor expansion for small x (similar to your epxm1Dx approach)
mask = torch.abs(denominator) < 5e-3
term1 = epxm1_x(-gamma_t)
@@ -133,17 +132,17 @@ def zeta1(gamma_t, delta):
term3 = expm1mxmhx2_x3(-gamma_t)
taylor = term1 + (1/2.+ term1-3*term2)*denominator + (-1/6. + term1/2 - 4 * term2 + 10 * term3) * denominator**2
result = torch.where(mask, taylor, result)
return result
def exp_cosh_minus_terms(gamma_t, delta):
"""
Compute E^(-tΓ) * (Cosh[tΓ] - 1 - (Cosh[tΓ√Δ] - 1)/Δ) / (tΓ(1 - Δ))
Parameters:
gamma_t: Γ*t term (could be a scalar or tensor)
delta: Δ term (could be a scalar or tensor)
Returns:
Result of the computation with numerical stability handling
"""
@@ -151,17 +150,17 @@ def exp_cosh_minus_terms(gamma_t, delta):
# Compute individual terms
exp_cosh_term = exp_cosh(gamma_t, gamma_t**0) - exp_term # E^(-tΓ) (Cosh[tΓ] - 1) term
exp_cosh_delta_term = - gamma_t**2 * exp_1mcosh_GD(gamma_t, delta) # E^(-tΓ) (Cosh[tΓ√Δ] - 1)/Δ term
#exp_1mcosh_GD e^(-Γt) * (1 - cosh(Γt√Δ))/ ( (Γt)**2 Δ )
# Main computation
numerator = exp_cosh_term - exp_cosh_delta_term
denominator = gamma_t * (1 - delta)
result = numerator / denominator
# Handle numerical instability
result = torch.where(torch.isfinite(result), result, torch.zeros_like(result))
# Taylor expansion for small gamma_t and delta near 1
mask = (torch.abs(denominator) < 1e-1)
exp_1mcosh_GD_term = exp_1mcosh_GD(gamma_t, delta**0)
@@ -170,7 +169,7 @@ def exp_cosh_minus_terms(gamma_t, delta):
- denominator / 4 * ( 0.5 * exp_cosh(gamma_t, delta**0) - 4 * exp_1mcosh_GD_term - 5 /2 * exp_sinh_GsqrtD(gamma_t, delta**0) )
)
result = torch.where(mask, taylor, result)
return result
@@ -185,7 +184,7 @@ def sig11(gamma_t, delta):
def Zcoefs(gamma_t, delta):
Zeta1 = zeta1(gamma_t, delta)
Zeta2 = zeta2(gamma_t, delta)
sq_total = 1 - Zeta1 + gamma_t * (delta - 1) * (Zeta1 - 1)**2 / 8
amplitude = torch.sqrt(sq_total)
Zcoef1 = ( gamma_t**0.5 * Zeta2 / 2 **0.5 ) / amplitude
@@ -208,7 +207,7 @@ class StochasticHarmonicOscillator:
dq(t) = -Γ A y(t) dt + Γ C dt + Γ D dw(t) - Γ q(t) dt
Also define v(t) = q(t) / √Γ, which is numerically more stable.
Where:
y(t) - Position variable
q(t) - Velocity variable
@@ -239,7 +238,7 @@ class StochasticHarmonicOscillator:
Returns:
tuple: (y(t), v(t))
"""
dummyzero = y0.new_zeros(1) # convert scalar to tensor with same device and dtype as y0
Delta = self.Delta + dummyzero
Gamma_hat = self.Gamma * t + dummyzero
@@ -254,12 +253,12 @@ class StochasticHarmonicOscillator:
if v0 is None:
v0 = torch.randn_like(y0) * D / 2 ** 0.5
#v0 = (C - A * y0)/Gamma**0.5
# Calculate mean position and velocity
term1 = (1 - zeta_1) * (C * t - A * t * y0) + zeta_2 * (Gamma ** 0.5) * v0 * t
y_mean = term1 + y0
v_mean = (1 - EE)*(C - A * y0) / (Gamma ** 0.5) + (EE - A * t * (1 - zeta_1)) * v0
cov_yy = D**2 * t * self.sig22(Gamma_hat, Delta)
cov_vv = D**2 * self.sig11(Gamma_hat, Delta) / 2
cov_yv = (zeta2(Gamma_hat, Delta) * Gamma_hat * D ) **2 / 2 / (Gamma ** 0.5)
@@ -274,7 +273,7 @@ class StochasticHarmonicOscillator:
cov_matrix[..., 1, 1] = cov_vv
# Compute the Cholesky decomposition to get scale_tril
#scale_tril = torch.linalg.cholesky(cov_matrix)
scale_tril = torch.zeros(*batch_shape, 2, 2, device=y0.device, dtype=y0.dtype)
@@ -298,4 +297,4 @@ class StochasticHarmonicOscillator:
scale_tril=scale_tril
).sample()
return new_yv[...,0], new_yv[...,1]
return new_yv[...,0], new_yv[...,1]
+4 -3
View File
@@ -1,4 +1,5 @@
[pytest]
testpaths = . # Run tests in the current directory
python_files = test_*.py # Run tests in files that start with "test_"
norecursedirs = .. # Don't run tests in the parent directory
# Keep settings value-only; pytest does not treat inline `# ...` as comments.
testpaths = .
python_files = test_*.py
norecursedirs = ..
+9 -17
View File
@@ -1,21 +1,13 @@
#!/usr/bin/env python
"""Basic import tests for LanPaint.
"""Tests for `LanPaint` package."""
The ComfyUI runtime dependencies (e.g. `comfy`) are intentionally optional for unit tests.
"""
import pytest
from src.LanPaint.nodes import Example
@pytest.fixture
def example_node():
"""Fixture to create an Example node instance."""
return Example()
def test_package_imports_without_comfy() -> None:
import LanPaint
def test_example_node_initialization(example_node):
"""Test that the node can be instantiated."""
assert isinstance(example_node, Example)
def test_return_types():
"""Test the node's metadata."""
assert Example.RETURN_TYPES == ("IMAGE",)
assert Example.FUNCTION == "test"
assert Example.CATEGORY == "Example"
assert isinstance(LanPaint.NODE_CLASS_MAPPINGS, dict)
assert isinstance(LanPaint.NODE_DISPLAY_NAME_MAPPINGS, dict)
assert "LanPaint_KSampler" in LanPaint.NODE_CLASS_MAPPINGS
assert LanPaint.WEB_DIRECTORY == "./web"
+104
View File
@@ -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
+74
View File
@@ -0,0 +1,74 @@
import importlib
import sys
import types
import pytest
import torch
def _repeat_to_batch_size(tensor: torch.Tensor, batch_size: int) -> torch.Tensor:
if tensor.shape[0] == batch_size:
return tensor
if tensor.shape[0] == 1:
return tensor.repeat((batch_size,) + (1,) * (tensor.ndim - 1))
repeats = (batch_size + tensor.shape[0] - 1) // tensor.shape[0]
return tensor.repeat((repeats,) + (1,) * (tensor.ndim - 1))[:batch_size]
def _import_nodes(monkeypatch, comfyui_version: str):
comfy_mod = types.ModuleType("comfy")
comfy_mod.__path__ = []
comfy_utils_mod = types.ModuleType("comfy.utils")
comfy_utils_mod.repeat_to_batch_size = _repeat_to_batch_size
comfy_samplers_mod = types.ModuleType("comfy.samplers")
class DummyKSAMPLER: ...
comfy_samplers_mod.KSAMPLER = DummyKSAMPLER
comfy_model_base_mod = types.ModuleType("comfy.model_base")
class ModelType:
FLUX = "FLUX"
FLOW = "FLOW"
class WAN22: ...
comfy_model_base_mod.ModelType = ModelType
comfy_model_base_mod.WAN22 = WAN22
comfyui_version_mod = types.ModuleType("comfyui_version")
comfyui_version_mod.__version__ = comfyui_version
comfy_mod.utils = comfy_utils_mod
comfy_mod.samplers = comfy_samplers_mod
comfy_mod.model_base = comfy_model_base_mod
monkeypatch.setitem(sys.modules, "comfy", comfy_mod)
monkeypatch.setitem(sys.modules, "comfy.utils", comfy_utils_mod)
monkeypatch.setitem(sys.modules, "comfy.samplers", comfy_samplers_mod)
monkeypatch.setitem(sys.modules, "comfy.model_base", comfy_model_base_mod)
monkeypatch.setitem(sys.modules, "nodes", types.ModuleType("nodes"))
monkeypatch.setitem(sys.modules, "latent_preview", types.ModuleType("latent_preview"))
monkeypatch.setitem(sys.modules, "comfyui_version", comfyui_version_mod)
sys.modules.pop("src.LanPaint.nodes", None)
return importlib.import_module("src.LanPaint.nodes")
@pytest.mark.parametrize("comfyui_version", ["0.5.0", "0.6.0"])
def test_reshape_mask_accepts_bhw_and_5d_output_shape(monkeypatch, comfyui_version: str) -> None:
lanpaint_nodes = _import_nodes(monkeypatch, comfyui_version)
input_mask = torch.zeros((1, 4, 4))
output_shape = (1, 16, 1, 8, 8)
out = lanpaint_nodes.reshape_mask(input_mask, output_shape, video_inpainting=False)
assert tuple(out.shape) == output_shape
def test_prepare_mask_accepts_hw_and_moves_device(monkeypatch) -> None:
lanpaint_nodes = _import_nodes(monkeypatch, "0.5.0")
input_mask = torch.zeros((4, 4))
output_shape = (2, 3, 8, 8)
out = lanpaint_nodes.prepare_mask(input_mask, output_shape, device=torch.device("cpu"), video_inpainting=False)
assert tuple(out.shape) == output_shape
assert out.device.type == "cpu"
+45
View File
@@ -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()