From 27cfb8373312f1e1bc96628024cdba00fe432922 Mon Sep 17 00:00:00 2001 From: John Pollock Date: Tue, 7 Oct 2025 13:19:09 -0500 Subject: [PATCH] feat: add MultiGPU logging enhancements and new CI scripts for workflow execution and log summarization, WanVidwoWrapper Model Loader/Sampler re-implemented --- __init__.py | 97 ++++++++++++++++++- ci/extract_allocation.py | 62 ++++++++++++ ci/run_workflows.py | 198 +++++++++++++++++++++++++++++++++++++++ ci/smoke_test.sh | 30 ++++++ ci/summarize_log.py | 59 ++++++++++++ wanvideo.py | 82 ++++++++-------- 6 files changed, 486 insertions(+), 42 deletions(-) create mode 100644 ci/extract_allocation.py create mode 100644 ci/run_workflows.py create mode 100644 ci/smoke_test.sh create mode 100644 ci/summarize_log.py diff --git a/__init__.py b/__init__.py index f110ab7..834a750 100644 --- a/__init__.py +++ b/__init__.py @@ -3,6 +3,8 @@ import logging import weakref import os import copy +import json +from datetime import datetime from pathlib import Path import folder_paths import comfy.model_management as mm @@ -27,6 +29,16 @@ DEBUG_LOG = False logger = logging.getLogger("MultiGPU") logger.propagate = False +FOCUS_LOG_LEVEL = logging.INFO + 5 +logging.addLevelName(FOCUS_LOG_LEVEL, "FOCUS") + +if not hasattr(logging.Logger, "focus"): + def focus(self, message, *args, **kwargs): + if self.isEnabledFor(FOCUS_LOG_LEVEL): + self._log(FOCUS_LOG_LEVEL, message, args, **kwargs) + + logging.Logger.focus = focus # type: ignore[attr-defined] + if not logger.handlers: log_level = logging.DEBUG if DEBUG_LOG else logging.INFO handler = logging.StreamHandler() @@ -35,10 +47,93 @@ if not logger.handlers: logger.addHandler(handler) logger.setLevel(log_level) + json_log_path = os.environ.get("MGPU_JSON_LOG_PATH") + json_static_fields = {} + if json_log_path: + try: + json_static_fields = json.loads(os.environ.get("MGPU_JSON_STATIC_FIELDS", "{}")) + except json.JSONDecodeError: + json_static_fields = {} + + level_aliases = { + "CRITICAL": logging.CRITICAL, + "ERROR": logging.ERROR, + "WARNING": logging.WARNING, + "FOCUS": FOCUS_LOG_LEVEL, + "INFO": logging.INFO, + "DEBUG": logging.DEBUG, + } + + json_min_level = FOCUS_LOG_LEVEL + configured_min_level = os.environ.get("MGPU_JSON_MIN_LEVEL") + if configured_min_level: + value = configured_min_level.strip() + upper_value = value.upper() + if upper_value in level_aliases: + json_min_level = level_aliases[upper_value] + else: + try: + json_min_level = int(value) + except ValueError: + json_min_level = FOCUS_LOG_LEVEL + + class JsonLineFileHandler(logging.Handler): + def __init__(self, path, static_fields, min_level, overwrite): + super().__init__() + self.path = Path(path) + self.path.parent.mkdir(parents=True, exist_ok=True) + self.static_fields = static_fields + self.setLevel(min_level) + if overwrite: + try: + with self.path.open("w", encoding="utf-8") as handle: + handle.write("") + except OSError: + pass + + def emit(self, record): + message = record.getMessage() + category = None + if message.startswith("[") and "]" in message: + bracket_split = message.split("]", 1) + category = bracket_split[0].strip("[]") + payload = { + "timestamp": datetime.utcnow().isoformat() + "Z", + "level": record.levelname, + "name": record.name, + "message": message, + } + if category: + payload["event_category"] = category + if hasattr(record, "mgpu_context") and isinstance(record.mgpu_context, dict): + payload.update(record.mgpu_context) + workflow_id = os.environ.get("MGPU_JSON_WORKFLOW") + prompt_id = os.environ.get("MGPU_JSON_PROMPT") + if workflow_id: + payload.setdefault("workflow_id", workflow_id) + if prompt_id: + payload.setdefault("prompt_id", prompt_id) + if self.static_fields: + payload.update(self.static_fields) + try: + with self.path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(payload, ensure_ascii=True) + "\n") + except OSError: + # Fail silently for JSON logging so primary logging continues. + pass + + overwrite_value = os.environ.get("MGPU_JSON_OVERWRITE", "true").strip().lower() + overwrite_enabled = overwrite_value not in {"0", "false", "no"} + + logger.addHandler(JsonLineFileHandler(json_log_path, json_static_fields, json_min_level, overwrite_enabled)) + def mgpu_mm_log_method(self, msg): """Add MultiGPU model management logging method to logger instance.""" if MGPU_MM_LOG: - self.info(f"[MultiGPU Model Management] {msg}") + self.focus( + f"[MultiGPU Model Management] {msg}", + extra={"mgpu_context": {"component": "model_management"}}, + ) logger.mgpu_mm_log = mgpu_mm_log_method.__get__(logger, type(logger)) def check_module_exists(module_path): diff --git a/ci/extract_allocation.py b/ci/extract_allocation.py new file mode 100644 index 0000000..a11c05e --- /dev/null +++ b/ci/extract_allocation.py @@ -0,0 +1,62 @@ +#!/usr/bin/env python3 +"""Filter MultiGPU JSON logs for allocation summaries.""" + +import argparse +import json +from pathlib import Path +from typing import Iterable, Iterator, Dict, Any + + +def load_json_lines(path: Path) -> Iterator[Dict[str, Any]]: + with path.open("r", encoding="utf-8") as handle: + for line in handle: + line = line.strip() + if not line: + continue + try: + yield json.loads(line) + except json.JSONDecodeError: + continue + + +def is_allocation_event(entry: Dict[str, Any], keywords: Iterable[str]) -> bool: + message = entry.get("message", "") + return any(keyword in message for keyword in keywords) + + +def main() -> int: + parser = argparse.ArgumentParser(description="Extract allocation-related events from MultiGPU JSON logs") + parser.add_argument("logfile", type=Path, help="Path to JSONL log produced by MGPU_JSON_LOG_PATH") + parser.add_argument( + "--keywords", + nargs="*", + default=["Final Allocation String", "Total memory", "Virtual VRAM"], + help="Keywords that mark allocation events", + ) + args = parser.parse_args() + + entries = list(load_json_lines(args.logfile)) + if not entries: + print("No entries found in log file.") + return 0 + + matched = [entry for entry in entries if is_allocation_event(entry, args.keywords)] + if not matched: + print("No allocation events matched provided keywords.") + return 0 + + for entry in matched: + timestamp = entry.get("timestamp", "unknown") + category = entry.get("event_category", "") + component = entry.get("component", "") + header_bits = [bit for bit in (timestamp, category, component) if bit] + header = " | ".join(header_bits) if header_bits else "allocation" + print(f"## {header}") + print(entry.get("message", "")) + print() + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/ci/run_workflows.py b/ci/run_workflows.py new file mode 100644 index 0000000..d7f4cd3 --- /dev/null +++ b/ci/run_workflows.py @@ -0,0 +1,198 @@ +#!/usr/bin/env python3 +"""Minimal ComfyUI workflow runner for CI smoke tests.""" + +import argparse +import json +import os +import sys +import time +import uuid +from pathlib import Path +from typing import Iterable, Optional + +import requests +import websocket + +DEFAULT_HOST = os.environ.get("COMFYUI_HOST", "127.0.0.1") +DEFAULT_PORT = int(os.environ.get("COMFYUI_PORT", "8188")) +DEFAULT_CONNECT_TIMEOUT = int(os.environ.get("COMFYUI_CONNECT_TIMEOUT", "60")) +DEFAULT_WORKFLOW_TIMEOUT = int(os.environ.get("COMFYUI_WORKFLOW_TIMEOUT", "900")) + + +class ComfyWorkflowRunner: + def __init__(self, host: str, port: int, connect_timeout: int, workflow_timeout: int) -> None: + self.host = host + self.port = port + self.base_http = f"http://{host}:{port}" + self.base_ws = f"ws://{host}:{port}/ws" + self.connect_timeout = connect_timeout + self.workflow_timeout = workflow_timeout + self.client_id = str(uuid.uuid4()) + self.session = requests.Session() + self.websocket: Optional[websocket.WebSocket] = None + + def wait_for_server(self) -> None: + deadline = time.monotonic() + self.connect_timeout + while time.monotonic() < deadline: + try: + response = self.session.get(f"{self.base_http}/system_stats", timeout=5) + if response.status_code == 200: + return + except requests.RequestException: + time.sleep(1) + raise TimeoutError(f"ComfyUI server not reachable at {self.base_http}") + + def open_websocket(self) -> None: + ws = websocket.WebSocket() + ws.settimeout(5) + ws.connect(f"{self.base_ws}?clientId={self.client_id}") + self.websocket = ws + + def close_websocket(self) -> None: + if self.websocket: + try: + self.websocket.close() + finally: + self.websocket = None + + def queue_prompt(self, prompt: dict) -> str: + payload = {"prompt": prompt, "client_id": self.client_id} + response = self.session.post(f"{self.base_http}/prompt", json=payload, timeout=15) + response.raise_for_status() + data = response.json() + prompt_id = data.get("prompt_id") + if not prompt_id: + raise RuntimeError("No prompt_id returned from ComfyUI") + return prompt_id + + def wait_for_completion(self, prompt_id: str) -> bool: + if not self.websocket: + raise RuntimeError("WebSocket connection not established") + deadline = time.monotonic() + self.workflow_timeout + ws = self.websocket + while time.monotonic() < deadline: + try: + message = ws.recv() + except websocket.WebSocketTimeoutException: + continue + except Exception as exc: # noqa: BLE001 + print(f"WebSocket error: {exc}", file=sys.stderr, flush=True) + return False + + if isinstance(message, bytes): + continue + + try: + payload = json.loads(message) + except json.JSONDecodeError: + continue + + message_type = payload.get("type") + data = payload.get("data", {}) + + if message_type == "execution_error": + if data.get("prompt_id") == prompt_id: + print(f"Execution error: {payload}", file=sys.stderr, flush=True) + return False + elif message_type == "status" and data.get("status") == "error": + if data.get("prompt_id") == prompt_id: + print(f"Status error: {payload}", file=sys.stderr, flush=True) + return False + elif message_type == "executing": + if data.get("prompt_id") == prompt_id and data.get("node") is None: + return True + print("Workflow timed out", file=sys.stderr, flush=True) + return False + + def run_workflow(self, workflow_path: Path) -> bool: + previous_workflow = os.environ.get("MGPU_JSON_WORKFLOW") + previous_prompt = os.environ.get("MGPU_JSON_PROMPT") + + def restore_env() -> None: + if previous_workflow is None: + os.environ.pop("MGPU_JSON_WORKFLOW", None) + else: + os.environ["MGPU_JSON_WORKFLOW"] = previous_workflow + if previous_prompt is None: + os.environ.pop("MGPU_JSON_PROMPT", None) + else: + os.environ["MGPU_JSON_PROMPT"] = previous_prompt + + if workflow_path: + os.environ["MGPU_JSON_WORKFLOW"] = workflow_path.name + try: + with workflow_path.open("r", encoding="utf-8") as handle: + workflow = json.load(handle) + except (OSError, json.JSONDecodeError) as exc: + print(f"Failed to load workflow {workflow_path}: {exc}", file=sys.stderr, flush=True) + restore_env() + return False + + print(f"Running workflow {workflow_path}", flush=True) + start = time.monotonic() + try: + prompt_id = self.queue_prompt(workflow) + os.environ["MGPU_JSON_PROMPT"] = prompt_id + except requests.HTTPError as exc: + print(f"HTTP error while queueing workflow: {exc}", file=sys.stderr, flush=True) + restore_env() + return False + except requests.RequestException as exc: + print(f"Request error while queueing workflow: {exc}", file=sys.stderr, flush=True) + restore_env() + return False + except RuntimeError as exc: + print(str(exc), file=sys.stderr, flush=True) + restore_env() + return False + + try: + if not self.wait_for_completion(prompt_id): + return False + duration = time.monotonic() - start + print(f"Workflow {workflow_path} completed in {duration:.2f}s", flush=True) + return True + finally: + restore_env() + + def run_suite(self, workflows: Iterable[Path], fail_fast: bool) -> bool: + self.wait_for_server() + self.open_websocket() + try: + overall = True + for workflow in workflows: + ok = self.run_workflow(workflow) + if not ok: + overall = False + if fail_fast: + break + return overall + finally: + self.close_websocket() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Run ComfyUI workflows via the HTTP/WebSocket API") + parser.add_argument("workflows", nargs="+", type=Path, help="Workflow files in ComfyUI API JSON format") + parser.add_argument("--host", default=DEFAULT_HOST, help="ComfyUI HTTP host") + parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="ComfyUI HTTP port") + parser.add_argument("--connect-timeout", type=int, default=DEFAULT_CONNECT_TIMEOUT, help="Seconds to wait for the server to come online") + parser.add_argument("--workflow-timeout", type=int, default=DEFAULT_WORKFLOW_TIMEOUT, help="Seconds to wait for each workflow to finish") + parser.add_argument("--fail-fast", action="store_true", help="Stop on first workflow failure") + return parser.parse_args() + + +def main() -> int: + args = parse_args() + runner = ComfyWorkflowRunner( + host=args.host, + port=args.port, + connect_timeout=args.connect_timeout, + workflow_timeout=args.workflow_timeout, + ) + success = runner.run_suite(args.workflows, fail_fast=args.fail_fast) + return 0 if success else 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/ci/smoke_test.sh b/ci/smoke_test.sh new file mode 100644 index 0000000..cfd9a79 --- /dev/null +++ b/ci/smoke_test.sh @@ -0,0 +1,30 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [[ $# -lt 1 ]]; then + echo "Usage: COMFYUI_HOME=/path/to/ComfyUI ci/smoke_test.sh [...]" >&2 + exit 1 +fi + +if [[ -z "${COMFYUI_HOME:-}" ]]; then + echo "COMFYUI_HOME environment variable must point to the ComfyUI checkout" >&2 + exit 1 +fi + +PYTHON_BIN=${PYTHON_BIN:-python3} +HOST=${COMFYUI_HOST:-127.0.0.1} +PORT=${COMFYUI_PORT:-8188} +LOG_FILE=${COMFYUI_LOG:-comfyui_ci.log} + +pushd "${COMFYUI_HOME}" >/dev/null + +${PYTHON_BIN} -m pip install --upgrade pip >/dev/null +${PYTHON_BIN} -m pip install -r requirements.txt >/dev/null + +${PYTHON_BIN} main.py --disable-auto-launch --listen "${HOST}" --port "${PORT}" >"${LOG_FILE}" 2>&1 & +SERVER_PID=$! +trap 'kill ${SERVER_PID} >/dev/null 2>&1 || true' EXIT + +popd >/dev/null + +"${PYTHON_BIN}" "$(dirname "$0")/run_workflows.py" --host "${HOST}" --port "${PORT}" "$@" diff --git a/ci/summarize_log.py b/ci/summarize_log.py new file mode 100644 index 0000000..f1d3b4b --- /dev/null +++ b/ci/summarize_log.py @@ -0,0 +1,59 @@ +#!/usr/bin/env python3 +"""Convert MultiGPU JSON log into a Markdown summary.""" + +import argparse +import json +from pathlib import Path +from typing import Iterator, Dict, Any + + +def load_json_lines(path: Path) -> Iterator[Dict[str, Any]]: + with path.open("r", encoding="utf-8") as handle: + for line in handle: + line = line.strip() + if not line: + continue + try: + yield json.loads(line) + except json.JSONDecodeError: + continue + + +def main() -> int: + parser = argparse.ArgumentParser(description="Summarize MultiGPU JSON logs into Markdown") + parser.add_argument("logfile", type=Path, help="Path to JSONL log produced by MGPU_JSON_LOG_PATH") + parser.add_argument("--severity", nargs="*", help="Optional severity levels to include (e.g. INFO WARN ERROR)") + parser.add_argument( + "--component", + nargs="*", + help="Optional component names to include (matches component or event_category fields)", + ) + args = parser.parse_args() + + entries = list(load_json_lines(args.logfile)) + if not entries: + print("No entries found in log file.") + return 0 + + print("| Timestamp | Level | Component | Message |") + print("| --- | --- | --- | --- |") + for entry in entries: + level = entry.get("level", "") + if args.severity and level not in args.severity: + continue + component_values = { + entry.get("component", ""), + entry.get("event_category", ""), + } + component = next((value for value in component_values if value), "") + if args.component and component not in args.component: + continue + timestamp = entry.get("timestamp", "") + message = entry.get("message", "").replace("|", "\u2502") + print(f"| {timestamp} | {level} | {component} | {message} |") + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/wanvideo.py b/wanvideo.py index 101d85a..e412a11 100644 --- a/wanvideo.py +++ b/wanvideo.py @@ -1,29 +1,8 @@ -""" - WanVideoControlnetLoader, controlnet/nodes.py - FantasyTalkingModelLoader, fantasytalking/nodes.py - MultiTalkModelLoader, multitalk/nodes.py - Wav2VecModelLoader, multitalk/nodes.py - WanVideoSetBlockSwap, nodes.py - WanVideoBlockList, nodes.py - WanVideoTextEncodeCached, nodes.py - X WanVideoTextEncode, nodes.py - WanVideoTextEncodeSingle, nodes.py - WanVideoClipVisionEncode, nodes.py - X WanVideoImageToVideoEncode, nodes.py - WanVideoVACEEncode, nodes.py - X WanVideoDecode, nodes.py - WanVideoImageClipEncode, nodes_deprecated.py - WanVideoModelLoader, nodes_model_loading.py - X WanVideoVAELoader, nodes_model_loading.py - WanVideoLoraBlockEdit, nodes_model_loading.py - X WanVideoTinyVAELoader, nodes_model_loading.py - WanVideoBlockSwap, nodes_model_loading.py - WanVideoVRAMManagement, nodes_model_loading.py - WanVideoTorchCompileSettings, nodes_model_loading.py - X LoadWanVideoT5TextEncoder, nodes_model_loading.py - LoadWanVideoClipTextEncoder, nodes_model_loading.py - WanVideoUni3C_ControlnetLoader, uni3c/nodes.py - """ +"""WanVideoWrapper integration helpers. + +For the current progress checklist and outstanding tasks, see +`.github/instructions/ComfyUI-MultiGPU.instructions.md`. +""" @@ -561,8 +540,8 @@ class WanVideoModelLoader: fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None, rms_norm_function="default"): from . import set_current_device - - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] User selected device: {compute_device}") + + logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] User selected device: {compute_device}") selected_device = torch.device(compute_device) @@ -574,7 +553,7 @@ class WanVideoModelLoader: loader_module = inspect.getmodule(original_loader) if loader_module: - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Patching '{loader_module.__name__}' to use device: {selected_device}") + logger.debug(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Patching '{loader_module.__name__}' to use device: {selected_device}") # Store original values to restore later if needed, though it's less critical in this workflow original_module_device = getattr(loader_module, 'device', None) @@ -586,7 +565,7 @@ class WanVideoModelLoader: if compute_device == "cpu": setattr(loader_module, 'offload_device', selected_device) - logger.info(f"[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Device patching complete. Calling original loader...") + logger.debug("[MultiGPU WanVideoWrapper][WanVideoModelLoaderMultiGPU] Device patching complete. Calling original loader...") # Call the original loader function with all the arguments it expects result = original_loader.loadmodel( @@ -634,18 +613,39 @@ class WanVideoSampler: def process(self, model, compute_device, **kwargs): from . import set_current_device - - logger.info(f"[MultiGPU WanVideoSampler] Received request to process on: {compute_device}") - # Set the global device context for the sampler's operations + logger.info(f"[MultiGPU WanVideoSampler] Received request to process on: {compute_device}") + + # Resolve the target device and update the global sampler context + target_device = None if compute_device: - set_current_device(torch.device(compute_device)) - - # The model is already on the correct device thanks to the patched loader. - # We no longer need to patch the model here. We just need to call the original sampler. - logger.info("[MultiGPU WanVideoSampler] Model is pre-configured. Calling original sampler.") + target_device = torch.device(compute_device) + set_current_device(target_device) + else: + target_device = mm.get_torch_device() original_sampler = NODE_CLASS_MAPPINGS["WanVideoSampler"]() - - # The original sampler will internally use mm.get_torch_device(), which is now correctly set. - return original_sampler.process(model=model, **kwargs) + sampler_module = inspect.getmodule(original_sampler) + + original_module_device = None + if sampler_module is not None: + original_module_device = getattr(sampler_module, "device", None) + setattr(sampler_module, "device", target_device) + + # Align offload device when running on CPU so intermediate tensors stay colocated. + if compute_device == "cpu": + setattr(sampler_module, "offload_device", target_device) + + if original_module_device != target_device: + logger.debug( + f"[MultiGPU WanVideoSampler] Patched sampler module device: {original_module_device} -> {target_device}" + ) + else: + logger.error("[MultiGPU WanVideoSampler] Unable to resolve sampler module for device patching.") + + try: + # The original sampler will internally use mm.get_torch_device(), which is now correctly set. + return original_sampler.process(model=model, **kwargs) + finally: + if sampler_module is not None and original_module_device is not None: + setattr(sampler_module, "device", original_module_device)