Address PR review comments for import fallback and mask handling

This commit is contained in:
xmarre
2026-05-01 16:32:20 +02:00
parent 2239bb7e73
commit a0b728ff13
4 changed files with 31 additions and 3 deletions
+10 -1
View File
@@ -1,6 +1,15 @@
try:
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
except ImportError:
except ModuleNotFoundError as exc:
if exc.name not in {f"{__package__}.nodes", "nodes"}:
raise
# Allows standalone pytest execution from a directory whose name is not a
# valid Python package identifier. ComfyUI imports this file as a package,
# so the relative import path above remains the normal runtime path.
from nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
except ImportError as exc:
if "attempted relative import with no known parent package" not in str(exc):
raise
# Allows standalone pytest execution from a directory whose name is not a
# valid Python package identifier. ComfyUI imports this file as a package,
# so the relative import path above remains the normal runtime path.
+7 -1
View File
@@ -4,7 +4,13 @@ from typing import Any
try:
from .tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
except ImportError:
except ModuleNotFoundError as exc:
if exc.name not in {f"{__package__}.tide_core", "tide_core"}:
raise
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
except ImportError as exc:
if "attempted relative import with no known parent package" not in str(exc):
raise
from tide_core import TIDEConfig, TIDEAttentionOverride, TIDEAttentionPatch, TIDEModelWrapper
+10
View File
@@ -33,6 +33,16 @@ class TIDEConfig:
preserve_existing_wrapper: bool = True
debug: bool = False
def __post_init__(self) -> None:
for name in ("width", "height", "base_width", "base_height", "token_px"):
value = getattr(self, name)
try:
numeric_value = int(value)
except (TypeError, ValueError) as exc:
raise ValueError(f"{name} must be a positive integer-like value, got {value!r}") from exc
if numeric_value <= 0:
raise ValueError(f"{name} must be > 0, got {value!r}")
@property
def target_image_tokens(self) -> int:
return max(1, (int(self.width) // self.token_px) * (int(self.height) // self.token_px))
+4 -1
View File
@@ -194,7 +194,10 @@ class TIDEAttentionOverride:
if mask is None:
mask = kwargs.get("attn_mask", None)
if self.config.force_pytorch_attention_with_mask and mask is not None:
return _sdpa_attention(*args, **kwargs)
sdpa_kwargs = dict(kwargs)
if "mask" not in sdpa_kwargs and "attn_mask" in sdpa_kwargs:
sdpa_kwargs["mask"] = sdpa_kwargs["attn_mask"]
return _sdpa_attention(*args, **sdpa_kwargs)
if self.old_override is not None:
return self.old_override(original_func, *args, **kwargs)
return original_func(*args, **kwargs)