Improvements for Wan 2.2 support

- add xet support and add the xet cache to manageable directories
 - xet is enabled by default
 - fix logging to root in various places
 - improve logging about model unloading and loading
 - TorchCompileNode now supports the VAE
 - torchaudio missing will cause less noise in the logs
 - feature flags will assume to be supporting everything in the distributed progress context
 - fixes progress notifications
This commit is contained in:
doctorpangloss
2025-07-28 14:36:27 -07:00
parent b5a50301f6
commit 03e5430121
19 changed files with 192 additions and 82 deletions
+4
View File
@@ -86,6 +86,9 @@ def init_default_paths(folder_names_and_paths: FolderNames, configuration: Optio
if "HF_HUB_CACHE" in os.environ:
hf_cache_paths.additional_absolute_directory_paths.append(os.environ.get("HF_HUB_CACHE"))
hf_xet = ModelPaths(["xet"], supported_extensions=set())
if "HF_XET_CACHE" in os.environ:
hf_xet.additional_absolute_directory_paths.append(os.environ.get("HF_XET_CACHE"))
model_paths_to_add = [
ModelPaths(["checkpoints"], supported_extensions=set(supported_pt_extensions)),
ModelPaths(["configs"], additional_absolute_directory_paths=[get_package_as_path("comfy.configs")], supported_extensions={".yaml"}),
@@ -107,6 +110,7 @@ def init_default_paths(folder_names_and_paths: FolderNames, configuration: Optio
ModelPaths(["classifiers"], supported_extensions=set()),
ModelPaths(["huggingface"], supported_extensions=set()),
hf_cache_paths,
hf_xet,
]
for model_paths in model_paths_to_add:
if replace_existing:
+1 -1
View File
@@ -18,7 +18,7 @@ from .. import options
from ..app import logger
os.environ["OPENCV_IO_ENABLE_OPENEXR"] = "1"
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
os.environ["HF_XET_HIGH_PERFORMANCE"] = "True"
os.environ["TORCHINDUCTOR_FX_GRAPH_CACHE"] = "1"
os.environ["TORCHINDUCTOR_AUTOGRAD_CACHE"] = "1"
os.environ["BITSANDBYTES_NOWELCOME"] = "1"
+12 -4
View File
@@ -221,7 +221,7 @@ class PromptServer(ExecutorToClientProgress):
handler_args={'max_field_size': 16380},
middlewares=middlewares)
self.sockets = dict()
self.sockets_metadata = dict()
self._sockets_metadata = dict()
self.web_root = (
FrontendManager.init_frontend(args.front_end_version)
if args.front_end_root is None
@@ -278,16 +278,16 @@ class PromptServer(ExecutorToClientProgress):
sid,
)
logging.info(
logger.info(
f"Feature flags negotiated for client {sid}: {client_flags}"
)
first_message = False
except json.JSONDecodeError:
logging.warning(
logger.warning(
f"Invalid JSON received from client {sid}: {msg.data}"
)
except Exception as e:
logging.error(f"Error processing WebSocket message: {e}")
logger.error(f"Error processing WebSocket message: {e}")
finally:
self.sockets.pop(sid, None)
self.sockets_metadata.pop(sid, None)
@@ -1236,3 +1236,11 @@ class PromptServer(ExecutorToClientProgress):
message = encode_text_for_progress(node_id, text)
self.send_sync(BinaryEventTypes.TEXT, message, sid)
@property
def sockets_metadata(self):
return self._sockets_metadata
@sockets_metadata.setter
def sockets_metadata(self, value):
self._sockets_metadata = value