feat: add MultiGPU logging enhancements and new CI scripts for workflow execution and log summarization, WanVidwoWrapper Model Loader/Sampler re-implemented

This commit is contained in:
John Pollock
2025-10-07 13:19:09 -05:00
parent b034901a05
commit 27cfb83733
6 changed files with 486 additions and 42 deletions
+96 -1
View File
@@ -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):
+62
View File
@@ -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())
+198
View File
@@ -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())
+30
View File
@@ -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 <workflow.json> [<workflow.json>...]" >&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}" "$@"
+59
View File
@@ -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())
+41 -41
View File
@@ -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)