Address PR review comments for import fallback and mask handling
This commit is contained in:
+10
-1
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user