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:
+96
-1
@@ -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):
|
||||
|
||||
@@ -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())
|
||||
@@ -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())
|
||||
@@ -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}" "$@"
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user