Compare commits
23
Commits
@@ -1,3 +1,4 @@
|
||||
# These are supported funding model platforms
|
||||
|
||||
github: robertvoy
|
||||
buy_me_a_coffee: robertvoy
|
||||
|
||||
@@ -19,6 +19,6 @@ jobs:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -7,3 +7,8 @@ __pycache__/
|
||||
node_modules/
|
||||
npm-debug.log*
|
||||
AGENTS.md
|
||||
.claude/
|
||||
|
||||
# Desloppify artifacts
|
||||
.desloppify/
|
||||
scorecard.png
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
|
||||
---
|
||||
|
||||
## Key Features
|
||||
## Key Features
|
||||
|
||||
#### Parallel Workflow Processing
|
||||
- Run your workflow on multiple GPUs simultaneously with varied seeds, collect results on the master
|
||||
@@ -27,12 +27,22 @@
|
||||
- Intelligent distribution
|
||||
- Handles single images and videos
|
||||
|
||||
#### Ease of Use
|
||||
- Auto-setup local workers; easily add remote/cloud ones
|
||||
- Convert any workflow to distributed with 2 nodes
|
||||
- JSON configuration with UI controls
|
||||
|
||||
---
|
||||
#### Ease of Use
|
||||
- Auto-setup local workers; easily add remote/cloud ones
|
||||
- Convert any workflow to distributed with 2 nodes
|
||||
- JSON configuration with UI controls
|
||||
|
||||
---
|
||||
|
||||
## Current Architecture
|
||||
|
||||
- Workflow-level load balancing is controlled by **Distributed Collector** via the `load_balance` toggle.
|
||||
- There is **no Distributed Queue node** anymore.
|
||||
- With `load_balance=true`, orchestration selects one least-busy execution participant:
|
||||
- If master participation is enabled, master is included as a candidate.
|
||||
- If master is in orchestrator-only mode, only workers are considered.
|
||||
|
||||
---
|
||||
|
||||
## Worker Types
|
||||
|
||||
@@ -51,6 +61,7 @@ ComfyUI Distributed supports three types of workers:
|
||||
## Requirements
|
||||
|
||||
- ComfyUI
|
||||
> Note: Desktop app not currently supported
|
||||
- Multiple NVIDIA GPUs
|
||||
> No additional GPUs? Use [Cloud Workers](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/worker-setup-guides.md#cloud-workers)
|
||||
- That's it
|
||||
@@ -81,19 +92,19 @@ Join Runpod with [this link](https://get.runpod.io/0bw29uf3ug0p) and unlock a sp
|
||||
|
||||
## Workflow Examples
|
||||
|
||||
### Basic Parallel Generation
|
||||
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
|
||||
### Basic Parallel Generation
|
||||
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
|
||||
|
||||

|
||||
|
||||
> [Download workflow](/workflows/distributed-txt2img.json)
|
||||
|
||||
1. Open your ComfyUI workflow
|
||||
2. Add **Distributed Seed** → connect to sampler's seed
|
||||
3. Add **Distributed Collector** → after VAE Decode
|
||||
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
|
||||
5. Enable workers in the UI
|
||||
6. Run the workflow!
|
||||
2. Add **Distributed Seed** → connect to sampler's seed
|
||||
3. Add **Distributed Collector** → after VAE Decode
|
||||
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
|
||||
5. Enable workers in the UI
|
||||
6. Run the workflow!
|
||||
|
||||
### Parallel WAN Generation
|
||||
Generate multiple videos in the time it takes to generate one. Each worker uses a different seed.
|
||||
@@ -111,12 +122,12 @@ Generate multiple videos in the time it takes to generate one. Each worker uses
|
||||
7. Enable workers in the UI
|
||||
8. Run the workflow!
|
||||
|
||||
### Distributed Image Upscaling
|
||||
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
|
||||
### Distributed Image Upscaling
|
||||
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
|
||||
|
||||

|
||||
|
||||
> [Download workflow](/workflows/distributed-upscale.json)
|
||||
> [Download workflow](/workflows/distributed-upscale.json)
|
||||
|
||||
1. Load your image
|
||||
2. Upscale with ESRGAN or similar
|
||||
@@ -142,40 +153,40 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
|
||||
|
||||
---
|
||||
|
||||
## Developer API
|
||||
## Developer API
|
||||
|
||||
Control your distributed cluster programmatically without opening the browser.
|
||||
|
||||
* **Endpoint:** `POST /distributed/queue`
|
||||
* **Functionality:** Accepts a ComfyUI API-format prompt, dispatches it to the requested reachable workers, and returns the master `prompt_id`.
|
||||
* **Documentation:** [See API Examples & Scripts](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/comfyui-distributed-api.md)
|
||||
* **Endpoint:** `POST /distributed/queue`
|
||||
* **Functionality:** Accepts a distributed queue payload (`prompt`, `client_id`, `enabled_worker_ids`, optional `workflow`/`delegate_master`/`trace_execution_id`), dispatches to healthy workers, and returns `{prompt_id, worker_count}`.
|
||||
* **Documentation:** [See API Examples & Scripts](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/comfyui-distributed-api.md)
|
||||
|
||||
> **⚠️ Security Warning:** Do not expose your ComfyUI port to the public internet. If you need remote access, run ComfyUI behind a secure proxy (like Cloudflare or a VPN).
|
||||
|
||||
---
|
||||
|
||||
## Distributed Value
|
||||
|
||||
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
|
||||
|
||||
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
|
||||
- The node shows only currently enabled workers.
|
||||
- If worker enablement changes, worker fields update automatically.
|
||||
- When disconnected, it resets to default string mode and clears per-worker overrides.
|
||||
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
|
||||
|
||||
---
|
||||
|
||||
## Nodes
|
||||
|
||||
| Node | Description |
|
||||
|------|-------------|
|
||||
| **Distributed Seed** | Generates unique seeds for each worker |
|
||||
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
|
||||
| **Distributed Value** | Outputs per-worker override values with fallback to default |
|
||||
| **Ultimate SD Upscale Distributed** | Distributes upscale tiles across workers |
|
||||
| **Image Batch Divider** | Splits image batches for multi-GPU output |
|
||||
| **Audio Segment Divider** | Splits an audio waveform into up to ten sequential time segments |
|
||||
> **⚠️ Security Warning:** Do not expose your ComfyUI port to the public internet. If you need remote access, run ComfyUI behind a secure proxy (like Cloudflare or a VPN).
|
||||
|
||||
---
|
||||
|
||||
## Distributed Value
|
||||
|
||||
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
|
||||
|
||||
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
|
||||
- The node shows only currently enabled workers.
|
||||
- If worker enablement changes, worker fields update automatically.
|
||||
- When disconnected, it resets to default string mode and clears per-worker overrides.
|
||||
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
|
||||
|
||||
---
|
||||
|
||||
## Nodes
|
||||
|
||||
| Node | Description |
|
||||
|------|-------------|
|
||||
| **Distributed Seed** | Generates unique seeds for each worker |
|
||||
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
|
||||
| **Distributed Value** | Outputs per-worker override values with fallback to default |
|
||||
| **Ultimate SD Upscale Distributed** | Distributes upscale tiles across workers |
|
||||
| **Image Batch Divider** | Splits image batches for multi-GPU output |
|
||||
| **Audio Batch Divider** | Splits audio batches for multi-GPU output |
|
||||
| **Distributed Model Name** | Passes model paths to workers, enabling workflows to use models not present on the master in orchestrator-only mode |
|
||||
| **Distributed Empty Image** | Produces an empty IMAGE batch used when the master delegates all work |
|
||||
|
||||
@@ -195,7 +206,7 @@ No, it does not speed up the generation of a single image or video. Instead, it
|
||||
|
||||
<details>
|
||||
<summary>Does it work with the ComfyUI desktop app?</summary>
|
||||
Yes, it does now.
|
||||
Currently, it is not compatible with the ComfyUI desktop app.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
@@ -238,4 +249,3 @@ Buy me a coffee at: https://buymeacoffee.com/robertvoy
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+7
-18
@@ -1,22 +1,11 @@
|
||||
"""ComfyUI-Distributed's native V3 extension entrypoint."""
|
||||
from comfy_api.v0_0_2 import ComfyExtension, io
|
||||
"""ComfyUI-Distributed package entrypoint."""
|
||||
from __future__ import annotations
|
||||
|
||||
from .nodes.v3 import NODES
|
||||
from .runtime.bootstrap import initialize
|
||||
from .bootstrap.entrypoint import build_node_mappings, initialize_runtime
|
||||
|
||||
WEB_DIRECTORY = './web'
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
|
||||
initialize_runtime()
|
||||
|
||||
class DistributedExtension(ComfyExtension):
|
||||
async def on_load(self) -> None:
|
||||
initialize()
|
||||
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return list(NODES)
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> DistributedExtension:
|
||||
return DistributedExtension()
|
||||
|
||||
|
||||
__all__ = ['comfy_entrypoint', 'WEB_DIRECTORY']
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
|
||||
|
||||
+21
-5
@@ -1,5 +1,21 @@
|
||||
from . import config_routes # noqa: F401
|
||||
from . import tunnel_routes # noqa: F401
|
||||
from . import worker_routes # noqa: F401
|
||||
from . import job_routes # noqa: F401
|
||||
from . import usdu_routes # noqa: F401
|
||||
"""Distributed API package with explicit route bootstrap."""
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
_ROUTE_MODULES = (
|
||||
"config_routes",
|
||||
"tunnel_routes",
|
||||
"worker_routes",
|
||||
"job_routes",
|
||||
"usdu_routes",
|
||||
)
|
||||
|
||||
|
||||
def bootstrap_routes() -> None:
|
||||
"""Import route modules once so aiohttp decorators register endpoints."""
|
||||
for module_name in _ROUTE_MODULES:
|
||||
import_module(f"{__name__}.{module_name}")
|
||||
|
||||
|
||||
__all__ = ["bootstrap_routes"]
|
||||
|
||||
+97
-79
@@ -1,32 +1,15 @@
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from __future__ import annotations
|
||||
|
||||
from aiohttp import web
|
||||
import server
|
||||
|
||||
try:
|
||||
from ..utils.config import config_transaction, load_config, save_config
|
||||
except ImportError:
|
||||
from ..utils.config import load_config
|
||||
|
||||
try:
|
||||
from ..utils.config import save_config
|
||||
except ImportError:
|
||||
def save_config(_config):
|
||||
return True
|
||||
|
||||
@asynccontextmanager
|
||||
async def config_transaction():
|
||||
config = load_config()
|
||||
original_snapshot = json.dumps(config, sort_keys=True)
|
||||
yield config
|
||||
if json.dumps(config, sort_keys=True) != original_snapshot:
|
||||
save_config(config)
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.config import config_transaction, load_config
|
||||
from ..utils.logging import debug_log
|
||||
from ..utils.network import handle_api_error, normalize_host
|
||||
from .endpoint_policy import run_authorized_endpoint
|
||||
|
||||
|
||||
def _positive_int(value):
|
||||
def _positive_int(value: int) -> bool:
|
||||
return value > 0
|
||||
|
||||
|
||||
@@ -86,68 +69,103 @@ def _apply_field_patch(target: dict, data: dict, field_rules: list) -> None:
|
||||
target[key] = normalizer(value) if (normalizer and value is not None) else value
|
||||
|
||||
|
||||
def _validate_workers_payload(workers_value: object) -> list[str]:
|
||||
"""Validate top-level workers payload shape before persisting config."""
|
||||
if not isinstance(workers_value, list):
|
||||
return ["workers: expected list"]
|
||||
|
||||
errors: list[str] = []
|
||||
seen_ids: set[str] = set()
|
||||
for index, worker in enumerate(workers_value):
|
||||
if not isinstance(worker, dict):
|
||||
errors.append(f"workers[{index}]: expected object")
|
||||
continue
|
||||
|
||||
worker_id = str(worker.get("id", "")).strip()
|
||||
if not worker_id:
|
||||
errors.append(f"workers[{index}].id: expected non-empty string")
|
||||
elif worker_id in seen_ids:
|
||||
errors.append(f"workers[{index}].id: duplicate id '{worker_id}'")
|
||||
else:
|
||||
seen_ids.add(worker_id)
|
||||
|
||||
if "port" in worker and worker.get("port") is not None:
|
||||
try:
|
||||
port = int(worker.get("port"))
|
||||
except (TypeError, ValueError):
|
||||
errors.append(f"workers[{index}].port: expected integer")
|
||||
else:
|
||||
if port <= 0:
|
||||
errors.append(f"workers[{index}].port: expected positive integer")
|
||||
|
||||
return errors
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/config")
|
||||
async def get_config_endpoint(request):
|
||||
config = load_config()
|
||||
return web.json_response(config)
|
||||
async def get_config_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
return web.json_response(load_config())
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config")
|
||||
async def update_config_endpoint(request):
|
||||
async def update_config_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Bulk config update with schema validation."""
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, f"Invalid JSON payload: {e}", 400)
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception as exc:
|
||||
return await handle_api_error(request, f"Invalid JSON payload: {exc}", 400)
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return await handle_api_error(request, "Config payload must be an object", 400)
|
||||
if not isinstance(data, dict):
|
||||
return await handle_api_error(request, "Config payload must be an object", 400)
|
||||
|
||||
validated_settings = {}
|
||||
validated_root = {}
|
||||
errors = []
|
||||
validated_settings = {}
|
||||
validated_root = {}
|
||||
errors = []
|
||||
|
||||
for key, value in data.items():
|
||||
if key not in CONFIG_SCHEMA:
|
||||
errors.append(f"Unknown field: {key}")
|
||||
continue
|
||||
for key, value in data.items():
|
||||
if key not in CONFIG_SCHEMA:
|
||||
errors.append(f"Unknown field: {key}")
|
||||
continue
|
||||
|
||||
expected_type, validator = CONFIG_SCHEMA[key]
|
||||
if not isinstance(value, expected_type):
|
||||
errors.append(f"{key}: expected {expected_type.__name__}")
|
||||
continue
|
||||
expected_type, validator = CONFIG_SCHEMA[key]
|
||||
if not isinstance(value, expected_type):
|
||||
errors.append(f"{key}: expected {expected_type.__name__}")
|
||||
continue
|
||||
|
||||
if validator and not validator(value):
|
||||
errors.append(f"{key}: value {value!r} failed validation")
|
||||
continue
|
||||
if validator and not validator(value):
|
||||
errors.append(f"{key}: value {value!r} failed validation")
|
||||
continue
|
||||
|
||||
if key in _SETTINGS_FIELDS:
|
||||
validated_settings[key] = value
|
||||
else:
|
||||
validated_root[key] = value
|
||||
if key == "workers":
|
||||
errors.extend(_validate_workers_payload(value))
|
||||
if errors:
|
||||
continue
|
||||
|
||||
if errors:
|
||||
return web.json_response({
|
||||
"status": "error",
|
||||
"error": errors,
|
||||
"message": "; ".join(errors),
|
||||
}, status=400)
|
||||
if key in _SETTINGS_FIELDS:
|
||||
validated_settings[key] = value
|
||||
else:
|
||||
validated_root[key] = value
|
||||
|
||||
if errors:
|
||||
return await handle_api_error(request, errors, 400)
|
||||
|
||||
try:
|
||||
async with config_transaction() as config:
|
||||
settings = config.setdefault("settings", {})
|
||||
settings.update(validated_settings)
|
||||
for key, value in validated_root.items():
|
||||
config[key] = value
|
||||
return web.json_response({"status": "success", "config": config})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/queue_status/{job_id}")
|
||||
async def queue_status_endpoint(request):
|
||||
async def queue_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Check if a job queue is initialized."""
|
||||
try:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
job_id = request.match_info['job_id']
|
||||
|
||||
# Import to ensure initialization
|
||||
@@ -159,12 +177,12 @@ async def queue_status_endpoint(request):
|
||||
|
||||
debug_log(f"Queue status check for job {job_id}: {'exists' if exists else 'not found'}")
|
||||
return web.json_response({"exists": exists, "job_id": job_id})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_worker")
|
||||
async def update_worker_endpoint(request):
|
||||
try:
|
||||
async def update_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
data = await request.json()
|
||||
worker_id = data.get("worker_id")
|
||||
|
||||
@@ -203,12 +221,12 @@ async def update_worker_endpoint(request):
|
||||
)
|
||||
|
||||
return web.json_response({"status": "success"})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/delete_worker")
|
||||
async def delete_worker_endpoint(request):
|
||||
try:
|
||||
async def delete_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
data = await request.json()
|
||||
worker_id = data.get("worker_id")
|
||||
|
||||
@@ -235,13 +253,13 @@ async def delete_worker_endpoint(request):
|
||||
"status": "success",
|
||||
"message": f"Worker {removed_worker.get('name', worker_id)} deleted"
|
||||
})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_setting")
|
||||
async def update_setting_endpoint(request):
|
||||
async def update_setting_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Updates a specific key in the settings object."""
|
||||
try:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
data = await request.json()
|
||||
key = data.get("key")
|
||||
value = data.get("value")
|
||||
@@ -258,13 +276,13 @@ async def update_setting_endpoint(request):
|
||||
config['settings'][key] = value
|
||||
|
||||
return web.json_response({"status": "success", "message": f"Setting '{key}' updated."})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_master")
|
||||
async def update_master_endpoint(request):
|
||||
async def update_master_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Updates master configuration."""
|
||||
try:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
data = await request.json()
|
||||
|
||||
async with config_transaction() as config:
|
||||
@@ -273,5 +291,5 @@ async def update_master_endpoint(request):
|
||||
_apply_field_patch(config['master'], data, _MASTER_FIELDS)
|
||||
|
||||
return web.json_response({"status": "success", "message": "Master configuration updated."})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Shared endpoint policy helpers for distributed API routes."""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from ..utils.network import handle_api_error
|
||||
from .request_guards import authorization_error_or_none
|
||||
|
||||
|
||||
async def run_authorized_endpoint(
|
||||
request: web.Request,
|
||||
operation: Callable[[], Awaitable[web.StreamResponse]],
|
||||
*,
|
||||
unexpected_status: int = 500,
|
||||
) -> web.StreamResponse:
|
||||
"""Run endpoint operation with shared auth guard and fallback error mapping."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
return await operation()
|
||||
except Exception as exc:
|
||||
return await handle_api_error(request, exc, unexpected_status)
|
||||
+94
-54
@@ -1,3 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import io
|
||||
@@ -5,6 +7,7 @@ import os
|
||||
import base64
|
||||
import binascii
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from aiohttp import web
|
||||
import server
|
||||
@@ -15,22 +18,21 @@ from ..utils.logging import debug_log
|
||||
from ..utils.image import pil_to_tensor, ensure_contiguous
|
||||
from ..utils.network import handle_api_error
|
||||
from ..utils.constants import JOB_INIT_GRACE_PERIOD, MEMORY_CLEAR_DELAY
|
||||
try:
|
||||
from .queue_orchestration import ensure_distributed_state, orchestrate_distributed_execution
|
||||
except ImportError:
|
||||
from .queue_orchestration import orchestrate_distributed_execution
|
||||
|
||||
def ensure_distributed_state():
|
||||
return None
|
||||
from ..utils.runtime_state import ensure_distributed_runtime_state
|
||||
from .request_guards import authorization_error_or_none
|
||||
from .schemas import require_bool_literal
|
||||
from .queue_orchestration import orchestrate_distributed_execution
|
||||
from .queue_request import parse_queue_request_payload
|
||||
|
||||
prompt_server = server.PromptServer.instance
|
||||
|
||||
# Canonical worker result envelope accepted by POST /distributed/job_complete:
|
||||
# { "job_id": str, "worker_id": str, "batch_idx": int, "image": <base64 PNG>, "is_last": bool }
|
||||
|
||||
|
||||
def _decode_image_sync(image_path):
|
||||
def _runtime_state():
|
||||
return ensure_distributed_runtime_state()
|
||||
|
||||
|
||||
def _decode_image_sync(image_path: str) -> dict[str, Any]:
|
||||
"""Decode image/video file and compute hash in a threadpool worker."""
|
||||
import base64
|
||||
import hashlib
|
||||
@@ -76,7 +78,7 @@ def _decode_image_sync(image_path):
|
||||
}
|
||||
|
||||
|
||||
def _check_file_sync(filename, expected_hash):
|
||||
def _check_file_sync(filename: str, expected_hash: str) -> dict[str, Any]:
|
||||
"""Check file presence and hash in a threadpool worker."""
|
||||
import hashlib
|
||||
import folder_paths
|
||||
@@ -101,7 +103,7 @@ def _check_file_sync(filename, expected_hash):
|
||||
}
|
||||
|
||||
|
||||
def _decode_canonical_png_tensor(image_payload):
|
||||
def _decode_canonical_png_tensor(image_payload: str) -> torch.Tensor:
|
||||
"""Decode canonical base64 PNG payload into a contiguous IMAGE tensor."""
|
||||
if not isinstance(image_payload, str) or not image_payload.strip():
|
||||
raise ValueError("Field 'image' must be a non-empty base64 PNG string.")
|
||||
@@ -132,7 +134,7 @@ def _decode_canonical_png_tensor(image_payload):
|
||||
raise ValueError(f"Failed to decode PNG image payload: {exc}") from exc
|
||||
|
||||
|
||||
def _decode_audio_payload(audio_payload):
|
||||
def _decode_audio_payload(audio_payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Decode canonical audio payload into an AUDIO dict."""
|
||||
from ..utils.audio_payload import decode_audio_payload
|
||||
|
||||
@@ -140,17 +142,20 @@ def _decode_audio_payload(audio_payload):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/prepare_job")
|
||||
async def prepare_job_endpoint(request):
|
||||
async def prepare_job_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
multi_job_id = data.get('multi_job_id')
|
||||
if not multi_job_id:
|
||||
return await handle_api_error(request, "Missing multi_job_id", 400)
|
||||
|
||||
ensure_distributed_state()
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_jobs:
|
||||
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
|
||||
runtime_state = _runtime_state()
|
||||
async with runtime_state.distributed_jobs_lock:
|
||||
if multi_job_id not in runtime_state.distributed_pending_jobs:
|
||||
runtime_state.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
|
||||
|
||||
debug_log(f"Prepared queue for job {multi_job_id}")
|
||||
return web.json_response({"status": "success"})
|
||||
@@ -158,9 +163,13 @@ async def prepare_job_endpoint(request):
|
||||
return await handle_api_error(request, e)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/clear_memory")
|
||||
async def clear_memory_endpoint(request):
|
||||
async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
debug_log("Received request to clear VRAM.")
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
warnings: list[str] = []
|
||||
# Use ComfyUI's prompt server queue system like the /free endpoint does
|
||||
if hasattr(server.PromptServer.instance, 'prompt_queue'):
|
||||
server.PromptServer.instance.prompt_queue.set_flag("unload_models", True)
|
||||
@@ -177,12 +186,16 @@ async def clear_memory_endpoint(request):
|
||||
try:
|
||||
mm.unload_all_models()
|
||||
except AttributeError as e:
|
||||
debug_log(f"Warning during model unload: {e}")
|
||||
warning = f"Model unload warning: {e}"
|
||||
warnings.append(warning)
|
||||
debug_log(warning)
|
||||
|
||||
try:
|
||||
mm.soft_empty_cache()
|
||||
except Exception as e:
|
||||
debug_log(f"Warning during cache clear: {e}")
|
||||
warning = f"Cache clear warning: {e}"
|
||||
warnings.append(warning)
|
||||
debug_log(warning)
|
||||
|
||||
for _ in range(3):
|
||||
gc.collect()
|
||||
@@ -191,6 +204,16 @@ async def clear_memory_endpoint(request):
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
if warnings:
|
||||
debug_log("VRAM cleared with warnings.")
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "partial",
|
||||
"message": "GPU memory cleared with warnings.",
|
||||
"warnings": warnings,
|
||||
}
|
||||
)
|
||||
|
||||
debug_log("VRAM cleared successfully.")
|
||||
return web.json_response({"status": "success", "message": "GPU memory cleared."})
|
||||
except Exception as e:
|
||||
@@ -199,13 +222,16 @@ async def clear_memory_endpoint(request):
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
debug_log(f"Partial VRAM clear completed with warning: {e}")
|
||||
return web.json_response({"status": "success", "message": "GPU memory cleared (with warnings)"})
|
||||
debug_log(f"VRAM clear failed: {e}")
|
||||
return await handle_api_error(request, f"GPU memory clear failed: {e}", 500)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/queue")
|
||||
async def distributed_queue_endpoint(request):
|
||||
async def distributed_queue_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Queue a distributed workflow, mirroring the UI orchestration pipeline."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
raw_payload = await request.json()
|
||||
except Exception as exc:
|
||||
@@ -217,7 +243,7 @@ async def distributed_queue_endpoint(request):
|
||||
return await handle_api_error(request, exc, 400)
|
||||
|
||||
try:
|
||||
prompt_id, prompt_number, worker_count, node_errors = await orchestrate_distributed_execution(
|
||||
prompt_id, worker_count = await orchestrate_distributed_execution(
|
||||
payload.prompt,
|
||||
payload.workflow_meta,
|
||||
payload.client_id,
|
||||
@@ -227,17 +253,17 @@ async def distributed_queue_endpoint(request):
|
||||
)
|
||||
return web.json_response({
|
||||
"prompt_id": prompt_id,
|
||||
"number": prompt_number,
|
||||
"node_errors": node_errors,
|
||||
"worker_count": worker_count,
|
||||
"auto_prepare_supported": True,
|
||||
})
|
||||
except Exception as exc:
|
||||
return await handle_api_error(request, exc, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/load_image")
|
||||
async def load_image_endpoint(request):
|
||||
async def load_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Load an image or video file and return it as base64 data with hash."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
image_path = data.get("image_path")
|
||||
@@ -253,8 +279,11 @@ async def load_image_endpoint(request):
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/check_file")
|
||||
async def check_file_endpoint(request):
|
||||
async def check_file_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Check if a file exists and matches the given hash."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
filename = data.get("filename")
|
||||
@@ -271,7 +300,10 @@ async def check_file_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/job_complete")
|
||||
async def job_complete_endpoint(request):
|
||||
async def job_complete_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception as exc:
|
||||
@@ -286,7 +318,7 @@ async def job_complete_endpoint(request):
|
||||
batch_idx = data.get("batch_idx")
|
||||
image_payload = data.get("image")
|
||||
audio_payload = data.get("audio")
|
||||
is_last = data.get("is_last")
|
||||
is_last_raw = data.get("is_last")
|
||||
|
||||
errors = []
|
||||
if not isinstance(job_id, str) or not job_id.strip():
|
||||
@@ -295,42 +327,50 @@ async def job_complete_endpoint(request):
|
||||
errors.append("worker_id: expected non-empty string")
|
||||
if not isinstance(batch_idx, int) or batch_idx < 0:
|
||||
errors.append("batch_idx: expected non-negative integer")
|
||||
image_provided = image_payload is not None
|
||||
audio_provided = audio_payload is not None
|
||||
if not image_provided and not audio_provided:
|
||||
errors.append("expected at least one of image or audio")
|
||||
if image_provided and (not isinstance(image_payload, str) or not image_payload.strip()):
|
||||
if not isinstance(image_payload, str) or not image_payload.strip():
|
||||
errors.append("image: expected non-empty base64 PNG string")
|
||||
if audio_provided and not isinstance(audio_payload, dict):
|
||||
if audio_payload is not None and not isinstance(audio_payload, dict):
|
||||
errors.append("audio: expected object when provided")
|
||||
if not isinstance(is_last, bool):
|
||||
errors.append("is_last: expected boolean")
|
||||
|
||||
try:
|
||||
is_last = require_bool_literal(is_last_raw, field_name="is_last")
|
||||
except ValueError as exc:
|
||||
errors.append(str(exc))
|
||||
is_last = False
|
||||
if errors:
|
||||
return await handle_api_error(request, errors, 400)
|
||||
|
||||
try:
|
||||
tensor = _decode_canonical_png_tensor(image_payload) if image_provided else None
|
||||
decoded_audio = _decode_audio_payload(audio_payload) if audio_provided else None
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 400)
|
||||
tensor = _decode_canonical_png_tensor(image_payload)
|
||||
decoded_audio = _decode_audio_payload(audio_payload) if audio_payload is not None else None
|
||||
multi_job_id = job_id.strip()
|
||||
worker_id = worker_id.strip()
|
||||
|
||||
runtime_state = _runtime_state()
|
||||
allowed_workers = runtime_state.distributed_job_allowed_workers.get(multi_job_id)
|
||||
if allowed_workers is not None and worker_id not in allowed_workers:
|
||||
return await handle_api_error(
|
||||
request,
|
||||
f"Unauthorized worker_id for job {multi_job_id}",
|
||||
403,
|
||||
)
|
||||
|
||||
pending = None
|
||||
queue_size = 0
|
||||
deadline = time.monotonic() + float(JOB_INIT_GRACE_PERIOD)
|
||||
while pending is None:
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
pending = prompt_server.distributed_pending_jobs.get(multi_job_id)
|
||||
async with runtime_state.distributed_jobs_lock:
|
||||
pending = runtime_state.distributed_pending_jobs.get(multi_job_id)
|
||||
if pending is not None:
|
||||
queue_item = {
|
||||
"tensor": tensor,
|
||||
"worker_id": worker_id,
|
||||
"image_index": int(batch_idx),
|
||||
"is_last": is_last,
|
||||
}
|
||||
if decoded_audio is not None:
|
||||
queue_item["audio"] = decoded_audio
|
||||
await pending.put(
|
||||
{
|
||||
"tensor": tensor,
|
||||
"worker_id": worker_id,
|
||||
"image_index": int(batch_idx),
|
||||
"is_last": is_last,
|
||||
"audio": decoded_audio,
|
||||
}
|
||||
queue_item
|
||||
)
|
||||
queue_size = pending.qsize()
|
||||
break
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
# Orchestration helpers split out from queue_orchestration.py:
|
||||
# - prompt_transform.py: graph pruning + hidden input overrides
|
||||
# - media_sync.py: remote media/path normalization
|
||||
# - dispatch.py: worker probe + prompt dispatch
|
||||
"""Orchestration helpers used by distributed queue execution."""
|
||||
|
||||
__all__ = [
|
||||
"dispatch",
|
||||
"media_sync",
|
||||
"prompt_transform",
|
||||
]
|
||||
|
||||
@@ -1,45 +1,33 @@
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import aiohttp
|
||||
|
||||
from ...utils.auth import distributed_auth_headers
|
||||
from ...utils.config import load_config
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.network import build_worker_url, get_client_session, probe_worker
|
||||
try:
|
||||
from ...utils.trace_logger import trace_debug, trace_info
|
||||
except ImportError:
|
||||
def trace_debug(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
def trace_info(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
try:
|
||||
from ..schemas import parse_positive_int
|
||||
except ImportError:
|
||||
def parse_positive_int(value, default):
|
||||
try:
|
||||
parsed = int(value)
|
||||
return parsed if parsed > 0 else default
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
from ...utils.trace_logger import trace_debug, trace_info
|
||||
from ..schemas import coerce_positive_int
|
||||
|
||||
_least_busy_rr_index = 0
|
||||
|
||||
|
||||
async def worker_is_active(worker):
|
||||
async def worker_is_active(worker: dict[str, Any]) -> bool:
|
||||
"""Ping worker's /prompt endpoint to confirm it's reachable."""
|
||||
url = build_worker_url(worker)
|
||||
return await probe_worker(url, timeout=3.0) is not None
|
||||
|
||||
|
||||
async def worker_ws_is_active(worker):
|
||||
async def worker_ws_is_active(worker: dict[str, Any]) -> bool:
|
||||
"""Ping worker's websocket endpoint to confirm it's reachable."""
|
||||
session = await get_client_session()
|
||||
url = build_worker_url(worker, "/distributed/worker_ws")
|
||||
headers = distributed_auth_headers(load_config())
|
||||
try:
|
||||
ws = await session.ws_connect(url, heartbeat=20, timeout=3)
|
||||
ws = await session.ws_connect(url, heartbeat=20, timeout=3, headers=headers)
|
||||
await ws.close()
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
@@ -59,7 +47,13 @@ async def _probe_worker_active(worker, use_websocket, semaphore):
|
||||
return worker, is_active
|
||||
|
||||
|
||||
async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
|
||||
async def _dispatch_via_websocket(
|
||||
worker_url,
|
||||
payload,
|
||||
client_id,
|
||||
headers: dict[str, str],
|
||||
timeout=60.0,
|
||||
):
|
||||
"""Open a fresh worker websocket, dispatch one prompt, wait for ack, then close."""
|
||||
request_id = uuid.uuid4().hex
|
||||
ws_payload = {
|
||||
@@ -73,7 +67,12 @@ async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
|
||||
ws_url = f"{ws_url}/distributed/worker_ws"
|
||||
session = await get_client_session()
|
||||
|
||||
async with session.ws_connect(ws_url, heartbeat=20, timeout=timeout) as ws:
|
||||
async with session.ws_connect(
|
||||
ws_url,
|
||||
heartbeat=20,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
) as ws:
|
||||
await ws.send_json(ws_payload)
|
||||
async for msg in ws:
|
||||
if msg.type == aiohttp.WSMsgType.TEXT:
|
||||
@@ -96,13 +95,13 @@ async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
|
||||
|
||||
|
||||
async def dispatch_worker_prompt(
|
||||
worker,
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id=None,
|
||||
use_websocket=False,
|
||||
trace_execution_id=None,
|
||||
):
|
||||
worker: dict[str, Any],
|
||||
prompt_obj: dict[str, Any],
|
||||
workflow_meta: dict[str, Any] | None,
|
||||
client_id: str | None = None,
|
||||
use_websocket: bool = False,
|
||||
trace_execution_id: str | None = None,
|
||||
) -> None:
|
||||
"""Send the prepared prompt to a worker ComfyUI instance."""
|
||||
worker_url = build_worker_url(worker)
|
||||
url = build_worker_url(worker, "/prompt")
|
||||
@@ -114,6 +113,7 @@ async def dispatch_worker_prompt(
|
||||
payload["extra_data"] = extra_data
|
||||
|
||||
if use_websocket:
|
||||
ws_headers = distributed_auth_headers(load_config())
|
||||
try:
|
||||
await _dispatch_via_websocket(
|
||||
worker_url,
|
||||
@@ -122,6 +122,7 @@ async def dispatch_worker_prompt(
|
||||
"workflow": workflow_meta,
|
||||
},
|
||||
client_id,
|
||||
ws_headers,
|
||||
)
|
||||
return
|
||||
except Exception as exc:
|
||||
@@ -142,14 +143,14 @@ async def dispatch_worker_prompt(
|
||||
|
||||
|
||||
async def select_active_workers(
|
||||
workers,
|
||||
use_websocket,
|
||||
delegate_master,
|
||||
trace_execution_id=None,
|
||||
probe_concurrency=8,
|
||||
):
|
||||
workers: list[dict[str, Any]],
|
||||
use_websocket: bool,
|
||||
delegate_master: bool,
|
||||
trace_execution_id: str | None = None,
|
||||
probe_concurrency: int = 8,
|
||||
) -> tuple[list[dict[str, Any]], bool]:
|
||||
"""Probe workers and return (active_workers, updated_delegate_master)."""
|
||||
probe_limit = parse_positive_int(probe_concurrency, 8)
|
||||
probe_limit = coerce_positive_int(probe_concurrency, 8)
|
||||
probe_semaphore = asyncio.Semaphore(probe_limit)
|
||||
|
||||
if trace_execution_id and workers:
|
||||
@@ -223,16 +224,16 @@ def _select_idle_round_robin(statuses):
|
||||
|
||||
|
||||
async def select_least_busy_worker(
|
||||
workers,
|
||||
trace_execution_id=None,
|
||||
probe_concurrency=8,
|
||||
probe_timeout=3.0,
|
||||
):
|
||||
workers: list[dict[str, Any]],
|
||||
trace_execution_id: str | None = None,
|
||||
probe_concurrency: int = 8,
|
||||
probe_timeout: float = 3.0,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Select one worker by queue depth, round-robin among idle workers."""
|
||||
if not workers:
|
||||
return None
|
||||
|
||||
probe_limit = parse_positive_int(probe_concurrency, 8)
|
||||
probe_limit = coerce_positive_int(probe_concurrency, 8)
|
||||
probe_semaphore = asyncio.Semaphore(probe_limit)
|
||||
statuses = await asyncio.gather(
|
||||
*[
|
||||
@@ -266,3 +267,50 @@ async def select_least_busy_worker(
|
||||
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}"
|
||||
)
|
||||
return worker
|
||||
|
||||
|
||||
async def rank_workers_by_load(
|
||||
workers: list[dict[str, Any]],
|
||||
trace_execution_id: str | None = None,
|
||||
probe_concurrency: int = 8,
|
||||
probe_timeout: float = 3.0,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return all workers sorted by queue depth (ascending), preserving unreachable workers at the end."""
|
||||
if not workers:
|
||||
return []
|
||||
|
||||
probe_limit = coerce_positive_int(probe_concurrency, 8)
|
||||
probe_semaphore = asyncio.Semaphore(probe_limit)
|
||||
statuses = await asyncio.gather(
|
||||
*[
|
||||
_probe_worker_queue(worker, probe_semaphore, probe_timeout)
|
||||
for worker in workers
|
||||
]
|
||||
)
|
||||
|
||||
ranked_statuses = []
|
||||
unreachable_workers = []
|
||||
for worker, status in zip(workers, statuses):
|
||||
if status is None:
|
||||
unreachable_workers.append(worker)
|
||||
continue
|
||||
ranked_statuses.append(status)
|
||||
|
||||
ranked_statuses.sort(key=lambda status: status["queue_remaining"])
|
||||
ranked_workers = [status["worker"] for status in ranked_statuses] + unreachable_workers
|
||||
|
||||
if trace_execution_id:
|
||||
trace_debug(
|
||||
trace_execution_id,
|
||||
"Ranked workers by load: "
|
||||
+ ", ".join(
|
||||
f"{status['worker'].get('id')}={status['queue_remaining']}"
|
||||
for status in ranked_statuses
|
||||
)
|
||||
+ (
|
||||
f"; unreachable={','.join(worker.get('id', '') for worker in unreachable_workers)}"
|
||||
if unreachable_workers
|
||||
else ""
|
||||
),
|
||||
)
|
||||
return ranked_workers
|
||||
|
||||
@@ -3,9 +3,12 @@ import hashlib
|
||||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
import aiohttp
|
||||
|
||||
from ...utils.auth import distributed_auth_headers
|
||||
from ...utils.config import load_config
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.network import build_worker_url, get_client_session
|
||||
from ...utils.trace_logger import trace_debug, trace_info
|
||||
@@ -33,7 +36,7 @@ def _normalize_media_reference(value):
|
||||
return None
|
||||
|
||||
|
||||
def convert_paths_for_platform(obj, target_separator):
|
||||
def convert_paths_for_platform(obj: Any, target_separator: str) -> Any:
|
||||
"""Recursively normalize likely file paths for the worker platform separator."""
|
||||
if target_separator not in ("/", "\\"):
|
||||
return obj
|
||||
@@ -124,12 +127,16 @@ def _load_media_file_sync(filename):
|
||||
return file_bytes, file_hash, mime_type
|
||||
|
||||
|
||||
async def fetch_worker_path_separator(worker, trace_execution_id=None):
|
||||
async def fetch_worker_path_separator(
|
||||
worker: dict[str, Any],
|
||||
trace_execution_id: str | None = None,
|
||||
) -> str | None:
|
||||
"""Best-effort fetch of a worker's path separator from /distributed/system_info."""
|
||||
url = build_worker_url(worker, "/distributed/system_info")
|
||||
session = await get_client_session()
|
||||
headers = distributed_auth_headers(load_config())
|
||||
try:
|
||||
async with session.get(url, timeout=aiohttp.ClientTimeout(total=5)) as resp:
|
||||
async with session.get(url, headers=headers, timeout=aiohttp.ClientTimeout(total=5)) as resp:
|
||||
if resp.status != 200:
|
||||
return None
|
||||
payload = await resp.json()
|
||||
@@ -146,6 +153,7 @@ async def fetch_worker_path_separator(worker, trace_execution_id=None):
|
||||
async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_type, trace_execution_id=None):
|
||||
"""Upload one media file to worker iff missing or hash-mismatched."""
|
||||
session = await get_client_session()
|
||||
headers = distributed_auth_headers(load_config())
|
||||
normalized = filename.replace("\\", "/")
|
||||
|
||||
check_url = build_worker_url(worker, "/distributed/check_file")
|
||||
@@ -153,6 +161,7 @@ async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_
|
||||
async with session.post(
|
||||
check_url,
|
||||
json={"filename": normalized, "hash": file_hash},
|
||||
headers=headers,
|
||||
timeout=aiohttp.ClientTimeout(total=6),
|
||||
) as resp:
|
||||
if resp.status == 200:
|
||||
@@ -193,7 +202,11 @@ async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_
|
||||
return True, worker_path
|
||||
|
||||
|
||||
async def sync_worker_media(worker, prompt_obj, trace_execution_id=None):
|
||||
async def sync_worker_media(
|
||||
worker: dict[str, Any],
|
||||
prompt_obj: dict[str, Any],
|
||||
trace_execution_id: str | None = None,
|
||||
) -> None:
|
||||
"""Sync referenced media files from master to a remote worker before dispatch."""
|
||||
media_refs = _find_media_references(prompt_obj)
|
||||
if not media_refs:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
|
||||
from ...utils.logging import debug_log
|
||||
|
||||
@@ -27,7 +28,7 @@ class PromptIndex:
|
||||
def nodes_for_class(self, class_name):
|
||||
return self.nodes_by_class.get(class_name, [])
|
||||
|
||||
def has_upstream(self, start_node_id, target_class):
|
||||
def has_upstream(self, start_node_id: str, target_class: str) -> bool:
|
||||
cache_key = (str(start_node_id), target_class)
|
||||
if cache_key in self._upstream_cache:
|
||||
return self._upstream_cache[cache_key]
|
||||
@@ -88,6 +89,38 @@ def _find_downstream_nodes(prompt_obj, start_ids):
|
||||
return connected
|
||||
|
||||
|
||||
def _find_downstream_of_output_slot(prompt_obj, node_id, slot_index):
|
||||
"""Return all nodes reachable from a specific output slot on one node."""
|
||||
source_id = str(node_id)
|
||||
try:
|
||||
slot = int(slot_index)
|
||||
except (TypeError, ValueError):
|
||||
return set()
|
||||
|
||||
direct_consumers = set()
|
||||
for candidate_id, candidate_node in _iter_prompt_nodes(prompt_obj):
|
||||
inputs = candidate_node.get("inputs", {})
|
||||
for value in inputs.values():
|
||||
if not (isinstance(value, list) and len(value) == 2):
|
||||
continue
|
||||
if str(value[0]) != source_id:
|
||||
continue
|
||||
try:
|
||||
input_slot = int(value[1])
|
||||
except (TypeError, ValueError):
|
||||
debug_log(
|
||||
f"Prompt transform: skipping malformed slot reference for node {candidate_id}: {value}"
|
||||
)
|
||||
continue
|
||||
if input_slot == slot:
|
||||
direct_consumers.add(str(candidate_id))
|
||||
break
|
||||
|
||||
if not direct_consumers:
|
||||
return set()
|
||||
return _find_downstream_nodes(prompt_obj, list(direct_consumers))
|
||||
|
||||
|
||||
def _create_numeric_id_generator(prompt_obj):
|
||||
"""Return a closure that yields new numeric string IDs."""
|
||||
max_id = 0
|
||||
@@ -95,6 +128,7 @@ def _create_numeric_id_generator(prompt_obj):
|
||||
try:
|
||||
numeric = int(node_id)
|
||||
except (TypeError, ValueError):
|
||||
debug_log(f"Prompt transform: ignoring non-numeric node id while generating ids: {node_id}")
|
||||
continue
|
||||
max_id = max(max_id, numeric)
|
||||
|
||||
@@ -125,218 +159,193 @@ def _find_upstream_nodes(prompt_obj, start_ids):
|
||||
return connected
|
||||
|
||||
|
||||
_DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES = {
|
||||
"PrimitiveBoolean",
|
||||
"PrimitiveFloat",
|
||||
"PrimitiveInt",
|
||||
"PrimitiveNode",
|
||||
"PrimitiveString",
|
||||
}
|
||||
|
||||
_DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES = {
|
||||
"LoadImage",
|
||||
}
|
||||
|
||||
_DELEGATE_MASTER_SAFE_SCALAR_TYPES = {"BOOLEAN", "FLOAT", "INT", "STRING"}
|
||||
_DELEGATE_MASTER_SAFE_LIST_TYPES = {"LIST"}
|
||||
|
||||
# ComfyUI 0.23 exposes CreateList via the newer schema API rather than the
|
||||
# legacy RETURN_TYPES/INPUT_TYPES attributes. Treat it as a safe config utility
|
||||
# only after its connected inputs recursively prove safe.
|
||||
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES = {"CreateList"}
|
||||
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES = {
|
||||
"CreateList": ("inputs.",),
|
||||
}
|
||||
|
||||
# Test hook. At runtime this stays None and the ComfyUI node registry is loaded lazily.
|
||||
_DELEGATE_MASTER_NODE_CLASS_MAPPINGS = None
|
||||
def _resolve_participants(enabled_worker_ids, delegate_master):
|
||||
worker_ids = [str(worker_id) for worker_id in (enabled_worker_ids or [])]
|
||||
if delegate_master:
|
||||
return worker_ids
|
||||
return ["master"] + worker_ids
|
||||
|
||||
|
||||
def _get_delegate_master_node_class_mappings():
|
||||
"""Return ComfyUI node-class mappings when available."""
|
||||
if _DELEGATE_MASTER_NODE_CLASS_MAPPINGS is not None:
|
||||
return _DELEGATE_MASTER_NODE_CLASS_MAPPINGS
|
||||
try:
|
||||
import nodes as comfy_nodes # type: ignore
|
||||
except Exception: # pragma: no cover - depends on ComfyUI runtime imports
|
||||
return {}
|
||||
return getattr(comfy_nodes, "NODE_CLASS_MAPPINGS", {}) or {}
|
||||
def _remove_dangling_input_refs(prompt_obj):
|
||||
"""Drop input links that point to nodes no longer present in prompt_obj."""
|
||||
existing_ids = set(prompt_obj.keys())
|
||||
for _node_id, node in _iter_prompt_nodes(prompt_obj):
|
||||
inputs = node.get("inputs", {})
|
||||
for input_name, input_value in list(inputs.items()):
|
||||
if isinstance(input_value, list) and len(input_value) == 2:
|
||||
source_id = str(input_value[0])
|
||||
if source_id not in existing_ids:
|
||||
inputs.pop(input_name, None)
|
||||
|
||||
|
||||
def _get_delegate_master_node_class(class_type):
|
||||
mappings = _get_delegate_master_node_class_mappings()
|
||||
return mappings.get(class_type) if isinstance(mappings, dict) else None
|
||||
|
||||
|
||||
def _normalize_delegate_master_return_type(return_type):
|
||||
if return_type is None:
|
||||
return ""
|
||||
return str(return_type).strip().upper()
|
||||
|
||||
|
||||
def _delegate_master_type_is_safe_scalar(type_name):
|
||||
return type_name in _DELEGATE_MASTER_SAFE_SCALAR_TYPES
|
||||
|
||||
|
||||
def _delegate_master_type_is_safe_config(type_name):
|
||||
return _delegate_master_type_is_safe_scalar(type_name) or type_name in _DELEGATE_MASTER_SAFE_LIST_TYPES
|
||||
|
||||
|
||||
def _delegate_master_output_is_safe_scalar(class_type, output_index):
|
||||
"""Return True when a registered node output is lightweight config data."""
|
||||
if class_type in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES:
|
||||
return True
|
||||
node_class = _get_delegate_master_node_class(class_type)
|
||||
return_types = getattr(node_class, "RETURN_TYPES", ()) if node_class is not None else ()
|
||||
try:
|
||||
output_type = return_types[int(output_index)]
|
||||
except (IndexError, TypeError, ValueError):
|
||||
return False
|
||||
return _delegate_master_type_is_safe_config(_normalize_delegate_master_return_type(output_type))
|
||||
|
||||
|
||||
def _get_delegate_master_input_types(class_type):
|
||||
node_class = _get_delegate_master_node_class(class_type)
|
||||
input_types = getattr(node_class, "INPUT_TYPES", None) if node_class is not None else None
|
||||
if callable(input_types):
|
||||
try:
|
||||
input_types = input_types()
|
||||
except TypeError:
|
||||
return {}
|
||||
return input_types if isinstance(input_types, dict) else {}
|
||||
|
||||
|
||||
def _normalize_delegate_master_input_type(input_spec):
|
||||
if isinstance(input_spec, (list, tuple)) and input_spec:
|
||||
return _normalize_delegate_master_return_type(input_spec[0])
|
||||
return _normalize_delegate_master_return_type(input_spec)
|
||||
|
||||
|
||||
def _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
"""Return True when a registered downstream input expects config data."""
|
||||
for prefix in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES.get(class_type, ()):
|
||||
if input_name.startswith(prefix):
|
||||
def _has_terminal_output_nodes(prompt_obj):
|
||||
"""Return True when the prompt already has at least one terminal output node."""
|
||||
for _node_id, node in _iter_prompt_nodes(prompt_obj):
|
||||
if node.get("class_type") in {"PreviewImage", "SaveImage"}:
|
||||
return True
|
||||
input_types = _get_delegate_master_input_types(class_type)
|
||||
for section_name in ("required", "optional"):
|
||||
section = input_types.get(section_name, {})
|
||||
if isinstance(section, dict) and input_name in section:
|
||||
input_type = _normalize_delegate_master_input_type(section[input_name])
|
||||
return _delegate_master_type_is_safe_config(input_type)
|
||||
return False
|
||||
|
||||
|
||||
def _is_delegate_master_always_retained_upstream_node(node):
|
||||
if not isinstance(node, dict):
|
||||
return False
|
||||
class_type = node.get("class_type")
|
||||
return isinstance(class_type, str) and class_type in _DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES
|
||||
def _add_preview_node(prompt_obj, next_id_fn, source_node_id, slot_index, title_suffix):
|
||||
"""Attach a PreviewImage output node to the given source output slot."""
|
||||
preview_id = next_id_fn()
|
||||
prompt_obj[preview_id] = {
|
||||
"inputs": {
|
||||
"images": [str(source_node_id), int(slot_index)],
|
||||
},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {
|
||||
"title": f"Preview Image ({title_suffix})",
|
||||
},
|
||||
}
|
||||
return preview_id
|
||||
|
||||
|
||||
def _is_delegate_master_retained_upstream_node(node, output_index=0):
|
||||
"""Return True for lightweight upstream nodes safe to keep on the master."""
|
||||
if not isinstance(node, dict):
|
||||
return False
|
||||
class_type = node.get("class_type")
|
||||
if not isinstance(class_type, str):
|
||||
return False
|
||||
return (
|
||||
class_type in _DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES
|
||||
or class_type.startswith("Primitive")
|
||||
or _delegate_master_output_is_safe_scalar(class_type, output_index)
|
||||
)
|
||||
def _ensure_worker_output_node(prompt_obj):
|
||||
"""Guarantee worker prompts retain at least one output node after pruning.
|
||||
|
||||
Branch pruning can remove all explicit output nodes for worker prompts, which
|
||||
causes ComfyUI validation to reject the prompt (`prompt_no_outputs`).
|
||||
"""
|
||||
if _has_terminal_output_nodes(prompt_obj):
|
||||
return prompt_obj
|
||||
|
||||
def _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
node_id,
|
||||
output_index,
|
||||
memo,
|
||||
visiting,
|
||||
):
|
||||
"""Return safe retained branch nodes, or None when the branch is not safe."""
|
||||
node_id = str(node_id)
|
||||
cache_key = (node_id, output_index)
|
||||
if cache_key in memo:
|
||||
cached = memo[cache_key]
|
||||
return None if cached is None else set(cached)
|
||||
if cache_key in visiting:
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
next_id = _create_numeric_id_generator(prompt_obj)
|
||||
|
||||
node = prompt_obj.get(node_id)
|
||||
if not _is_delegate_master_retained_upstream_node(node, output_index):
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
|
||||
visiting.add(cache_key)
|
||||
retained = {node_id}
|
||||
inputs = node.get("inputs", {}) if isinstance(node, dict) else {}
|
||||
class_type = node.get("class_type") if isinstance(node, dict) else None
|
||||
for input_name, value in inputs.items():
|
||||
if not (isinstance(value, list) and len(value) == 2):
|
||||
# Prefer a DistributedBranchCollector output tied to the assigned branch for this worker.
|
||||
for node_id, node in _iter_prompt_nodes(prompt_obj):
|
||||
if node.get("class_type") != "DistributedBranchCollector":
|
||||
continue
|
||||
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
source_id = str(value[0])
|
||||
branch = _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
source_id,
|
||||
value[1],
|
||||
memo,
|
||||
visiting,
|
||||
)
|
||||
if branch is None:
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
retained.update(branch)
|
||||
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = frozenset(retained)
|
||||
return retained
|
||||
|
||||
|
||||
def _find_delegate_master_retained_upstream_nodes(prompt_obj, start_ids):
|
||||
"""Return lightweight upstream nodes needed by kept delegate-master nodes."""
|
||||
connected = set()
|
||||
memo = {}
|
||||
for node_id in start_ids:
|
||||
node = prompt_obj.get(str(node_id)) or {}
|
||||
inputs = node.get("inputs", {})
|
||||
class_type = node.get("class_type") if isinstance(node, dict) else None
|
||||
for input_name, value in inputs.items():
|
||||
if not (isinstance(value, list) and len(value) == 2):
|
||||
continue
|
||||
source_node = prompt_obj.get(str(value[0]))
|
||||
if _is_delegate_master_always_retained_upstream_node(source_node):
|
||||
connected.add(str(value[0]))
|
||||
continue
|
||||
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
continue
|
||||
branch = _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
value[0],
|
||||
value[1],
|
||||
memo,
|
||||
set(),
|
||||
)
|
||||
if branch is not None:
|
||||
connected.update(branch)
|
||||
return connected
|
||||
try:
|
||||
assigned_branch = int(inputs.get("assigned_branch", -1))
|
||||
except (TypeError, ValueError):
|
||||
assigned_branch = -1
|
||||
if assigned_branch >= 0:
|
||||
_add_preview_node(prompt_obj, next_id, node_id, assigned_branch, "auto-added worker output")
|
||||
return prompt_obj
|
||||
|
||||
# Otherwise, anchor to the assigned DistributedBranch slot if available.
|
||||
for node_id, node in _iter_prompt_nodes(prompt_obj):
|
||||
if node.get("class_type") != "DistributedBranch":
|
||||
continue
|
||||
inputs = node.get("inputs", {})
|
||||
try:
|
||||
assigned_branch = int(inputs.get("assigned_branch", -1))
|
||||
except (TypeError, ValueError):
|
||||
assigned_branch = -1
|
||||
if assigned_branch >= 0:
|
||||
_add_preview_node(prompt_obj, next_id, node_id, assigned_branch, "auto-added worker output")
|
||||
return prompt_obj
|
||||
|
||||
# Idle participants (assigned_branch=-1) still need a terminal node to pass
|
||||
# validation; keep this cheap by emitting a tiny synthetic image preview.
|
||||
empty_id = next_id()
|
||||
prompt_obj[empty_id] = {
|
||||
"class_type": "DistributedEmptyImage",
|
||||
"inputs": {
|
||||
"height": 64,
|
||||
"width": 64,
|
||||
"channels": 3,
|
||||
},
|
||||
"_meta": {
|
||||
"title": "Distributed Empty Image (auto-added worker output)",
|
||||
},
|
||||
}
|
||||
_add_preview_node(prompt_obj, next_id, empty_id, 0, "auto-added worker output")
|
||||
return prompt_obj
|
||||
|
||||
|
||||
def prune_prompt_for_worker(prompt_obj):
|
||||
def _prune_worker_downstream_of_branch_collectors(prompt_obj):
|
||||
"""Drop worker-side nodes downstream of branch collectors."""
|
||||
collector_ids = find_nodes_by_class(prompt_obj, "DistributedBranchCollector")
|
||||
if not collector_ids:
|
||||
return prompt_obj
|
||||
|
||||
downstream = _find_downstream_nodes(prompt_obj, collector_ids)
|
||||
for node_id in downstream:
|
||||
if node_id in collector_ids:
|
||||
continue
|
||||
prompt_obj.pop(node_id, None)
|
||||
|
||||
_remove_dangling_input_refs(prompt_obj)
|
||||
return prompt_obj
|
||||
|
||||
|
||||
def prune_prompt_for_branch_worker(
|
||||
prompt_obj: dict[str, Any],
|
||||
branch_node_id: str,
|
||||
assigned_branch: int | list[int] | tuple[int, ...] | set[int],
|
||||
num_branches: int,
|
||||
) -> dict[str, Any]:
|
||||
"""Prune non-assigned branch paths while keeping shared downstream nodes."""
|
||||
branch_id = str(branch_node_id)
|
||||
|
||||
assigned_slots = set()
|
||||
if isinstance(assigned_branch, (list, tuple, set)):
|
||||
for value in assigned_branch:
|
||||
try:
|
||||
idx = int(value)
|
||||
except (TypeError, ValueError):
|
||||
debug_log(f"Prompt transform: invalid assigned branch slot ignored: {value}")
|
||||
continue
|
||||
if idx >= 0:
|
||||
assigned_slots.add(idx)
|
||||
else:
|
||||
try:
|
||||
idx = int(assigned_branch)
|
||||
except (TypeError, ValueError):
|
||||
idx = -1
|
||||
if idx >= 0:
|
||||
assigned_slots.add(idx)
|
||||
|
||||
try:
|
||||
total_slots = int(num_branches)
|
||||
except (TypeError, ValueError):
|
||||
total_slots = 2
|
||||
total_slots = max(2, min(total_slots, 10))
|
||||
|
||||
downstream_by_slot = {}
|
||||
for slot_idx in range(total_slots):
|
||||
downstream = _find_downstream_of_output_slot(prompt_obj, branch_id, slot_idx)
|
||||
downstream.discard(branch_id)
|
||||
downstream_by_slot[slot_idx] = downstream
|
||||
|
||||
keep = {branch_id}
|
||||
for slot_idx in assigned_slots:
|
||||
keep.update(downstream_by_slot.get(slot_idx, set()))
|
||||
|
||||
remove = set()
|
||||
for slot_idx in range(total_slots):
|
||||
if slot_idx in assigned_slots:
|
||||
continue
|
||||
remove.update(downstream_by_slot.get(slot_idx, set()))
|
||||
|
||||
remove -= keep
|
||||
for node_id in remove:
|
||||
prompt_obj.pop(node_id, None)
|
||||
|
||||
_remove_dangling_input_refs(prompt_obj)
|
||||
return prompt_obj
|
||||
|
||||
|
||||
def prune_prompt_for_worker(prompt_obj: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Prune worker prompt to distributed nodes and their upstream dependencies."""
|
||||
collector_ids = find_nodes_by_class(prompt_obj, "DistributedCollector")
|
||||
branch_collector_ids = find_nodes_by_class(prompt_obj, "DistributedBranchCollector")
|
||||
upscale_ids = find_nodes_by_class(prompt_obj, "UltimateSDUpscaleDistributed")
|
||||
distributed_ids = collector_ids + upscale_ids
|
||||
branch_ids = find_nodes_by_class(prompt_obj, "DistributedBranch")
|
||||
distributed_ids = collector_ids + branch_collector_ids + upscale_ids + branch_ids
|
||||
if not distributed_ids:
|
||||
return prompt_obj
|
||||
|
||||
connected = _find_upstream_nodes(prompt_obj, distributed_ids)
|
||||
if branch_ids:
|
||||
connected.update(_find_downstream_nodes(prompt_obj, branch_ids))
|
||||
if branch_collector_ids:
|
||||
downstream_of_collectors = _find_downstream_nodes(prompt_obj, branch_collector_ids)
|
||||
connected -= (downstream_of_collectors - set(branch_collector_ids))
|
||||
|
||||
pruned_prompt = {}
|
||||
for node_id in connected:
|
||||
node = prompt_obj.get(node_id)
|
||||
@@ -349,47 +358,30 @@ def prune_prompt_for_worker(prompt_obj):
|
||||
if dist_id not in pruned_prompt:
|
||||
continue
|
||||
downstream = _find_downstream_nodes(prompt_obj, [dist_id])
|
||||
has_removed_downstream = any(node_id != dist_id for node_id in downstream)
|
||||
has_removed_downstream = any(node_id != dist_id and node_id not in connected for node_id in downstream)
|
||||
if has_removed_downstream:
|
||||
original_node = prompt_obj.get(str(dist_id), {})
|
||||
class_type = original_node.get("class_type")
|
||||
inputs = original_node.get("inputs", {})
|
||||
image_connected = class_type != "DistributedCollector" or (
|
||||
isinstance(inputs.get("images"), list)
|
||||
and len(inputs["images"]) == 2
|
||||
)
|
||||
audio_connected = (
|
||||
class_type == "DistributedCollector"
|
||||
and isinstance(inputs.get("audio"), list)
|
||||
and len(inputs["audio"]) == 2
|
||||
)
|
||||
|
||||
if image_connected:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {"images": [dist_id, 0]},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {"title": "Preview Image (auto-added)"},
|
||||
}
|
||||
elif audio_connected:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {"audio": [dist_id, 1]},
|
||||
"class_type": "PreviewAudio",
|
||||
"_meta": {"title": "Preview Audio (auto-added)"},
|
||||
}
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {
|
||||
"images": [dist_id, 0],
|
||||
},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {
|
||||
"title": "Preview Image (auto-added)",
|
||||
},
|
||||
}
|
||||
|
||||
return pruned_prompt
|
||||
|
||||
|
||||
def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
def prepare_delegate_master_prompt(
|
||||
prompt_obj: dict[str, Any],
|
||||
collector_ids: list[str],
|
||||
) -> dict[str, Any]:
|
||||
"""Prune master prompt so it only executes post-collector nodes in delegate mode."""
|
||||
downstream = _find_downstream_nodes(prompt_obj, collector_ids)
|
||||
nodes_to_keep = set(collector_ids)
|
||||
nodes_to_keep.update(downstream)
|
||||
nodes_to_keep.update(
|
||||
_find_delegate_master_retained_upstream_nodes(prompt_obj, nodes_to_keep)
|
||||
)
|
||||
|
||||
pruned_prompt = {}
|
||||
for node_id in nodes_to_keep:
|
||||
@@ -417,9 +409,16 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
collector_entry = pruned_prompt.get(collector_id)
|
||||
if not collector_entry:
|
||||
continue
|
||||
original_inputs = (prompt_obj.get(collector_id) or {}).get("inputs", {})
|
||||
original_images = original_inputs.get("images")
|
||||
if not (isinstance(original_images, list) and len(original_images) == 2):
|
||||
collector_class = collector_entry.get("class_type")
|
||||
if collector_class == "DistributedBranchCollector":
|
||||
# Branch collector does not require an image input in delegate-only mode.
|
||||
continue
|
||||
if collector_class == "UltimateSDUpscaleDistributed":
|
||||
# Swap to lightweight delegate collector — no model inputs needed.
|
||||
collector_entry["class_type"] = "USDUDelegateCollector"
|
||||
debug_log(
|
||||
f"Swapped USDU node {collector_id} to USDUDelegateCollector for delegate-only master prompt."
|
||||
)
|
||||
continue
|
||||
placeholder_id = next_id()
|
||||
pruned_prompt[placeholder_id] = {
|
||||
@@ -441,32 +440,145 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
return pruned_prompt
|
||||
|
||||
|
||||
def generate_job_id_map(prompt_index, prefix):
|
||||
def generate_job_id_map(prompt_index: PromptIndex, prefix: str) -> dict[str, str]:
|
||||
"""Create stable per-node job IDs for distributed nodes."""
|
||||
job_map = {}
|
||||
distributed_nodes = prompt_index.nodes_for_class("DistributedCollector") + prompt_index.nodes_for_class(
|
||||
"UltimateSDUpscaleDistributed"
|
||||
distributed_nodes = (
|
||||
prompt_index.nodes_for_class("DistributedCollector")
|
||||
+ prompt_index.nodes_for_class("DistributedBranchCollector")
|
||||
+ prompt_index.nodes_for_class("UltimateSDUpscaleDistributed")
|
||||
+ prompt_index.nodes_for_class("DistributedBranch")
|
||||
)
|
||||
for node_id in distributed_nodes:
|
||||
job_map[node_id] = f"{prefix}_{node_id}"
|
||||
return job_map
|
||||
|
||||
|
||||
def _override_seed_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map):
|
||||
"""Configure DistributedSeed nodes for master or worker role."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedSeed"):
|
||||
def _override_simple_nodes(prompt_copy, prompt_index, is_master, participant_id, enabled_json, class_names):
|
||||
"""Configure simple distributed nodes (DistributedSeed, DistributedValue)."""
|
||||
for class_name in class_names:
|
||||
for node_id in prompt_index.nodes_for_class(class_name):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["is_worker"] = not is_master
|
||||
inputs["enabled_worker_ids"] = enabled_json
|
||||
inputs["worker_id"] = "" if is_master else str(participant_id)
|
||||
|
||||
|
||||
def _override_job_nodes(
|
||||
prompt_copy, prompt_index, is_master, participant_id,
|
||||
job_id_map, master_url, enabled_json, delegate_master,
|
||||
class_names, skip_if_upstream=None,
|
||||
):
|
||||
"""Configure job-aware distributed nodes (collector, upscale)."""
|
||||
for class_name in class_names:
|
||||
for node_id in prompt_index.nodes_for_class(class_name):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
if skip_if_upstream and prompt_index.has_upstream(node_id, skip_if_upstream):
|
||||
node.setdefault("inputs", {})["pass_through"] = True
|
||||
continue
|
||||
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
|
||||
inputs["is_worker"] = not is_master
|
||||
inputs["enabled_worker_ids"] = enabled_json
|
||||
if is_master:
|
||||
inputs["delegate_only"] = bool(delegate_master)
|
||||
inputs.pop("master_url", None)
|
||||
inputs.pop("worker_id", None)
|
||||
else:
|
||||
inputs["master_url"] = master_url
|
||||
inputs["worker_id"] = participant_id
|
||||
inputs["delegate_only"] = False
|
||||
|
||||
|
||||
def _override_branch_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
participant_id,
|
||||
enabled_worker_ids,
|
||||
job_id_map,
|
||||
delegate_master,
|
||||
):
|
||||
"""Assign branch slots per participant and prune non-assigned branch paths."""
|
||||
participants = _resolve_participants(enabled_worker_ids, delegate_master)
|
||||
participant_id = str(participant_id)
|
||||
participant_count = len(participants)
|
||||
participant_pos = participants.index(participant_id) if participant_id in participants else None
|
||||
|
||||
for node_id in prompt_index.nodes_for_class("DistributedBranch"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["is_worker"] = not is_master
|
||||
if is_master:
|
||||
inputs["worker_id"] = ""
|
||||
try:
|
||||
num_branches = int(inputs.get("num_branches", 2))
|
||||
except (TypeError, ValueError):
|
||||
num_branches = 2
|
||||
num_branches = max(2, min(num_branches, 10))
|
||||
|
||||
assigned_slots = []
|
||||
if participant_pos is not None and participant_count > 0:
|
||||
assigned_slots = [slot for slot in range(num_branches) if slot % participant_count == participant_pos]
|
||||
|
||||
if len(assigned_slots) == 1:
|
||||
assigned_branch = assigned_slots[0]
|
||||
else:
|
||||
inputs["worker_id"] = f"worker_{worker_index_map.get(participant_id, 0)}"
|
||||
assigned_branch = -1
|
||||
|
||||
inputs["is_worker"] = not is_master
|
||||
inputs["worker_id"] = "" if is_master else participant_id
|
||||
inputs["assigned_branch"] = assigned_branch
|
||||
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
|
||||
|
||||
if is_master and delegate_master:
|
||||
continue
|
||||
|
||||
prompt_copy = prune_prompt_for_branch_worker(
|
||||
prompt_copy,
|
||||
node_id,
|
||||
assigned_slots,
|
||||
num_branches,
|
||||
)
|
||||
|
||||
return prompt_copy
|
||||
|
||||
|
||||
def _override_collector_nodes(
|
||||
def _find_upstream_branch_node(prompt_copy, start_node_id):
|
||||
"""Locate the nearest upstream DistributedBranch node for the provided node."""
|
||||
visited = set()
|
||||
stack = [str(start_node_id)]
|
||||
while stack:
|
||||
node_id = stack.pop()
|
||||
if node_id in visited:
|
||||
continue
|
||||
visited.add(node_id)
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
inputs = node.get("inputs", {})
|
||||
for input_value in inputs.values():
|
||||
if not (isinstance(input_value, list) and len(input_value) == 2):
|
||||
continue
|
||||
source_id = str(input_value[0])
|
||||
source_node = prompt_copy.get(source_id)
|
||||
if not isinstance(source_node, dict):
|
||||
continue
|
||||
if source_node.get("class_type") == "DistributedBranch":
|
||||
return source_id
|
||||
stack.append(source_id)
|
||||
return None
|
||||
|
||||
|
||||
def _override_branch_collector_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
@@ -476,87 +588,77 @@ def _override_collector_nodes(
|
||||
enabled_json,
|
||||
delegate_master,
|
||||
):
|
||||
"""Configure DistributedCollector nodes for master or worker role."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedCollector"):
|
||||
"""Configure DistributedBranchCollector nodes for branch convergence."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedBranchCollector"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
if prompt_index.has_upstream(node_id, "UltimateSDUpscaleDistributed"):
|
||||
node.setdefault("inputs", {})["pass_through"] = True
|
||||
continue
|
||||
upstream_branch_id = _find_upstream_branch_node(prompt_copy, node_id)
|
||||
assigned_branch = -1
|
||||
if upstream_branch_id is not None:
|
||||
branch_node = prompt_copy.get(upstream_branch_id, {})
|
||||
branch_inputs = branch_node.get("inputs", {}) if isinstance(branch_node, dict) else {}
|
||||
try:
|
||||
assigned_branch = int(branch_inputs.get("assigned_branch", -1))
|
||||
except (TypeError, ValueError):
|
||||
assigned_branch = -1
|
||||
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
|
||||
# Group all branch collectors under the same DistributedBranch into one queue.
|
||||
job_source_id = upstream_branch_id or node_id
|
||||
inputs["multi_job_id"] = job_id_map.get(job_source_id, job_source_id)
|
||||
inputs["is_worker"] = not is_master
|
||||
inputs["enabled_worker_ids"] = enabled_json
|
||||
inputs["assigned_branch"] = assigned_branch
|
||||
if is_master:
|
||||
inputs["delegate_only"] = bool(delegate_master)
|
||||
inputs.pop("master_url", None)
|
||||
inputs.pop("worker_id", None)
|
||||
else:
|
||||
inputs["master_url"] = master_url
|
||||
inputs["worker_id"] = participant_id
|
||||
inputs["worker_id"] = str(participant_id)
|
||||
inputs["delegate_only"] = False
|
||||
|
||||
|
||||
def _override_upscale_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
participant_id,
|
||||
job_id_map,
|
||||
master_url,
|
||||
enabled_json,
|
||||
):
|
||||
"""Configure UltimateSDUpscaleDistributed nodes for master or worker role."""
|
||||
for node_id in prompt_index.nodes_for_class("UltimateSDUpscaleDistributed"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
|
||||
inputs["is_worker"] = not is_master
|
||||
inputs["enabled_worker_ids"] = enabled_json
|
||||
if is_master:
|
||||
inputs.pop("master_url", None)
|
||||
inputs.pop("worker_id", None)
|
||||
else:
|
||||
inputs["master_url"] = master_url
|
||||
inputs["worker_id"] = participant_id
|
||||
|
||||
|
||||
def _override_value_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map):
|
||||
"""Configure DistributedValue nodes for master or worker role."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedValue"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
inputs = node.setdefault("inputs", {})
|
||||
inputs["is_worker"] = not is_master
|
||||
if is_master:
|
||||
inputs["worker_id"] = ""
|
||||
else:
|
||||
inputs["worker_id"] = f"worker_{worker_index_map.get(participant_id, 0)}"
|
||||
|
||||
|
||||
def apply_participant_overrides(
|
||||
prompt_copy,
|
||||
participant_id,
|
||||
enabled_worker_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
delegate_master,
|
||||
prompt_index,
|
||||
):
|
||||
prompt_copy: dict[str, Any],
|
||||
participant_id: str,
|
||||
enabled_worker_ids: list[str],
|
||||
job_id_map: dict[str, str],
|
||||
master_url: str,
|
||||
delegate_master: bool,
|
||||
prompt_index: PromptIndex,
|
||||
) -> dict[str, Any]:
|
||||
"""Return a prompt copy with hidden inputs configured for master/worker."""
|
||||
is_master = participant_id == "master"
|
||||
worker_index_map = {wid: idx for idx, wid in enumerate(enabled_worker_ids)}
|
||||
enabled_json = json.dumps(enabled_worker_ids)
|
||||
|
||||
_override_seed_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map)
|
||||
_override_value_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map)
|
||||
_override_collector_nodes(
|
||||
_override_simple_nodes(
|
||||
prompt_copy, prompt_index, is_master, participant_id, enabled_json,
|
||||
["DistributedSeed", "DistributedValue"],
|
||||
)
|
||||
_override_job_nodes(
|
||||
prompt_copy, prompt_index, is_master, participant_id,
|
||||
job_id_map, master_url, enabled_json, delegate_master,
|
||||
["DistributedCollector"],
|
||||
skip_if_upstream="UltimateSDUpscaleDistributed",
|
||||
)
|
||||
_override_job_nodes(
|
||||
prompt_copy, prompt_index, is_master, participant_id,
|
||||
job_id_map, master_url, enabled_json, delegate_master,
|
||||
["UltimateSDUpscaleDistributed"],
|
||||
)
|
||||
prompt_copy = _override_branch_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
participant_id,
|
||||
enabled_worker_ids,
|
||||
job_id_map,
|
||||
delegate_master,
|
||||
)
|
||||
_override_branch_collector_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
@@ -566,14 +668,9 @@ def apply_participant_overrides(
|
||||
enabled_json,
|
||||
delegate_master,
|
||||
)
|
||||
_override_upscale_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
participant_id,
|
||||
job_id_map,
|
||||
master_url,
|
||||
enabled_json,
|
||||
)
|
||||
|
||||
if not is_master:
|
||||
prompt_copy = _prune_worker_downstream_of_branch_collectors(prompt_copy)
|
||||
prompt_copy = _ensure_worker_output_node(prompt_copy)
|
||||
|
||||
return prompt_copy
|
||||
|
||||
+245
-165
@@ -1,8 +1,7 @@
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import server
|
||||
from typing import Any
|
||||
|
||||
from ..utils.async_helpers import queue_prompt_payload
|
||||
from ..utils.config import load_config
|
||||
@@ -13,11 +12,13 @@ from ..utils.constants import (
|
||||
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
|
||||
)
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.network import build_master_url, build_master_callback_url
|
||||
from ..utils.network import build_master_url
|
||||
from ..utils.runtime_state import ensure_distributed_runtime_state, get_prompt_server_instance
|
||||
from ..utils.trace_logger import trace_debug
|
||||
from .schemas import parse_positive_float, parse_positive_int
|
||||
from .schemas import coerce_positive_float, coerce_positive_int
|
||||
from .orchestration.dispatch import (
|
||||
dispatch_worker_prompt,
|
||||
rank_workers_by_load,
|
||||
select_active_workers,
|
||||
select_least_busy_worker,
|
||||
)
|
||||
@@ -31,33 +32,21 @@ from .orchestration.prompt_transform import (
|
||||
prune_prompt_for_worker,
|
||||
)
|
||||
|
||||
|
||||
prompt_server = server.PromptServer.instance
|
||||
|
||||
|
||||
def _generate_execution_trace_id():
|
||||
return f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
|
||||
|
||||
|
||||
def ensure_distributed_state(server_instance=None):
|
||||
"""Ensure prompt_server has the state used by distributed queue orchestration."""
|
||||
ps = server_instance or prompt_server
|
||||
if not hasattr(ps, "distributed_pending_jobs"):
|
||||
ps.distributed_pending_jobs = {}
|
||||
if not hasattr(ps, "distributed_jobs_lock"):
|
||||
ps.distributed_jobs_lock = asyncio.Lock()
|
||||
|
||||
|
||||
# Initialize top-level distributed queue state at module import time.
|
||||
ensure_distributed_state()
|
||||
ensure_distributed_runtime_state(server_instance)
|
||||
|
||||
|
||||
async def _ensure_distributed_queue(job_id):
|
||||
"""Ensure a queue exists for the given distributed job ID."""
|
||||
ensure_distributed_state()
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if job_id not in prompt_server.distributed_pending_jobs:
|
||||
prompt_server.distributed_pending_jobs[job_id] = asyncio.Queue()
|
||||
runtime_state = ensure_distributed_runtime_state()
|
||||
async with runtime_state.distributed_jobs_lock:
|
||||
if job_id not in runtime_state.distributed_pending_jobs:
|
||||
runtime_state.distributed_pending_jobs[job_id] = asyncio.Queue()
|
||||
|
||||
|
||||
def _resolve_enabled_workers(config, requested_ids=None):
|
||||
@@ -96,19 +85,19 @@ def _resolve_enabled_workers(config, requested_ids=None):
|
||||
def _resolve_orchestration_limits(config):
|
||||
"""Resolve bounded concurrency/timeouts for worker preparation pipeline."""
|
||||
settings = (config or {}).get("settings", {}) or {}
|
||||
worker_probe_concurrency = parse_positive_int(
|
||||
worker_probe_concurrency = coerce_positive_int(
|
||||
settings.get("worker_probe_concurrency"),
|
||||
ORCHESTRATION_WORKER_PROBE_CONCURRENCY,
|
||||
)
|
||||
worker_prep_concurrency = parse_positive_int(
|
||||
worker_prep_concurrency = coerce_positive_int(
|
||||
settings.get("worker_prep_concurrency"),
|
||||
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
|
||||
)
|
||||
media_sync_concurrency = parse_positive_int(
|
||||
media_sync_concurrency = coerce_positive_int(
|
||||
settings.get("media_sync_concurrency"),
|
||||
ORCHESTRATION_MEDIA_SYNC_CONCURRENCY,
|
||||
)
|
||||
media_sync_timeout_seconds = parse_positive_float(
|
||||
media_sync_timeout_seconds = coerce_positive_float(
|
||||
settings.get("media_sync_timeout_seconds"),
|
||||
ORCHESTRATION_MEDIA_SYNC_TIMEOUT,
|
||||
)
|
||||
@@ -144,7 +133,6 @@ async def _prepare_worker_payload(
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
config,
|
||||
delegate_master,
|
||||
trace_execution_id,
|
||||
worker_prep_semaphore,
|
||||
@@ -154,11 +142,6 @@ async def _prepare_worker_payload(
|
||||
"""Prepare one worker prompt payload with bounded concurrency and media-sync timeout."""
|
||||
async with worker_prep_semaphore:
|
||||
worker_prompt = prompt_index.copy_prompt()
|
||||
worker_master_url = build_master_callback_url(
|
||||
worker,
|
||||
config=config,
|
||||
prompt_server_instance=prompt_server,
|
||||
)
|
||||
|
||||
worker_type = str(worker.get("type") or "local").strip().lower()
|
||||
is_remote_like = bool(worker.get("host")) and worker_type != "local"
|
||||
@@ -173,7 +156,7 @@ async def _prepare_worker_payload(
|
||||
worker["id"],
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
worker_master_url,
|
||||
master_url,
|
||||
delegate_master,
|
||||
prompt_index,
|
||||
)
|
||||
@@ -197,60 +180,16 @@ async def _prepare_worker_payload(
|
||||
return worker, worker_prompt
|
||||
|
||||
|
||||
async def orchestrate_distributed_execution(
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
enabled_worker_ids=None,
|
||||
delegate_master=None,
|
||||
trace_execution_id=None,
|
||||
async def _select_execution_workers(
|
||||
workers,
|
||||
use_websocket,
|
||||
delegate_master,
|
||||
load_balance_requested,
|
||||
has_branch_nodes,
|
||||
master_url,
|
||||
execution_trace_id,
|
||||
worker_probe_concurrency,
|
||||
):
|
||||
"""Core orchestration logic for the /distributed/queue endpoint.
|
||||
|
||||
Returns:
|
||||
tuple[str, int, int, dict]: (prompt_id, number, worker_count, node_errors)
|
||||
"""
|
||||
ensure_distributed_state()
|
||||
execution_trace_id = trace_execution_id or _generate_execution_trace_id()
|
||||
|
||||
config = load_config()
|
||||
use_websocket = bool(config.get("settings", {}).get("websocket_orchestration", False))
|
||||
master_url = build_master_url(config=config, prompt_server_instance=prompt_server)
|
||||
(
|
||||
worker_probe_concurrency,
|
||||
worker_prep_concurrency,
|
||||
media_sync_concurrency,
|
||||
media_sync_timeout_seconds,
|
||||
) = _resolve_orchestration_limits(config)
|
||||
requested_ids = enabled_worker_ids if enabled_worker_ids is not None else None
|
||||
workers = _resolve_enabled_workers(config, requested_ids)
|
||||
prompt_index = PromptIndex(prompt_obj)
|
||||
load_balance_requested = _prompt_requests_load_balance(prompt_index)
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
(
|
||||
f"Orchestration start: requested_workers={len(workers)}, "
|
||||
f"requested_ids={requested_ids if requested_ids is not None else 'enabled_only'}, "
|
||||
f"websocket={use_websocket}, "
|
||||
f"probe_concurrency={worker_probe_concurrency}, "
|
||||
f"prep_concurrency={worker_prep_concurrency}, "
|
||||
f"media_sync_concurrency={media_sync_concurrency}, "
|
||||
f"media_sync_timeout={media_sync_timeout_seconds:.1f}s, "
|
||||
f"load_balance={load_balance_requested}"
|
||||
),
|
||||
)
|
||||
|
||||
# Respect master delegate-only configuration
|
||||
if delegate_master is None:
|
||||
delegate_master = bool(config.get("settings", {}).get("master_delegate_only", False))
|
||||
|
||||
if not workers and delegate_master:
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
"Delegate-only requested but no workers are enabled. Falling back to master execution.",
|
||||
)
|
||||
delegate_master = False
|
||||
|
||||
active_workers, delegate_master = await select_active_workers(
|
||||
workers,
|
||||
use_websocket,
|
||||
@@ -262,7 +201,6 @@ async def orchestrate_distributed_execution(
|
||||
if load_balance_requested:
|
||||
candidate_workers = list(active_workers)
|
||||
if not delegate_master:
|
||||
# Include master in load balancing only when master participation is enabled.
|
||||
candidate_workers.append(
|
||||
{
|
||||
"id": "master",
|
||||
@@ -287,7 +225,6 @@ async def orchestrate_distributed_execution(
|
||||
selected_worker = candidate_workers[0]
|
||||
|
||||
if selected_worker is not None and str(selected_worker.get("id")) == "master":
|
||||
# Master selected as least busy; run master workload only.
|
||||
active_workers = []
|
||||
delegate_master = False
|
||||
trace_debug(
|
||||
@@ -296,7 +233,6 @@ async def orchestrate_distributed_execution(
|
||||
)
|
||||
elif selected_worker is not None:
|
||||
active_workers = [selected_worker]
|
||||
# Worker selected as least busy; keep master orchestrator-only for this run.
|
||||
delegate_master = True
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
@@ -310,29 +246,47 @@ async def orchestrate_distributed_execution(
|
||||
active_workers = []
|
||||
delegate_master = False
|
||||
|
||||
enabled_ids = [worker["id"] for worker in active_workers]
|
||||
|
||||
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
|
||||
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
|
||||
|
||||
if not job_id_map:
|
||||
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
|
||||
queue_result = await queue_prompt_payload(
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
include_queue_metadata=True,
|
||||
)
|
||||
return (
|
||||
queue_result["prompt_id"],
|
||||
queue_result["number"],
|
||||
0,
|
||||
queue_result.get("node_errors", {}),
|
||||
if has_branch_nodes and len(active_workers) > 1:
|
||||
active_workers = await rank_workers_by_load(
|
||||
active_workers,
|
||||
trace_execution_id=execution_trace_id,
|
||||
probe_concurrency=worker_probe_concurrency,
|
||||
)
|
||||
|
||||
return active_workers, delegate_master
|
||||
|
||||
|
||||
def _distributed_queue_nodes(prompt_index):
|
||||
return (
|
||||
prompt_index.nodes_for_class("DistributedCollector")
|
||||
+ prompt_index.nodes_for_class("DistributedBranchCollector")
|
||||
+ prompt_index.nodes_for_class("DistributedBranch")
|
||||
+ prompt_index.nodes_for_class("UltimateSDUpscaleDistributed")
|
||||
)
|
||||
|
||||
|
||||
def _register_job_allowed_workers(job_id_map, enabled_ids):
|
||||
runtime_state = ensure_distributed_runtime_state()
|
||||
allowed = {str(worker_id) for worker_id in (enabled_ids or []) if str(worker_id).strip()}
|
||||
for job_id in job_id_map.values():
|
||||
await _ensure_distributed_queue(job_id)
|
||||
if job_id:
|
||||
runtime_state.distributed_job_allowed_workers[str(job_id)] = set(allowed)
|
||||
|
||||
|
||||
async def _ensure_job_queues(prompt_index, job_id_map):
|
||||
for node_id in _distributed_queue_nodes(prompt_index):
|
||||
job_id = job_id_map.get(node_id)
|
||||
if job_id:
|
||||
await _ensure_distributed_queue(job_id)
|
||||
|
||||
|
||||
def _build_master_prompt(
|
||||
prompt_index,
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
delegate_master,
|
||||
):
|
||||
master_prompt = prompt_index.copy_prompt()
|
||||
master_prompt = apply_participant_overrides(
|
||||
master_prompt,
|
||||
@@ -344,19 +298,171 @@ async def orchestrate_distributed_execution(
|
||||
prompt_index,
|
||||
)
|
||||
|
||||
if delegate_master:
|
||||
collector_ids = find_nodes_by_class(master_prompt, "DistributedCollector")
|
||||
upscale_nodes = find_nodes_by_class(master_prompt, "UltimateSDUpscaleDistributed")
|
||||
if upscale_nodes:
|
||||
debug_log(
|
||||
"Delegate-only master mode currently does not support UltimateSDUpscaleDistributed nodes; running full prompt on master."
|
||||
if not delegate_master:
|
||||
return master_prompt
|
||||
|
||||
collector_ids = (
|
||||
find_nodes_by_class(master_prompt, "DistributedCollector")
|
||||
+ find_nodes_by_class(master_prompt, "DistributedBranchCollector")
|
||||
)
|
||||
upscale_nodes = find_nodes_by_class(master_prompt, "UltimateSDUpscaleDistributed")
|
||||
# Include USDU nodes as collector-like for delegate pruning
|
||||
collector_ids.extend(upscale_nodes)
|
||||
if not collector_ids:
|
||||
debug_log(
|
||||
"Delegate-only master mode requested but no collector/branch-collector nodes found in master prompt. Running full prompt on master."
|
||||
)
|
||||
return master_prompt
|
||||
return prepare_delegate_master_prompt(master_prompt, collector_ids)
|
||||
|
||||
|
||||
async def _prepare_worker_payloads_for_dispatch(
|
||||
active_workers,
|
||||
prompt_index,
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
delegate_master,
|
||||
execution_trace_id,
|
||||
worker_prep_concurrency,
|
||||
media_sync_concurrency,
|
||||
media_sync_timeout_seconds,
|
||||
):
|
||||
if not active_workers:
|
||||
return []
|
||||
|
||||
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
|
||||
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
|
||||
return await asyncio.gather(
|
||||
*[
|
||||
_prepare_worker_payload(
|
||||
worker,
|
||||
prompt_index,
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
delegate_master,
|
||||
execution_trace_id,
|
||||
worker_prep_semaphore,
|
||||
media_sync_semaphore,
|
||||
media_sync_timeout_seconds,
|
||||
)
|
||||
elif not collector_ids:
|
||||
debug_log(
|
||||
"Delegate-only master mode requested but no collectors found in master prompt. Running full prompt on master."
|
||||
for worker in active_workers
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
async def _dispatch_worker_payloads(
|
||||
worker_payloads,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
use_websocket,
|
||||
execution_trace_id,
|
||||
):
|
||||
if not worker_payloads:
|
||||
return
|
||||
|
||||
await asyncio.gather(
|
||||
*[
|
||||
dispatch_worker_prompt(
|
||||
worker,
|
||||
wprompt,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
use_websocket=use_websocket,
|
||||
trace_execution_id=execution_trace_id,
|
||||
)
|
||||
else:
|
||||
master_prompt = prepare_delegate_master_prompt(master_prompt, collector_ids)
|
||||
for worker, wprompt in worker_payloads
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
async def orchestrate_distributed_execution(
|
||||
prompt_obj: dict[str, Any],
|
||||
workflow_meta: dict[str, Any] | None,
|
||||
client_id: str | None,
|
||||
enabled_worker_ids: list[str] | set[str] | None = None,
|
||||
delegate_master: bool | None = None,
|
||||
trace_execution_id: str | None = None,
|
||||
) -> tuple[str, int]:
|
||||
"""Core orchestration logic for the /distributed/queue endpoint.
|
||||
|
||||
Returns:
|
||||
tuple[str, int]: (prompt_id, worker_count)
|
||||
"""
|
||||
ensure_distributed_state()
|
||||
execution_trace_id = trace_execution_id or _generate_execution_trace_id()
|
||||
|
||||
config = load_config()
|
||||
use_websocket = bool(config.get("settings", {}).get("websocket_orchestration", False))
|
||||
master_url = build_master_url(config=config, prompt_server_instance=get_prompt_server_instance())
|
||||
(
|
||||
worker_probe_concurrency,
|
||||
worker_prep_concurrency,
|
||||
media_sync_concurrency,
|
||||
media_sync_timeout_seconds,
|
||||
) = _resolve_orchestration_limits(config)
|
||||
requested_ids = enabled_worker_ids if enabled_worker_ids is not None else None
|
||||
workers = _resolve_enabled_workers(config, requested_ids)
|
||||
prompt_index = PromptIndex(prompt_obj)
|
||||
load_balance_requested = _prompt_requests_load_balance(prompt_index)
|
||||
has_branch_nodes = bool(prompt_index.nodes_for_class("DistributedBranch"))
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
(
|
||||
f"Orchestration start: requested_workers={len(workers)}, "
|
||||
f"requested_ids={requested_ids if requested_ids is not None else 'enabled_only'}, "
|
||||
f"websocket={use_websocket}, "
|
||||
f"probe_concurrency={worker_probe_concurrency}, "
|
||||
f"prep_concurrency={worker_prep_concurrency}, "
|
||||
f"media_sync_concurrency={media_sync_concurrency}, "
|
||||
f"media_sync_timeout={media_sync_timeout_seconds:.1f}s, "
|
||||
f"load_balance={load_balance_requested}, has_branch_nodes={has_branch_nodes}"
|
||||
),
|
||||
)
|
||||
|
||||
# Respect master delegate-only configuration
|
||||
if delegate_master is None:
|
||||
delegate_master = bool(config.get("settings", {}).get("master_delegate_only", False))
|
||||
|
||||
if not workers and delegate_master:
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
"Delegate-only requested but no workers are enabled. Falling back to master execution.",
|
||||
)
|
||||
delegate_master = False
|
||||
|
||||
active_workers, delegate_master = await _select_execution_workers(
|
||||
workers=workers,
|
||||
use_websocket=use_websocket,
|
||||
delegate_master=delegate_master,
|
||||
load_balance_requested=load_balance_requested,
|
||||
has_branch_nodes=has_branch_nodes,
|
||||
master_url=master_url,
|
||||
execution_trace_id=execution_trace_id,
|
||||
worker_probe_concurrency=worker_probe_concurrency,
|
||||
)
|
||||
|
||||
enabled_ids = [worker["id"] for worker in active_workers]
|
||||
|
||||
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
|
||||
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
|
||||
|
||||
if not job_id_map:
|
||||
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
|
||||
prompt_id = await queue_prompt_payload(prompt_obj, workflow_meta, client_id)
|
||||
return prompt_id, 0
|
||||
|
||||
_register_job_allowed_workers(job_id_map, enabled_ids)
|
||||
await _ensure_job_queues(prompt_index, job_id_map)
|
||||
|
||||
master_prompt = _build_master_prompt(
|
||||
prompt_index=prompt_index,
|
||||
enabled_ids=enabled_ids,
|
||||
job_id_map=job_id_map,
|
||||
master_url=master_url,
|
||||
delegate_master=delegate_master,
|
||||
)
|
||||
|
||||
if active_workers:
|
||||
trace_debug(
|
||||
@@ -364,55 +470,29 @@ async def orchestrate_distributed_execution(
|
||||
"Active distributed workers: "
|
||||
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers),
|
||||
)
|
||||
worker_payloads = []
|
||||
if active_workers:
|
||||
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
|
||||
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
|
||||
worker_payloads = await asyncio.gather(
|
||||
*[
|
||||
_prepare_worker_payload(
|
||||
worker,
|
||||
prompt_index,
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
config,
|
||||
delegate_master,
|
||||
execution_trace_id,
|
||||
worker_prep_semaphore,
|
||||
media_sync_semaphore,
|
||||
media_sync_timeout_seconds,
|
||||
)
|
||||
for worker in active_workers
|
||||
]
|
||||
)
|
||||
|
||||
if worker_payloads:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
dispatch_worker_prompt(
|
||||
worker,
|
||||
wprompt,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
use_websocket=use_websocket,
|
||||
trace_execution_id=execution_trace_id,
|
||||
)
|
||||
for worker, wprompt in worker_payloads
|
||||
]
|
||||
)
|
||||
|
||||
queue_result = await queue_prompt_payload(
|
||||
master_prompt,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
include_queue_metadata=True,
|
||||
worker_payloads = await _prepare_worker_payloads_for_dispatch(
|
||||
active_workers=active_workers,
|
||||
prompt_index=prompt_index,
|
||||
enabled_ids=enabled_ids,
|
||||
job_id_map=job_id_map,
|
||||
master_url=master_url,
|
||||
delegate_master=delegate_master,
|
||||
execution_trace_id=execution_trace_id,
|
||||
worker_prep_concurrency=worker_prep_concurrency,
|
||||
media_sync_concurrency=media_sync_concurrency,
|
||||
media_sync_timeout_seconds=media_sync_timeout_seconds,
|
||||
)
|
||||
prompt_id = queue_result["prompt_id"]
|
||||
prompt_number = queue_result["number"]
|
||||
node_errors = queue_result.get("node_errors", {})
|
||||
await _dispatch_worker_payloads(
|
||||
worker_payloads=worker_payloads,
|
||||
workflow_meta=workflow_meta,
|
||||
client_id=client_id,
|
||||
use_websocket=use_websocket,
|
||||
execution_trace_id=execution_trace_id,
|
||||
)
|
||||
|
||||
prompt_id = await queue_prompt_payload(master_prompt, workflow_meta, client_id)
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
f"Orchestration complete: prompt_id={prompt_id}, dispatched_workers={len(worker_payloads)}, delegate_master={delegate_master}",
|
||||
)
|
||||
return prompt_id, prompt_number, len(worker_payloads), node_errors
|
||||
return prompt_id, len(worker_payloads)
|
||||
|
||||
+20
-15
@@ -1,16 +1,20 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
try:
|
||||
from ..utils.worker_ids import require_enabled_worker_ids
|
||||
except ImportError: # pragma: no cover - supports direct module loading in isolated tests.
|
||||
from utils.worker_ids import require_enabled_worker_ids
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QueueRequestPayload:
|
||||
prompt: Dict[str, Any]
|
||||
workflow_meta: Any
|
||||
prompt: dict[str, Any]
|
||||
workflow_meta: dict[str, Any] | None
|
||||
client_id: str
|
||||
delegate_master: Optional[bool]
|
||||
enabled_worker_ids: List[str]
|
||||
auto_prepare: bool
|
||||
trace_execution_id: Optional[str]
|
||||
delegate_master: bool | None
|
||||
enabled_worker_ids: list[str]
|
||||
trace_execution_id: str | None
|
||||
|
||||
|
||||
def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
@@ -18,11 +22,6 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("Expected a JSON object body")
|
||||
|
||||
auto_prepare_raw = data.get("auto_prepare", True)
|
||||
if not isinstance(auto_prepare_raw, bool):
|
||||
raise ValueError("auto_prepare must be a boolean when provided")
|
||||
auto_prepare = auto_prepare_raw
|
||||
|
||||
prompt = data.get("prompt")
|
||||
# Auto-prepare is always on server-side; keep the field for wire compatibility.
|
||||
if prompt is None:
|
||||
@@ -35,6 +34,10 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
if not isinstance(prompt, dict):
|
||||
raise ValueError("Field 'prompt' must be an object")
|
||||
|
||||
workflow_meta = data.get("workflow")
|
||||
if workflow_meta is not None and not isinstance(workflow_meta, dict):
|
||||
raise ValueError("Field 'workflow' must be an object when provided")
|
||||
|
||||
enabled_ids_raw = data.get("enabled_worker_ids")
|
||||
workers_field = data.get("workers")
|
||||
if enabled_ids_raw is None and workers_field is not None:
|
||||
@@ -51,7 +54,10 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
else:
|
||||
if not isinstance(enabled_ids_raw, list):
|
||||
raise ValueError("enabled_worker_ids must be a list of worker IDs")
|
||||
enabled_ids = [str(worker_id).strip() for worker_id in enabled_ids_raw if str(worker_id).strip()]
|
||||
try:
|
||||
enabled_ids = require_enabled_worker_ids(enabled_ids_raw)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"enabled_worker_ids invalid: {exc}") from exc
|
||||
|
||||
delegate_master = data.get("delegate_master")
|
||||
if delegate_master is not None and not isinstance(delegate_master, bool):
|
||||
@@ -70,10 +76,9 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
|
||||
return QueueRequestPayload(
|
||||
prompt=prompt,
|
||||
workflow_meta=data.get("workflow"),
|
||||
workflow_meta=workflow_meta,
|
||||
client_id=client_id,
|
||||
delegate_master=delegate_master,
|
||||
enabled_worker_ids=enabled_ids,
|
||||
auto_prepare=auto_prepare,
|
||||
trace_execution_id=trace_execution_id,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Shared request guards for distributed API endpoints."""
|
||||
from __future__ import annotations
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from ..utils.config import load_config
|
||||
from ..utils.network import handle_api_error
|
||||
from .schemas import is_authorized_request
|
||||
|
||||
|
||||
async def authorization_error_or_none(request: web.Request) -> web.StreamResponse | None:
|
||||
"""Return a 403 response when distributed API auth fails, else None."""
|
||||
if is_authorized_request(request, load_config()):
|
||||
return None
|
||||
return await handle_api_error(request, "Unauthorized", 403)
|
||||
+61
-15
@@ -1,3 +1,29 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..utils.auth import (
|
||||
AUTH_HEADER_NAME,
|
||||
distributed_auth_headers,
|
||||
is_authorized_request,
|
||||
)
|
||||
from ..utils.parsing import coerce_bool, require_bool_literal
|
||||
|
||||
__all__ = [
|
||||
"AUTH_HEADER_NAME",
|
||||
"coerce_bool",
|
||||
"coerce_positive_float",
|
||||
"coerce_positive_int",
|
||||
"distributed_auth_headers",
|
||||
"is_authorized_request",
|
||||
"require_bool_literal",
|
||||
"require_fields",
|
||||
"require_worker_id",
|
||||
"require_positive_int",
|
||||
"validate_worker_id",
|
||||
]
|
||||
|
||||
|
||||
def require_fields(data: dict, *fields) -> list[str]:
|
||||
"""Return field names that are missing or empty in a JSON object."""
|
||||
if not isinstance(data, dict):
|
||||
@@ -18,26 +44,43 @@ def require_fields(data: dict, *fields) -> list[str]:
|
||||
return missing
|
||||
|
||||
|
||||
def validate_worker_id(worker_id: str, config: dict) -> bool:
|
||||
def validate_worker_id(worker_id: str, config: dict[str, Any]) -> bool:
|
||||
"""Return True when worker_id exists in config['workers']."""
|
||||
worker_id_str = str(worker_id)
|
||||
try:
|
||||
require_worker_id(worker_id, config)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def require_worker_id(worker_id: str, config: dict[str, Any], field_name: str = "worker_id") -> str:
|
||||
"""Strict worker-id validator for request boundaries."""
|
||||
worker_id_str = str(worker_id).strip()
|
||||
if not worker_id_str:
|
||||
raise ValueError(f"Field '{field_name}' must be a non-empty worker id.")
|
||||
|
||||
workers = (config or {}).get("workers", [])
|
||||
return any(str(worker.get("id")) == worker_id_str for worker in workers)
|
||||
for worker in workers:
|
||||
if not isinstance(worker, dict):
|
||||
continue
|
||||
if str(worker.get("id")).strip() == worker_id_str:
|
||||
return worker_id_str
|
||||
raise ValueError(f"Worker {worker_id_str} not found.")
|
||||
|
||||
|
||||
def validate_positive_int(value, field_name: str) -> str | None:
|
||||
"""Validate positive integers and return an error string when invalid."""
|
||||
def require_positive_int(value, field_name: str = "value") -> int:
|
||||
"""Strict positive-int parser for request boundary validation."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return f"Field '{field_name}' must be a positive integer."
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"Field '{field_name}' must be a positive integer.") from exc
|
||||
if parsed <= 0:
|
||||
return f"Field '{field_name}' must be a positive integer."
|
||||
return None
|
||||
raise ValueError(f"Field '{field_name}' must be a positive integer.")
|
||||
return parsed
|
||||
|
||||
|
||||
def parse_positive_int(value, default: int) -> int:
|
||||
"""Parse value as positive int, returning default on failure."""
|
||||
def coerce_positive_int(value, default: int) -> int:
|
||||
"""Permissive positive-int coercion used for non-boundary defaults."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError):
|
||||
@@ -45,10 +88,13 @@ def parse_positive_int(value, default: int) -> int:
|
||||
return max(1, parsed)
|
||||
|
||||
|
||||
def parse_positive_float(value, default: float) -> float:
|
||||
"""Parse value as positive float, returning default on failure."""
|
||||
def coerce_positive_float(value, default: float) -> float:
|
||||
"""Permissive positive-float coercion used for non-boundary defaults."""
|
||||
fallback = float(default)
|
||||
if fallback <= 0.0:
|
||||
fallback = 0.1
|
||||
try:
|
||||
parsed = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return max(0.0, float(default))
|
||||
return max(0.0, parsed)
|
||||
return fallback
|
||||
return parsed if parsed > 0.0 else fallback
|
||||
|
||||
+41
-29
@@ -1,51 +1,63 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from aiohttp import web
|
||||
import server
|
||||
|
||||
from .endpoint_policy import run_authorized_endpoint
|
||||
from ..utils.cloudflare import cloudflare_tunnel_manager
|
||||
from ..utils.config import load_config
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.network import handle_api_error
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/tunnel/status")
|
||||
async def tunnel_status_endpoint(request):
|
||||
async def tunnel_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Return Cloudflare tunnel status and last known details."""
|
||||
try:
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
status = cloudflare_tunnel_manager.get_status()
|
||||
config = load_config()
|
||||
master_host = (config.get("master") or {}).get("host")
|
||||
return web.json_response({
|
||||
"status": "success",
|
||||
"tunnel": status,
|
||||
"master_host": master_host
|
||||
})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "success",
|
||||
"tunnel": status,
|
||||
"master_host": master_host,
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/tunnel/start")
|
||||
async def tunnel_start_endpoint(request):
|
||||
async def tunnel_start_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Start a Cloudflare tunnel pointing at the current ComfyUI server."""
|
||||
try:
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
result = await cloudflare_tunnel_manager.start_tunnel()
|
||||
config = load_config()
|
||||
return web.json_response({
|
||||
"status": "success",
|
||||
"tunnel": result,
|
||||
"master_host": (config.get("master") or {}).get("host")
|
||||
})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "success",
|
||||
"tunnel": result,
|
||||
"master_host": (config.get("master") or {}).get("host"),
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/tunnel/stop")
|
||||
async def tunnel_stop_endpoint(request):
|
||||
async def tunnel_stop_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Stop the managed Cloudflare tunnel if running."""
|
||||
try:
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
result = await cloudflare_tunnel_manager.stop_tunnel()
|
||||
config = load_config()
|
||||
return web.json_response({
|
||||
"status": "success",
|
||||
"tunnel": result,
|
||||
"master_host": (config.get("master") or {}).get("host")
|
||||
})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
return web.json_response(
|
||||
{
|
||||
"status": "success",
|
||||
"tunnel": result,
|
||||
"master_host": (config.get("master") or {}).get("host"),
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
+137
-21
@@ -1,3 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import io
|
||||
import time
|
||||
@@ -6,15 +8,78 @@ from aiohttp import web
|
||||
from PIL import Image
|
||||
import server
|
||||
|
||||
from .request_guards import authorization_error_or_none
|
||||
from .schemas import require_bool_literal
|
||||
from ..upscale.job_models import BaseJobState, ImageJobState, TileJobState
|
||||
from ..upscale.job_store import MAX_PAYLOAD_SIZE, ensure_tile_jobs_initialized
|
||||
from ..upscale.payload_parsers import _parse_tiles_from_form
|
||||
from ..upscale.job_store import ensure_tile_jobs_initialized, init_dynamic_job
|
||||
from ..upscale.payload_parsers import parse_tiles_from_form
|
||||
from ..utils.logging import debug_log
|
||||
from ..utils.network import handle_api_error
|
||||
from ..utils.usdu_management import MAX_PAYLOAD_SIZE
|
||||
|
||||
|
||||
def _parse_int_field(value, field_name: str, *, minimum: int | None = None) -> int:
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"{field_name} must be an integer") from exc
|
||||
if minimum is not None and parsed < minimum:
|
||||
raise ValueError(f"{field_name} must be >= {minimum}")
|
||||
return parsed
|
||||
|
||||
|
||||
def _parse_bool_field(
|
||||
value,
|
||||
field_name: str,
|
||||
*,
|
||||
default: bool = False,
|
||||
) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
return require_bool_literal(value, field_name=field_name)
|
||||
|
||||
|
||||
def _is_worker_allowed(job_data: BaseJobState, worker_id: str) -> bool:
|
||||
known_workers = {str(wid) for wid in getattr(job_data, "worker_status", {}).keys()}
|
||||
known_workers.update(str(wid) for wid in getattr(job_data, "assigned_to_workers", {}).keys())
|
||||
if not known_workers:
|
||||
return True
|
||||
return str(worker_id) in known_workers
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/init_dynamic_job")
|
||||
async def init_dynamic_job_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Allow workers to initialize a dynamic job queue on the master.
|
||||
|
||||
Called by the first worker to reach the USDU node in delegate-only mode.
|
||||
Idempotent — subsequent calls are no-ops if the queue already exists.
|
||||
"""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
multi_job_id = data.get("multi_job_id")
|
||||
if not multi_job_id:
|
||||
return await handle_api_error(request, "Missing multi_job_id", 400)
|
||||
|
||||
batch_size = _parse_int_field(data.get("batch_size"), "batch_size", minimum=1)
|
||||
enabled_workers = data.get("enabled_workers") or []
|
||||
if isinstance(enabled_workers, str):
|
||||
import json as _json
|
||||
enabled_workers = _json.loads(enabled_workers)
|
||||
|
||||
await init_dynamic_job(multi_job_id, batch_size, enabled_workers)
|
||||
return web.json_response({"status": "success"})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/heartbeat")
|
||||
async def heartbeat_endpoint(request):
|
||||
async def heartbeat_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
worker_id = data.get('worker_id')
|
||||
@@ -28,6 +93,8 @@ async def heartbeat_endpoint(request):
|
||||
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if isinstance(job_data, BaseJobState):
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
job_data.worker_status[worker_id] = time.time()
|
||||
debug_log(f"Heartbeat from worker {worker_id}")
|
||||
return web.json_response({"status": "success"})
|
||||
@@ -38,24 +105,38 @@ async def heartbeat_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/submit_tiles")
|
||||
async def submit_tiles_endpoint(request):
|
||||
async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Endpoint for workers to submit processed tiles in static mode."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
content_length = request.headers.get('content-length')
|
||||
if content_length and int(content_length) > MAX_PAYLOAD_SIZE:
|
||||
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
|
||||
if content_length:
|
||||
try:
|
||||
payload_size = int(content_length)
|
||||
except ValueError:
|
||||
return await handle_api_error(request, "Invalid content-length header", 400)
|
||||
if payload_size > MAX_PAYLOAD_SIZE:
|
||||
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
|
||||
|
||||
data = await request.post()
|
||||
multi_job_id = data.get('multi_job_id')
|
||||
worker_id = data.get('worker_id')
|
||||
is_last = data.get('is_last', 'False').lower() == 'true'
|
||||
try:
|
||||
is_last = _parse_bool_field(data.get('is_last', False), "is_last", default=False)
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, str(e), 400)
|
||||
|
||||
if multi_job_id is None or worker_id is None:
|
||||
return await handle_api_error(request, "Missing multi_job_id or worker_id", 400)
|
||||
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
|
||||
batch_size = int(data.get('batch_size', 0))
|
||||
try:
|
||||
batch_size = _parse_int_field(data.get('batch_size', 0), "batch_size", minimum=0)
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, str(e), 400)
|
||||
|
||||
# Handle completion signal
|
||||
if batch_size == 0 and is_last:
|
||||
@@ -64,6 +145,8 @@ async def submit_tiles_endpoint(request):
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, TileJobState):
|
||||
return await handle_api_error(request, "Job not configured for tile submissions", 400)
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
await job_data.queue.put({
|
||||
'worker_id': worker_id,
|
||||
'is_last': True,
|
||||
@@ -73,7 +156,7 @@ async def submit_tiles_endpoint(request):
|
||||
return web.json_response({"status": "success"})
|
||||
|
||||
try:
|
||||
tiles = _parse_tiles_from_form(data)
|
||||
tiles = parse_tiles_from_form(data)
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, str(e), 400)
|
||||
|
||||
@@ -83,6 +166,8 @@ async def submit_tiles_endpoint(request):
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, TileJobState):
|
||||
return await handle_api_error(request, "Job not configured for tile submissions", 400)
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
|
||||
q = job_data.queue
|
||||
if batch_size > 0 or len(tiles) > 0:
|
||||
@@ -106,17 +191,28 @@ async def submit_tiles_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/submit_image")
|
||||
async def submit_image_endpoint(request):
|
||||
async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Endpoint for workers to submit processed images in dynamic mode."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
content_length = request.headers.get('content-length')
|
||||
if content_length and int(content_length) > MAX_PAYLOAD_SIZE:
|
||||
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
|
||||
if content_length:
|
||||
try:
|
||||
payload_size = int(content_length)
|
||||
except ValueError:
|
||||
return await handle_api_error(request, "Invalid content-length header", 400)
|
||||
if payload_size > MAX_PAYLOAD_SIZE:
|
||||
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
|
||||
|
||||
data = await request.post()
|
||||
multi_job_id = data.get('multi_job_id')
|
||||
worker_id = data.get('worker_id')
|
||||
is_last = data.get('is_last', 'False').lower() == 'true'
|
||||
try:
|
||||
is_last = _parse_bool_field(data.get('is_last', False), "is_last", default=False)
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, str(e), 400)
|
||||
|
||||
if multi_job_id is None or worker_id is None:
|
||||
return await handle_api_error(request, "Missing multi_job_id or worker_id", 400)
|
||||
@@ -125,7 +221,10 @@ async def submit_image_endpoint(request):
|
||||
|
||||
# Handle image submission
|
||||
if 'full_image' in data and 'image_idx' in data:
|
||||
image_idx = int(data.get('image_idx'))
|
||||
try:
|
||||
image_idx = _parse_int_field(data.get('image_idx'), "image_idx", minimum=0)
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, str(e), 400)
|
||||
img_data = data['full_image'].file.read()
|
||||
img = Image.open(io.BytesIO(img_data)).convert("RGB")
|
||||
|
||||
@@ -136,6 +235,8 @@ async def submit_image_endpoint(request):
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, ImageJobState):
|
||||
return await handle_api_error(request, "Job not configured for image submissions", 400)
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
await job_data.queue.put({
|
||||
'worker_id': worker_id,
|
||||
'image_idx': image_idx,
|
||||
@@ -151,6 +252,8 @@ async def submit_image_endpoint(request):
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, ImageJobState):
|
||||
return await handle_api_error(request, "Job not configured for image submissions", 400)
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
await job_data.queue.put({
|
||||
'worker_id': worker_id,
|
||||
'is_last': True,
|
||||
@@ -166,8 +269,11 @@ async def submit_image_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/request_image")
|
||||
async def request_image_endpoint(request):
|
||||
async def request_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Endpoint for workers to request tasks (images in dynamic mode, tiles in static mode)."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
data = await request.json()
|
||||
worker_id = data.get('worker_id')
|
||||
@@ -190,6 +296,8 @@ async def request_image_endpoint(request):
|
||||
pending_queue = job_data.pending_tasks
|
||||
else:
|
||||
return await handle_api_error(request, "Invalid job configuration", 400)
|
||||
if not _is_worker_allowed(job_data, worker_id):
|
||||
return await handle_api_error(request, "Unauthorized worker_id", 403)
|
||||
|
||||
try:
|
||||
task_idx = await asyncio.wait_for(pending_queue.get(), timeout=0.1)
|
||||
@@ -199,25 +307,33 @@ async def request_image_endpoint(request):
|
||||
|
||||
if mode == 'dynamic':
|
||||
debug_log(f"UltimateSDUpscale API - Assigned image {task_idx} to worker {worker_id}")
|
||||
return web.json_response({"image_idx": task_idx, "estimated_remaining": remaining})
|
||||
return web.json_response(
|
||||
{
|
||||
"kind": "image",
|
||||
"task_idx": task_idx,
|
||||
"estimated_remaining": remaining,
|
||||
}
|
||||
)
|
||||
debug_log(f"UltimateSDUpscale API - Assigned tile {task_idx} to worker {worker_id}")
|
||||
return web.json_response({
|
||||
"tile_idx": task_idx,
|
||||
"kind": "tile",
|
||||
"task_idx": task_idx,
|
||||
"estimated_remaining": remaining,
|
||||
"batched_static": job_data.batched_static,
|
||||
})
|
||||
except asyncio.TimeoutError:
|
||||
if mode == 'dynamic':
|
||||
return web.json_response({"image_idx": None})
|
||||
return web.json_response({"tile_idx": None})
|
||||
return web.json_response({"kind": "none", "task_idx": None, "estimated_remaining": 0})
|
||||
return await handle_api_error(request, "Job not found", 404)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/job_status")
|
||||
async def job_status_endpoint(request):
|
||||
async def job_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Endpoint to check if a job is ready."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
multi_job_id = request.query.get('multi_job_id')
|
||||
if not multi_job_id:
|
||||
return web.json_response({"ready": False})
|
||||
|
||||
+109
-61
@@ -1,10 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
import platform
|
||||
import subprocess
|
||||
import subprocess # nosec B404 - commands are fixed and never shell-expanded
|
||||
import socket
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import aiohttp
|
||||
@@ -16,35 +19,33 @@ from ..utils.logging import debug_log, log
|
||||
from ..utils.network import (
|
||||
build_worker_url,
|
||||
get_client_session,
|
||||
get_server_port,
|
||||
handle_api_error,
|
||||
normalize_host,
|
||||
probe_worker,
|
||||
)
|
||||
from ..utils.constants import CHUNK_SIZE
|
||||
from ..workers import get_worker_manager
|
||||
from ..workers.ports import allocate_worker_ports
|
||||
from .schemas import require_fields, validate_worker_id
|
||||
from .request_guards import authorization_error_or_none
|
||||
from .schemas import (
|
||||
distributed_auth_headers,
|
||||
require_fields,
|
||||
require_worker_id,
|
||||
)
|
||||
from ..workers.detection import (
|
||||
get_machine_id,
|
||||
is_docker_environment,
|
||||
is_runpod_environment,
|
||||
)
|
||||
try:
|
||||
from ..utils.async_helpers import PromptValidationError, queue_prompt_payload
|
||||
except ImportError:
|
||||
from ..utils.async_helpers import queue_prompt_payload
|
||||
|
||||
class PromptValidationError(RuntimeError):
|
||||
def __init__(self, message, validation_error=None, node_errors=None):
|
||||
super().__init__(str(message))
|
||||
self.validation_error = validation_error if isinstance(validation_error, dict) else {}
|
||||
self.node_errors = node_errors if isinstance(node_errors, dict) else {}
|
||||
from ..utils.async_helpers import PromptValidationError, queue_prompt_payload
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/worker_ws")
|
||||
async def worker_ws_endpoint(request):
|
||||
async def worker_ws_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""WebSocket endpoint for worker prompt dispatch."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
|
||||
ws = web.WebSocketResponse(heartbeat=30)
|
||||
await ws.prepare(request)
|
||||
|
||||
@@ -115,8 +116,11 @@ async def worker_ws_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/worker/clear_launching")
|
||||
async def clear_launching_state(request):
|
||||
async def clear_launching_state(request: web.Request) -> web.StreamResponse:
|
||||
"""Clear the launching flag when worker is confirmed running."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
wm = get_worker_manager()
|
||||
data = await request.json()
|
||||
@@ -126,8 +130,10 @@ async def clear_launching_state(request):
|
||||
|
||||
worker_id = str(data.get("worker_id")).strip()
|
||||
config = load_config()
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
|
||||
# Clear launching flag in managed processes
|
||||
if worker_id in wm.processes:
|
||||
@@ -141,8 +147,9 @@ async def clear_launching_state(request):
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
|
||||
def get_network_ips():
|
||||
def get_network_ips() -> list[str]:
|
||||
"""Get all network IPs, trying multiple methods."""
|
||||
command_timeout = 5.0
|
||||
ips = []
|
||||
hostname = socket.gethostname()
|
||||
|
||||
@@ -153,8 +160,8 @@ def get_network_ips():
|
||||
ip = info[4][0]
|
||||
if ip and ip not in ips and not ip.startswith('::'): # Skip IPv6 for now
|
||||
ips.append(ip)
|
||||
except (socket.gaierror, OSError):
|
||||
pass
|
||||
except (socket.gaierror, OSError) as exc:
|
||||
debug_log(f"get_network_ips: getaddrinfo failed for hostname {hostname}: {exc}")
|
||||
|
||||
# Method 2: Try to connect to external server and get local IP
|
||||
try:
|
||||
@@ -164,14 +171,19 @@ def get_network_ips():
|
||||
s.close()
|
||||
if local_ip not in ips:
|
||||
ips.append(local_ip)
|
||||
except (OSError, socket.error):
|
||||
pass
|
||||
except (OSError, socket.error) as exc:
|
||||
debug_log(f"get_network_ips: UDP local IP probe failed: {exc}")
|
||||
|
||||
# Method 3: Platform-specific commands
|
||||
try:
|
||||
if platform.system() == "Windows":
|
||||
# Windows ipconfig
|
||||
result = subprocess.run(["ipconfig"], capture_output=True, text=True)
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ipconfig"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
lines = result.stdout.split('\n')
|
||||
for i, line in enumerate(lines):
|
||||
if 'IPv4' in line and i + 1 < len(lines):
|
||||
@@ -181,10 +193,20 @@ def get_network_ips():
|
||||
else:
|
||||
# Unix/Linux/Mac ifconfig or ip addr
|
||||
try:
|
||||
result = subprocess.run(["ip", "addr"], capture_output=True, text=True)
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ip", "addr"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
except (FileNotFoundError, OSError):
|
||||
try:
|
||||
result = subprocess.run(["ifconfig"], capture_output=True, text=True)
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ifconfig"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
except (FileNotFoundError, OSError):
|
||||
result = None
|
||||
|
||||
@@ -195,13 +217,13 @@ def get_network_ips():
|
||||
ip = match.group(1)
|
||||
if ip and ip not in ips:
|
||||
ips.append(ip)
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
pass
|
||||
except (OSError, subprocess.SubprocessError) as exc:
|
||||
debug_log(f"get_network_ips: platform command probe failed: {exc}")
|
||||
|
||||
return ips
|
||||
|
||||
|
||||
def get_recommended_ip(ips):
|
||||
def get_recommended_ip(ips: list[str]) -> str | None:
|
||||
"""Choose the best IP for master-worker communication."""
|
||||
# Priority order:
|
||||
# 1. Private network ranges (192.168.x.x, 10.x.x.x, 172.16-31.x.x)
|
||||
@@ -236,7 +258,7 @@ def get_recommended_ip(ips):
|
||||
return None
|
||||
|
||||
|
||||
def _get_cuda_info():
|
||||
def _get_cuda_info() -> tuple[int | None, int, int]:
|
||||
"""Detect CUDA device index and total physical GPU count.
|
||||
|
||||
Returns (cuda_device, cuda_device_count, physical_device_count).
|
||||
@@ -252,7 +274,7 @@ def _get_cuda_info():
|
||||
if visible_devices:
|
||||
cuda_device = visible_devices[0]
|
||||
try:
|
||||
result = subprocess.run(
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
['nvidia-smi', '--query-gpu=name', '--format=csv,noheader'],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
@@ -276,16 +298,9 @@ def _get_cuda_info():
|
||||
return None, 0, 0
|
||||
|
||||
|
||||
def _collect_network_info_sync():
|
||||
def _collect_network_info_sync() -> dict[str, Any]:
|
||||
"""Collect network/cuda info in a worker thread to avoid blocking route handlers."""
|
||||
cuda_device, cuda_device_count, physical_device_count = _get_cuda_info()
|
||||
device_count = physical_device_count if physical_device_count > 0 else cuda_device_count
|
||||
master_port = get_server_port()
|
||||
config = load_config()
|
||||
worker_ports = []
|
||||
if not config.get("settings", {}).get("has_auto_populated_workers") and not config.get("workers"):
|
||||
worker_count = device_count - (1 if cuda_device is not None and 0 <= cuda_device < device_count else 0)
|
||||
worker_ports = allocate_worker_ports(master_port, config.get("workers", []), worker_count)
|
||||
hostname = socket.gethostname()
|
||||
all_ips = get_network_ips()
|
||||
recommended_ip = get_recommended_ip(all_ips)
|
||||
@@ -294,13 +309,11 @@ def _collect_network_info_sync():
|
||||
"all_ips": all_ips,
|
||||
"recommended_ip": recommended_ip,
|
||||
"cuda_device": cuda_device,
|
||||
"cuda_device_count": device_count,
|
||||
"master_port": master_port,
|
||||
"local_worker_ports": worker_ports,
|
||||
"cuda_device_count": physical_device_count if physical_device_count > 0 else cuda_device_count,
|
||||
}
|
||||
|
||||
|
||||
def _read_worker_log_sync(log_file, lines_to_read):
|
||||
def _read_worker_log_sync(log_file: str, lines_to_read: int) -> dict[str, Any]:
|
||||
"""Read worker log content from disk in a threadpool worker."""
|
||||
file_size = os.path.getsize(log_file)
|
||||
|
||||
@@ -336,7 +349,12 @@ def _read_worker_log_sync(log_file, lines_to_read):
|
||||
}
|
||||
|
||||
|
||||
def _parse_positive_int_query(value, default, minimum=1, maximum=10000):
|
||||
def _parse_positive_int_query(
|
||||
value: Any,
|
||||
default: int,
|
||||
minimum: int = 1,
|
||||
maximum: int | None = 10000,
|
||||
) -> int:
|
||||
"""Parse bounded positive integer query params with sane fallback."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
@@ -348,7 +366,7 @@ def _parse_positive_int_query(value, default, minimum=1, maximum=10000):
|
||||
return parsed
|
||||
|
||||
|
||||
def _find_worker_by_id(config, worker_id):
|
||||
def _find_worker_by_id(config: dict[str, Any], worker_id: str) -> dict[str, Any] | None:
|
||||
worker_id_str = str(worker_id).strip()
|
||||
for worker in config.get("workers", []):
|
||||
if str(worker.get("id")).strip() == worker_id_str:
|
||||
@@ -357,8 +375,11 @@ def _find_worker_by_id(config, worker_id):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/local_log")
|
||||
async def get_local_log_endpoint(request):
|
||||
async def get_local_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Return this instance's in-memory ComfyUI log buffer."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
from app.logger import get_logs
|
||||
except Exception as e:
|
||||
@@ -402,8 +423,11 @@ async def get_local_log_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/network_info")
|
||||
async def get_network_info_endpoint(request):
|
||||
async def get_network_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Get network interfaces and recommend best IP for master."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
info = await loop.run_in_executor(None, _collect_network_info_sync)
|
||||
@@ -417,8 +441,11 @@ async def get_network_info_endpoint(request):
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/system_info")
|
||||
async def get_system_info_endpoint(request):
|
||||
async def get_system_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Get system information including machine ID for local worker detection."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
import socket
|
||||
|
||||
@@ -441,8 +468,11 @@ async def get_system_info_endpoint(request):
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/launch_worker")
|
||||
async def launch_worker_endpoint(request):
|
||||
async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Launch a worker process from the UI."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
wm = get_worker_manager()
|
||||
data = await request.json()
|
||||
@@ -454,8 +484,10 @@ async def launch_worker_endpoint(request):
|
||||
|
||||
# Find worker config
|
||||
config = load_config()
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
worker = next((w for w in config.get("workers", []) if str(w.get("id")) == worker_id), None)
|
||||
if not worker:
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
@@ -474,7 +506,7 @@ async def launch_worker_endpoint(request):
|
||||
is_running = process.poll() is None
|
||||
else:
|
||||
# Restored process without subprocess object
|
||||
is_running = wm._is_process_running(proc_info['pid'])
|
||||
is_running = wm.is_process_running(proc_info['pid'])
|
||||
|
||||
if is_running:
|
||||
return await handle_api_error(request, "Worker already running (managed by UI)", 409)
|
||||
@@ -494,8 +526,6 @@ async def launch_worker_endpoint(request):
|
||||
"message": f"Worker {worker['name']} launched",
|
||||
"log_file": log_file
|
||||
})
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 400)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 500)
|
||||
|
||||
@@ -504,8 +534,11 @@ async def launch_worker_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/stop_worker")
|
||||
async def stop_worker_endpoint(request):
|
||||
async def stop_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Stop a worker process that was launched from the UI."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
wm = get_worker_manager()
|
||||
data = await request.json()
|
||||
@@ -515,8 +548,10 @@ async def stop_worker_endpoint(request):
|
||||
|
||||
worker_id = str(data.get("worker_id")).strip()
|
||||
config = load_config()
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
|
||||
success, message = wm.stop_worker(worker_id)
|
||||
|
||||
@@ -534,8 +569,11 @@ async def stop_worker_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/managed_workers")
|
||||
async def get_managed_workers_endpoint(request):
|
||||
async def get_managed_workers_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Get list of workers managed by this UI instance."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
managed = get_worker_manager().get_managed_workers()
|
||||
return web.json_response({
|
||||
@@ -547,8 +585,11 @@ async def get_managed_workers_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/local-worker-status")
|
||||
async def get_local_worker_status_endpoint(request):
|
||||
async def get_local_worker_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Check status of all local workers (localhost/no host specified)."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
config = load_config()
|
||||
worker_statuses = {}
|
||||
@@ -617,8 +658,11 @@ async def get_local_worker_status_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/worker_log/{worker_id}")
|
||||
async def get_worker_log_endpoint(request):
|
||||
async def get_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Get log content for a specific worker."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
wm = get_worker_manager()
|
||||
worker_id = request.match_info['worker_id']
|
||||
@@ -660,8 +704,11 @@ async def get_worker_log_endpoint(request):
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/remote_worker_log/{worker_id}")
|
||||
async def get_remote_worker_log_endpoint(request):
|
||||
async def get_remote_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"""Proxy a remote worker log request to the worker's local in-memory log endpoint."""
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
try:
|
||||
worker_id = str(request.match_info["worker_id"]).strip()
|
||||
config = load_config()
|
||||
@@ -684,6 +731,7 @@ async def get_remote_worker_log_endpoint(request):
|
||||
async with session.get(
|
||||
worker_url,
|
||||
params={"lines": str(lines_to_read)},
|
||||
headers=distributed_auth_headers(config),
|
||||
timeout=aiohttp.ClientTimeout(total=5),
|
||||
) as resp:
|
||||
if resp.status >= 400:
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Bootstrap entrypoints for ComfyUI-Distributed."""
|
||||
|
||||
from .entrypoint import build_node_mappings, initialize_runtime
|
||||
|
||||
__all__ = ["build_node_mappings", "initialize_runtime"]
|
||||
@@ -0,0 +1,66 @@
|
||||
"""Single startup owner for node surface and runtime initialization."""
|
||||
from __future__ import annotations
|
||||
|
||||
import atexit
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def build_node_mappings() -> tuple[dict[str, Any], dict[str, str]]:
|
||||
"""Build and cache the full exported node mapping surface."""
|
||||
from ..nodes import (
|
||||
NODE_CLASS_MAPPINGS as distributed_class_mappings,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as distributed_display_name_mappings,
|
||||
)
|
||||
from ..nodes.distributed_upscale import (
|
||||
NODE_CLASS_MAPPINGS as upscale_class_mappings,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as upscale_display_name_mappings,
|
||||
)
|
||||
|
||||
return (
|
||||
{**distributed_class_mappings, **upscale_class_mappings},
|
||||
{**distributed_display_name_mappings, **upscale_display_name_mappings},
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _runtime_initialized() -> dict[str, bool]:
|
||||
return {"done": False}
|
||||
|
||||
|
||||
def initialize_runtime(prompt_server: Any | None = None) -> None:
|
||||
"""Initialize distributed runtime side effects once."""
|
||||
state = _runtime_initialized()
|
||||
if state["done"]:
|
||||
return
|
||||
|
||||
import server
|
||||
|
||||
from ..api import bootstrap_routes
|
||||
from ..api.queue_orchestration import ensure_distributed_state
|
||||
from ..upscale.job_store import ensure_tile_jobs_initialized
|
||||
from ..utils.config import CONFIG_FILE, ensure_config_exists
|
||||
from ..utils.logging import debug_log
|
||||
from ..workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
|
||||
ensure_config_exists()
|
||||
bootstrap_routes()
|
||||
|
||||
server_instance = prompt_server if prompt_server is not None else server.PromptServer.instance
|
||||
ensure_distributed_state(server_instance)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
|
||||
if not os.environ.get("COMFYUI_IS_WORKER"):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
|
||||
state["done"] = True
|
||||
|
||||
node_mappings, _ = build_node_mappings()
|
||||
debug_log("Loaded Distributed nodes.")
|
||||
debug_log(f"Config file: {CONFIG_FILE}")
|
||||
debug_log(f"Available nodes: {list(node_mappings.keys())}")
|
||||
+1
-25
@@ -1,28 +1,4 @@
|
||||
# conftest.py — project-level pytest configuration.
|
||||
#
|
||||
# Problem: custom_nodes/ComfyUI-Distributed/__init__.py uses relative imports
|
||||
# (from .nodes.v3 import ...) that fail when pytest tries to import it as a
|
||||
# standalone module during Package.setup() for the root package node.
|
||||
#
|
||||
# Fix: patch Package.setup() to skip the root-package's __init__.py import.
|
||||
# All actual package context is provided by each test module via
|
||||
# importlib.util.spec_from_file_location with synthetic stub packages.
|
||||
|
||||
from _pytest.python import Package
|
||||
|
||||
_orig_pkg_setup = Package.setup
|
||||
|
||||
|
||||
def _patched_pkg_setup(self) -> None:
|
||||
# Skip the root package setup — its __init__.py uses relative imports
|
||||
# that require a parent package (ComfyUI's plugin loader) which is not
|
||||
# available in the test environment.
|
||||
if self.path == self.config.rootpath:
|
||||
return
|
||||
_orig_pkg_setup(self)
|
||||
|
||||
|
||||
Package.setup = _patched_pkg_setup
|
||||
"""Project-level pytest collection rules for plugin-style package layout."""
|
||||
|
||||
collect_ignore = [
|
||||
"__init__.py",
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Compatibility facade for legacy imports from `distributed`."""
|
||||
from __future__ import annotations
|
||||
|
||||
from .bootstrap.entrypoint import (
|
||||
build_node_mappings,
|
||||
initialize_runtime,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"initialize_runtime",
|
||||
]
|
||||
@@ -19,8 +19,8 @@ This document describes the **public HTTP API** added to ComfyUI-Distributed to
|
||||
- `POST /distributed/queue` — queues a workflow using the same distributed orchestration rules as the UI:
|
||||
- Detects distributed nodes in the prompt (`DistributedCollector`, `UltimateSDUpscaleDistributed`).
|
||||
- Resolves enabled/selected workers.
|
||||
- Probes and dispatches workers through `/distributed/worker_ws` by default.
|
||||
- If `settings.websocket_orchestration=false`, probes with `GET /prompt` and dispatches with `POST /prompt` instead.
|
||||
- Pings workers (`GET /prompt`) to include only reachable ones.
|
||||
- Dispatches the workflow to workers (`POST /prompt`).
|
||||
- Queues the master workflow in ComfyUI’s prompt queue.
|
||||
- If any `DistributedCollector` has `load_balance=true`, selects one least-busy participant for this run.
|
||||
|
||||
@@ -53,7 +53,6 @@ Queue a workflow for distributed execution.
|
||||
"delegate_master": false,
|
||||
"enabled_worker_ids": ["1", "2"],
|
||||
"workers": ["1", "2"],
|
||||
"auto_prepare": true,
|
||||
"trace_execution_id": "exec_1700000000_ab12cd"
|
||||
}
|
||||
```
|
||||
@@ -61,8 +60,7 @@ Queue a workflow for distributed execution.
|
||||
#### Fields
|
||||
|
||||
- `prompt` (required unless `workflow.prompt` is present, object)
|
||||
- A complete ComfyUI API-format prompt graph, using the same shape as `POST /prompt`.
|
||||
- This is not the normal visual workflow export from the ComfyUI editor.
|
||||
- The ComfyUI prompt/workflow graph, same shape as used by `POST /prompt`.
|
||||
- `workflow` (optional, object)
|
||||
- Workflow metadata that ComfyUI normally stores in `extra_pnginfo.workflow`.
|
||||
- If you don’t care about UI metadata, you can omit it.
|
||||
@@ -75,10 +73,6 @@ Queue a workflow for distributed execution.
|
||||
- The explicit worker IDs to consider for this run.
|
||||
- `workers` (optional, array of strings or objects with `id`)
|
||||
- Transitional alias for `enabled_worker_ids` used by older clients.
|
||||
- `auto_prepare` (optional, boolean)
|
||||
- Kept for wire compatibility.
|
||||
- Backend orchestration always runs with auto-prepare semantics.
|
||||
- If top-level `prompt` is omitted, backend will attempt `workflow.prompt`.
|
||||
- `trace_execution_id` (optional, string)
|
||||
- Passed through to orchestration logs.
|
||||
- Server log lines include the marker as `[exec:<trace_execution_id>]`.
|
||||
@@ -106,8 +100,7 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
|
||||
```json
|
||||
{
|
||||
"prompt_id": "<uuid>",
|
||||
"worker_count": 2,
|
||||
"auto_prepare_supported": true
|
||||
"worker_count": 2
|
||||
}
|
||||
```
|
||||
|
||||
@@ -124,14 +117,10 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
|
||||
|
||||
## Worker requirements (important)
|
||||
|
||||
For a worker to participate, it must be reachable from the master. By default:
|
||||
|
||||
- WebSocket probe and dispatch: `<worker-base>/distributed/worker_ws` must accept the connection.
|
||||
|
||||
If `settings.websocket_orchestration=false`:
|
||||
For a worker to participate, it must be reachable from the master:
|
||||
|
||||
- Health check: `GET <worker-base>/prompt` must return HTTP 200.
|
||||
- Dispatch: `POST <worker-base>/prompt` must accept the prompt.
|
||||
- Dispatch: `POST <worker-base>/prompt` must accept the workflow.
|
||||
|
||||
Also, for collector-based flows:
|
||||
|
||||
@@ -206,7 +195,7 @@ Proxy endpoint on master that fetches logs from a configured remote/cloud worker
|
||||
|
||||
## Examples
|
||||
|
||||
### 1) Minimal `curl` request envelope
|
||||
### 1) Minimal `curl`
|
||||
|
||||
```bash
|
||||
curl -X POST "http://127.0.0.1:8188/distributed/queue" \
|
||||
@@ -214,23 +203,18 @@ curl -X POST "http://127.0.0.1:8188/distributed/queue" \
|
||||
-d @payload.json
|
||||
```
|
||||
|
||||
`payload.json` must contain a complete ComfyUI API-format prompt. The abbreviated envelope below illustrates the request shape but is not directly executable:
|
||||
Where `payload.json` contains at least:
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": {
|
||||
"<node_id>": {
|
||||
"class_type": "<node class>",
|
||||
"inputs": {"<required_input>": "<value or connection>"}
|
||||
}
|
||||
"1": {"class_type": "KSampler", "inputs": {} }
|
||||
},
|
||||
"enabled_worker_ids": [],
|
||||
"client_id": "external-client"
|
||||
}
|
||||
```
|
||||
|
||||
Export or construct a valid API-format prompt with all required node inputs and at least one output node before submitting it.
|
||||
|
||||
### 2) Python (`requests`)
|
||||
|
||||
```python
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
# Native V3 node API
|
||||
|
||||
The package registers one `ComfyExtension` through `comfy_entrypoint()` and
|
||||
imports the versioned `comfy_api.v0_0_2` API. Use a ComfyUI version that provides
|
||||
that API; there is no V1 registration fallback.
|
||||
|
||||
## Compatibility
|
||||
|
||||
- All eight node IDs, display names, categories, visible input order, defaults
|
||||
and output order are retained. Existing execution algorithms remain in the
|
||||
private collector, utilities and upscale modules.
|
||||
- The image/audio dividers explicitly declare ten typed outputs. Their existing
|
||||
frontend extensions still show only the selected number of outputs. This
|
||||
replaces the V1 `ByPassTypeTuple` indexing workaround without changing saved
|
||||
IMAGE/AUDIO socket indices or the ten returned values.
|
||||
- Standard hidden context uses `cls.hidden`. Worker/orchestration metadata keeps
|
||||
its existing prompt-input names and Python defaults via `accept_all_inputs`;
|
||||
it does not become visible widgets.
|
||||
- The collector retains list-input handling. Upscale retains its always-changing
|
||||
fingerprint and creates a private runtime helper for each V3 execution, so
|
||||
mutable helper state is not attached to sanitized V3 class clones.
|
||||
- Routes, distributed state and the existing worker startup/shutdown hooks are
|
||||
initialized by `runtime/bootstrap.py` during `ComfyExtension.on_load()`.
|
||||
Worker mode still suppresses automatic worker launch. `distributed.py` is no
|
||||
longer an entrypoint.
|
||||
|
||||
## Verification
|
||||
|
||||
Run ordinary unit tests from this repository:
|
||||
|
||||
```bash
|
||||
python -m pytest tests -q -o addopts=
|
||||
```
|
||||
|
||||
Opt into real-framework acceptance using a ComfyUI checkout and its interpreter:
|
||||
|
||||
```bash
|
||||
COMFYUI_SOURCE_ROOT=/path/to/ComfyUI \
|
||||
/path/to/ComfyUI/.venv/bin/python -m pytest tests -q -o addopts=
|
||||
```
|
||||
|
||||
The acceptance subprocess uses CPU mode and the real ComfyUI loader, input
|
||||
parser, V3 class preparation, prompt validation and `PromptExecutor`. It compares
|
||||
all node schemas with the V1 fixture, checks all five bundled workflows' node
|
||||
IDs and socket/link contracts, exercises injected worker metadata, image/audio
|
||||
lists, collector aggregation and divider outputs, rejects invalid upscale enums,
|
||||
and decodes an actual preview PNG referenced by executor history.
|
||||
|
||||
The upscale GPU/model boundary is mocked to check argument forwarding and
|
||||
per-execution helper isolation. Full checkpoint inference, browser canvas
|
||||
acceptance and multi-host HTTP transport are not exercised. No HTTP listener,
|
||||
workers or model downloads are started; preview files use a temporary scratch
|
||||
directory. The test does not install this branch into a live custom-node folder.
|
||||
@@ -79,7 +79,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
|
||||
📺 [Watch Tutorial](https://www.youtube.com/watch?v=wxKKWMQhYTk)
|
||||
|
||||
**On Runpod:**
|
||||
> If using your own template, launch ComfyUI with `--listen --enable-cors-header` and clone `ComfyUI-Distributed` into `custom_nodes`. ⚠️ **Required!**
|
||||
> If using your own template, make sure you launch ComfyUI with the `--enable-cors-header` argument and you `git clone ComfyUI-Distributed` into custom_nodes. ⚠️ **Required!**
|
||||
|
||||
1. Register a [Runpod](https://get.runpod.io/0bw29uf3ug0p) account.
|
||||
2. On Runpod, go to Storage > New Network Volume and create a volume that will store the models you need. Start with 40 GB, you can always add more later. Learn more [about Network Volumes](https://docs.runpod.io/pods/storage/create-network-volumes).
|
||||
@@ -92,7 +92,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
|
||||
- SAGE_ATTENTION: optional optimisation (set to true/false)
|
||||
5. Deploy your pod.
|
||||
6. Connect to your pod using JupyterLabs. This gives us access to the pod's file system.
|
||||
7. Download models into `/workspace/ComfyUI/models/` (these will remain on your network drive even after you terminate the pod). Example commands below:
|
||||
7. Download models into /workspaces/ComfyUI/models/ (these will remain on your network drive even after you terminate the pod). Example commands below:
|
||||
```
|
||||
# Download from CivitAI
|
||||
comfy model download --url https://civitai.com/api/download/models/1759168 --relative-path /workspace/ComfyUI/models/checkpoints --set-civitai-api-token $CIVITAI_API_TOKEN
|
||||
|
||||
+37
-1
@@ -1 +1,37 @@
|
||||
"""Private execution helpers; public registration lives in nodes.v3."""
|
||||
from .utilities import (
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
ImageBatchDivider,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from .collector import DistributedCollectorNode
|
||||
from .branch import DistributedBranch
|
||||
from .branch_collector import DistributedBranchCollector
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DistributedCollector": DistributedCollectorNode,
|
||||
"DistributedBranch": DistributedBranch,
|
||||
"DistributedBranchCollector": DistributedBranchCollector,
|
||||
"DistributedSeed": DistributedSeed,
|
||||
"DistributedModelName": DistributedModelName,
|
||||
"DistributedValue": DistributedValue,
|
||||
"ImageBatchDivider": ImageBatchDivider,
|
||||
"AudioBatchDivider": AudioBatchDivider,
|
||||
"DistributedEmptyImage": DistributedEmptyImage,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DistributedCollector": "Distributed Collector",
|
||||
"DistributedBranch": "Distributed Branch",
|
||||
"DistributedBranchCollector": "Distributed Branch Collector",
|
||||
"DistributedSeed": "Distributed Seed",
|
||||
"DistributedModelName": "Distributed Model Name",
|
||||
"DistributedValue": "Distributed Value",
|
||||
"ImageBatchDivider": "Image Batch Divider",
|
||||
"AudioBatchDivider": "Audio Batch Divider",
|
||||
"DistributedEmptyImage": "Distributed Empty Image",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
from typing import Any
|
||||
|
||||
from .utilities import any_type
|
||||
|
||||
|
||||
class DistributedBranch:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedBranch"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"input": (any_type,),
|
||||
"num_branches": ("INT", {"default": 2, "min": 2, "max": 10, "step": 1}),
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
"assigned_branch": ("INT", {"default": -1, "min": -1, "max": 9}),
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = tuple([any_type] * 10)
|
||||
RETURN_NAMES = tuple([f"branch_{idx + 1}" for idx in range(10)])
|
||||
FUNCTION = "branch"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def branch(
|
||||
self,
|
||||
input: Any,
|
||||
num_branches: int = 2,
|
||||
**_kwargs: Any,
|
||||
) -> tuple[Any, ...]:
|
||||
return tuple([input] * 10)
|
||||
@@ -0,0 +1,298 @@
|
||||
import asyncio
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import aiohttp
|
||||
import torch
|
||||
|
||||
from ..utils.auth import distributed_auth_headers
|
||||
from ..utils.async_helpers import run_async_in_server_loop
|
||||
from ..utils.config import is_master_delegate_only, load_config
|
||||
from ..utils.image import encode_tensor_png_data_url, ensure_contiguous
|
||||
from ..utils.logging import log
|
||||
from ..utils.network import get_client_session
|
||||
from ..utils.worker_ids import coerce_enabled_worker_ids
|
||||
from .context_kwargs import parse_distributed_hidden_context
|
||||
from .hidden_inputs import build_distributed_hidden_inputs
|
||||
from .queue_wait import collect_worker_queue_results
|
||||
from .runtime_helpers import (
|
||||
get_prompt_server_instance as _get_prompt_server_instance,
|
||||
throw_if_processing_interrupted as _throw_if_processing_interrupted,
|
||||
)
|
||||
from .utilities import any_type
|
||||
|
||||
|
||||
MAX_BRANCH_OUTPUTS = 10
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BranchRunContext:
|
||||
multi_job_id: str = ""
|
||||
is_worker: bool = False
|
||||
master_url: str = ""
|
||||
enabled_worker_ids: list[str] = field(default_factory=list)
|
||||
worker_id: str = ""
|
||||
assigned_branch: int = -1
|
||||
delegate_only: bool = False
|
||||
|
||||
class DistributedBranchCollector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedBranchCollector"]) -> dict[str, Any]:
|
||||
optional_inputs = {
|
||||
f"branch_{idx + 1}": (any_type,)
|
||||
for idx in range(MAX_BRANCH_OUTPUTS)
|
||||
}
|
||||
return {
|
||||
"required": {
|
||||
"num_branches": ("INT", {"default": 2, "min": 2, "max": MAX_BRANCH_OUTPUTS, "step": 1}),
|
||||
},
|
||||
"optional": optional_inputs,
|
||||
"hidden": build_distributed_hidden_inputs(
|
||||
include_assigned_branch=True,
|
||||
assigned_branch_max=MAX_BRANCH_OUTPUTS - 1,
|
||||
),
|
||||
}
|
||||
|
||||
RETURN_TYPES = tuple([any_type] * MAX_BRANCH_OUTPUTS)
|
||||
RETURN_NAMES = tuple([f"branch_{idx + 1}" for idx in range(MAX_BRANCH_OUTPUTS)])
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def _branch_values(self, **kwargs):
|
||||
values = []
|
||||
for idx in range(MAX_BRANCH_OUTPUTS):
|
||||
values.append(kwargs.get(f"branch_{idx + 1}", None))
|
||||
return values
|
||||
|
||||
def _build_run_context(self, **kwargs: Any) -> BranchRunContext:
|
||||
try:
|
||||
assigned_branch = int(kwargs.get("assigned_branch", -1))
|
||||
except (TypeError, ValueError):
|
||||
assigned_branch = -1
|
||||
|
||||
common_context = parse_distributed_hidden_context(kwargs)
|
||||
return BranchRunContext(
|
||||
assigned_branch=assigned_branch,
|
||||
**common_context,
|
||||
)
|
||||
|
||||
def _resolve_local_branch_value(self, branch_values, assigned_branch):
|
||||
try:
|
||||
assigned_idx = int(assigned_branch)
|
||||
except (TypeError, ValueError):
|
||||
assigned_idx = -1
|
||||
|
||||
if 0 <= assigned_idx < MAX_BRANCH_OUTPUTS:
|
||||
assigned_value = branch_values[assigned_idx]
|
||||
if assigned_value is not None:
|
||||
return assigned_value, assigned_idx
|
||||
|
||||
# In participant-pruned prompts, each participant should typically have
|
||||
# exactly one local branch input. If socket names drifted (e.g. branch_2
|
||||
# connected for a participant assigned to slot 0), remap that sole value
|
||||
# to the assigned slot so the participant still contributes correctly.
|
||||
non_none_indices = [idx for idx, value in enumerate(branch_values) if value is not None]
|
||||
if len(non_none_indices) == 1:
|
||||
return branch_values[non_none_indices[0]], assigned_idx
|
||||
|
||||
for idx, value in enumerate(branch_values):
|
||||
if value is not None:
|
||||
return value, idx
|
||||
return None, assigned_idx
|
||||
|
||||
def _build_outputs(self, branch_values, num_branches):
|
||||
outputs = [None] * MAX_BRANCH_OUTPUTS
|
||||
try:
|
||||
branch_count = int(num_branches)
|
||||
except (TypeError, ValueError):
|
||||
branch_count = 2
|
||||
branch_count = max(2, min(branch_count, MAX_BRANCH_OUTPUTS))
|
||||
|
||||
for idx in range(branch_count):
|
||||
value = branch_values[idx]
|
||||
if value is not None:
|
||||
outputs[idx] = value
|
||||
return tuple(outputs)
|
||||
|
||||
def run(
|
||||
self,
|
||||
num_branches: int = 2,
|
||||
**kwargs: Any,
|
||||
) -> tuple[Any, ...]:
|
||||
branch_values = self._branch_values(**kwargs)
|
||||
context = self._build_run_context(**kwargs)
|
||||
|
||||
if not context.multi_job_id:
|
||||
return self._build_outputs(branch_values, num_branches)
|
||||
|
||||
return run_async_in_server_loop(
|
||||
self.execute(
|
||||
branch_values,
|
||||
num_branches=num_branches,
|
||||
context=context,
|
||||
)
|
||||
)
|
||||
|
||||
def _participant_branch_slots(self, enabled_workers, delegate_mode, num_branches):
|
||||
participants = list(enabled_workers) if delegate_mode else ["master"] + list(enabled_workers)
|
||||
participant_count = len(participants)
|
||||
if participant_count <= 0:
|
||||
return {}
|
||||
|
||||
slots_by_participant = {}
|
||||
for branch_slot in range(int(num_branches)):
|
||||
participant_id = str(participants[branch_slot % participant_count])
|
||||
slots_by_participant.setdefault(participant_id, []).append(branch_slot)
|
||||
return slots_by_participant
|
||||
|
||||
def _expected_worker_ids_for_branches(self, enabled_workers, delegate_mode, num_branches):
|
||||
slots_by_participant = self._participant_branch_slots(enabled_workers, delegate_mode, num_branches)
|
||||
expected = []
|
||||
for participant_id, slots in slots_by_participant.items():
|
||||
if participant_id == "master" or not slots:
|
||||
continue
|
||||
expected.append(str(participant_id))
|
||||
return expected
|
||||
|
||||
def _fallback_for_missing_branch(self, source_value):
|
||||
if isinstance(source_value, torch.Tensor):
|
||||
tensor = source_value
|
||||
if tensor.ndim == 3:
|
||||
tensor = tensor.unsqueeze(0)
|
||||
if tensor.ndim == 4:
|
||||
fallback = torch.zeros_like(tensor)
|
||||
if fallback.is_cuda:
|
||||
fallback = fallback.cpu()
|
||||
return ensure_contiguous(fallback)
|
||||
return torch.zeros((1, 64, 64, 3), dtype=torch.float32)
|
||||
|
||||
async def _send_branch_result_to_master(self, value, branch_idx, multi_job_id, master_url, worker_id):
|
||||
if not isinstance(value, torch.Tensor):
|
||||
raise ValueError("DistributedBranchCollector currently supports torch.Tensor results for worker transfer.")
|
||||
|
||||
image_batch = value
|
||||
if image_batch.ndim == 3:
|
||||
image_batch = image_batch.unsqueeze(0)
|
||||
if image_batch.ndim != 4 or image_batch.shape[0] <= 0:
|
||||
raise ValueError(
|
||||
"DistributedBranchCollector tensor result must be IMAGE-like with shape [B,H,W,C] or [H,W,C]."
|
||||
)
|
||||
|
||||
payload = {
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": int(branch_idx),
|
||||
"image": encode_tensor_png_data_url(image_batch, 0),
|
||||
"is_last": True,
|
||||
}
|
||||
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/job_complete"
|
||||
async with session.post(
|
||||
url,
|
||||
json=payload,
|
||||
headers=distributed_auth_headers(load_config()),
|
||||
timeout=aiohttp.ClientTimeout(total=60),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
|
||||
async def execute(self, *args: Any, **kwargs: Any) -> tuple[Any, ...]:
|
||||
"""Compatibility wrapper with a uniform execute() signature across collectors."""
|
||||
return await self._execute_branch(*args, **kwargs)
|
||||
|
||||
async def _execute_branch(
|
||||
self,
|
||||
branch_values: list[Any],
|
||||
num_branches: int = 2,
|
||||
context: BranchRunContext | None = None,
|
||||
) -> tuple[Any, ...]:
|
||||
run_context = context or BranchRunContext()
|
||||
try:
|
||||
branch_count = int(num_branches)
|
||||
except (TypeError, ValueError):
|
||||
branch_count = 2
|
||||
branch_count = max(2, min(branch_count, MAX_BRANCH_OUTPUTS))
|
||||
|
||||
outputs = [None] * MAX_BRANCH_OUTPUTS
|
||||
for idx in range(branch_count):
|
||||
if branch_values[idx] is not None:
|
||||
outputs[idx] = branch_values[idx]
|
||||
|
||||
local_value, local_branch_idx = self._resolve_local_branch_value(branch_values, run_context.assigned_branch)
|
||||
if 0 <= local_branch_idx < MAX_BRANCH_OUTPUTS and local_value is not None:
|
||||
outputs[local_branch_idx] = local_value
|
||||
|
||||
if run_context.is_worker:
|
||||
if 0 <= local_branch_idx < MAX_BRANCH_OUTPUTS and local_value is not None:
|
||||
try:
|
||||
await self._send_branch_result_to_master(
|
||||
local_value,
|
||||
local_branch_idx,
|
||||
run_context.multi_job_id,
|
||||
run_context.master_url,
|
||||
run_context.worker_id,
|
||||
)
|
||||
outputs[local_branch_idx] = local_value
|
||||
except Exception as exc:
|
||||
log(f"Worker - DistributedBranchCollector failed to send branch result: {exc}")
|
||||
return tuple(outputs)
|
||||
|
||||
delegate_mode = bool(run_context.delegate_only or is_master_delegate_only())
|
||||
enabled_workers = coerce_enabled_worker_ids(run_context.enabled_worker_ids)
|
||||
slots_by_participant = self._participant_branch_slots(enabled_workers, delegate_mode, branch_count)
|
||||
expected_worker_ids = self._expected_worker_ids_for_branches(enabled_workers, delegate_mode, branch_count)
|
||||
expected_workers = set(expected_worker_ids)
|
||||
|
||||
if not expected_workers:
|
||||
return tuple(outputs)
|
||||
|
||||
prompt_server = _get_prompt_server_instance()
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if run_context.multi_job_id not in prompt_server.distributed_pending_jobs:
|
||||
prompt_server.distributed_pending_jobs[run_context.multi_job_id] = asyncio.Queue()
|
||||
|
||||
workers_done: set[str] = set()
|
||||
|
||||
def _handle_queue_result(result: dict[str, Any]) -> None:
|
||||
image_index = result.get("image_index")
|
||||
tensor = result.get("tensor")
|
||||
try:
|
||||
slot_idx = int(image_index)
|
||||
except (TypeError, ValueError):
|
||||
slot_idx = -1
|
||||
|
||||
if 0 <= slot_idx < MAX_BRANCH_OUTPUTS and tensor is not None:
|
||||
if isinstance(tensor, torch.Tensor):
|
||||
if tensor.is_cuda:
|
||||
tensor = tensor.cpu()
|
||||
tensor = ensure_contiguous(tensor)
|
||||
outputs[slot_idx] = tensor
|
||||
|
||||
try:
|
||||
workers_done = await collect_worker_queue_results(
|
||||
prompt_server=prompt_server,
|
||||
multi_job_id=run_context.multi_job_id,
|
||||
expected_workers=expected_workers,
|
||||
on_result=_handle_queue_result,
|
||||
timeout_log_prefix=(
|
||||
"Master - DistributedBranchCollector heartbeat timeout. "
|
||||
"Still waiting for workers: "
|
||||
),
|
||||
throw_if_interrupted=_throw_if_processing_interrupted,
|
||||
)
|
||||
|
||||
finally:
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
prompt_server.distributed_pending_jobs.pop(run_context.multi_job_id, None)
|
||||
if hasattr(prompt_server, "distributed_job_allowed_workers"):
|
||||
prompt_server.distributed_job_allowed_workers.pop(run_context.multi_job_id, None)
|
||||
|
||||
missing_workers = sorted(expected_workers - workers_done)
|
||||
if missing_workers:
|
||||
fallback = self._fallback_for_missing_branch(local_value)
|
||||
for missing_worker_id in missing_workers:
|
||||
for slot_idx in slots_by_participant.get(str(missing_worker_id), []):
|
||||
if 0 <= slot_idx < MAX_BRANCH_OUTPUTS and outputs[slot_idx] is None:
|
||||
outputs[slot_idx] = fallback
|
||||
|
||||
return tuple(outputs)
|
||||
+338
-361
@@ -1,34 +1,52 @@
|
||||
import torch
|
||||
import io
|
||||
import json
|
||||
import asyncio
|
||||
import time
|
||||
import base64
|
||||
import io
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import aiohttp
|
||||
import server as _server
|
||||
import comfy.model_management
|
||||
import torch
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
from ..utils.auth import distributed_auth_headers
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.config import get_worker_timeout_seconds, load_config, is_master_delegate_only
|
||||
from ..utils.constants import HEARTBEAT_INTERVAL
|
||||
from ..utils.image import tensor_to_pil, pil_to_tensor, ensure_contiguous
|
||||
from ..utils.network import build_worker_url, get_client_session, probe_worker
|
||||
from ..utils.worker_ids import coerce_enabled_worker_ids
|
||||
from ..utils.audio_payload import encode_audio_payload
|
||||
from ..utils.async_helpers import run_async_in_server_loop
|
||||
from .context_kwargs import parse_distributed_hidden_context
|
||||
from .hidden_inputs import build_distributed_hidden_inputs
|
||||
|
||||
prompt_server = _server.PromptServer.instance
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CollectorRunContext:
|
||||
multi_job_id: str = ""
|
||||
is_worker: bool = False
|
||||
master_url: str = ""
|
||||
enabled_worker_ids: list[str] = field(default_factory=list)
|
||||
worker_batch_size: int = 1
|
||||
worker_id: str = ""
|
||||
pass_through: bool = False
|
||||
delegate_only: bool = False
|
||||
|
||||
|
||||
class DistributedCollectorNode:
|
||||
INPUT_IS_LIST = True
|
||||
EMPTY_AUDIO = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls: type["DistributedCollectorNode"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"load_balance": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
@@ -37,171 +55,102 @@ class DistributedCollectorNode:
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"images": ("IMAGE",),
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"master_url": ("STRING", {"default": ""}),
|
||||
"enabled_worker_ids": ("STRING", {"default": "[]"}),
|
||||
"worker_batch_size": ("INT", {"default": 1, "min": 1, "max": 1024}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
"pass_through": ("BOOLEAN", {"default": False}),
|
||||
"delegate_only": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": { "audio": ("AUDIO",) },
|
||||
"hidden": build_distributed_hidden_inputs(
|
||||
include_worker_batch_size=True,
|
||||
include_pass_through=True,
|
||||
),
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "AUDIO")
|
||||
RETURN_NAMES = ("images", "audio")
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "image"
|
||||
|
||||
@staticmethod
|
||||
def _unwrap_list_input(value):
|
||||
"""Unwrap scalar inputs when ComfyUI passes them via INPUT_IS_LIST."""
|
||||
if isinstance(value, (list, tuple)) and len(value) == 1:
|
||||
return value[0]
|
||||
return value
|
||||
|
||||
def _normalize_images_input(self, images):
|
||||
"""Collapse ComfyUI list IMAGE inputs into a normal batched IMAGE tensor."""
|
||||
if isinstance(images, (list, tuple)):
|
||||
if not images:
|
||||
raise ValueError("Collector received an empty image list")
|
||||
if not all(isinstance(image, torch.Tensor) for image in images):
|
||||
raise TypeError("Collector expected IMAGE list items to be torch.Tensor instances")
|
||||
if len(images) == 1:
|
||||
return ensure_contiguous(images[0])
|
||||
return ensure_contiguous(torch.cat([ensure_contiguous(image) for image in images], dim=0))
|
||||
return ensure_contiguous(images)
|
||||
def _build_run_context(self, **kwargs: Any) -> CollectorRunContext:
|
||||
try:
|
||||
worker_batch_size = int(kwargs.get("worker_batch_size", 1))
|
||||
except (TypeError, ValueError):
|
||||
worker_batch_size = 1
|
||||
|
||||
def _normalize_audio_input(self, audio):
|
||||
"""Collapse ComfyUI list AUDIO inputs into a single AUDIO payload when present."""
|
||||
if not isinstance(audio, (list, tuple)):
|
||||
return audio
|
||||
|
||||
audio_items = [item for item in audio if item is not None]
|
||||
if not audio_items:
|
||||
return None
|
||||
if len(audio_items) == 1:
|
||||
return audio_items[0]
|
||||
|
||||
waveforms = []
|
||||
sample_rate = 44100
|
||||
for item in audio_items:
|
||||
if not isinstance(item, dict):
|
||||
raise TypeError("Collector expected AUDIO list items to be dictionaries")
|
||||
waveform = item.get("waveform")
|
||||
if waveform is None or waveform.numel() == 0:
|
||||
continue
|
||||
waveforms.append(waveform)
|
||||
if sample_rate == 44100:
|
||||
sample_rate = item.get("sample_rate", 44100)
|
||||
|
||||
if not waveforms:
|
||||
return None
|
||||
return {"waveform": torch.cat(waveforms, dim=-1), "sample_rate": sample_rate}
|
||||
|
||||
def run(self, images=None, load_balance=False, audio=None, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False, delegate_only=False):
|
||||
if images is not None:
|
||||
images = self._normalize_images_input(images)
|
||||
audio = self._normalize_audio_input(audio)
|
||||
load_balance = self._unwrap_list_input(load_balance)
|
||||
multi_job_id = self._unwrap_list_input(multi_job_id)
|
||||
is_worker = self._unwrap_list_input(is_worker)
|
||||
master_url = self._unwrap_list_input(master_url)
|
||||
enabled_worker_ids = self._unwrap_list_input(enabled_worker_ids)
|
||||
worker_batch_size = self._unwrap_list_input(worker_batch_size)
|
||||
worker_id = self._unwrap_list_input(worker_id)
|
||||
pass_through = self._unwrap_list_input(pass_through)
|
||||
delegate_only = self._unwrap_list_input(delegate_only)
|
||||
|
||||
remote_only_master = (
|
||||
bool(multi_job_id)
|
||||
and not is_worker
|
||||
and (delegate_only or is_master_delegate_only())
|
||||
common_context = parse_distributed_hidden_context(kwargs)
|
||||
return CollectorRunContext(
|
||||
worker_batch_size=max(worker_batch_size, 1),
|
||||
pass_through=bool(kwargs.get("pass_through", False)),
|
||||
**common_context,
|
||||
)
|
||||
if images is None and audio is None and not remote_only_master:
|
||||
raise ValueError("DistributedCollector requires at least one image or audio input")
|
||||
|
||||
def run(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
load_balance: bool = False,
|
||||
audio: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
context = self._build_run_context(**kwargs)
|
||||
# Create empty audio if not provided
|
||||
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
|
||||
|
||||
if not multi_job_id or pass_through:
|
||||
if pass_through:
|
||||
if not context.multi_job_id or context.pass_through:
|
||||
if context.pass_through:
|
||||
debug_log("Collector: pass-through mode enabled, returning images unchanged")
|
||||
return (images, audio if audio is not None else empty_audio)
|
||||
|
||||
# Use async helper to run in server loop
|
||||
result = run_async_in_server_loop(
|
||||
self.execute(
|
||||
images,
|
||||
audio,
|
||||
load_balance,
|
||||
multi_job_id,
|
||||
is_worker,
|
||||
master_url,
|
||||
enabled_worker_ids,
|
||||
worker_batch_size,
|
||||
worker_id,
|
||||
delegate_only,
|
||||
images=images,
|
||||
audio=audio,
|
||||
load_balance=load_balance,
|
||||
context=context,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
async def send_batch_to_master(self, image_batch, audio, multi_job_id, master_url, worker_id):
|
||||
"""Send an image batch, optionally with audio, or an audio-only completion."""
|
||||
async def send_batch_to_master(
|
||||
self,
|
||||
image_batch: torch.Tensor,
|
||||
audio: dict[str, Any] | None,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str,
|
||||
) -> None:
|
||||
"""Send image batch to master via canonical JSON envelopes."""
|
||||
batch_size = image_batch.shape[0]
|
||||
if batch_size == 0:
|
||||
return
|
||||
|
||||
encoded_audio = encode_audio_payload(audio)
|
||||
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/job_complete"
|
||||
for batch_idx in range(batch_size):
|
||||
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
|
||||
byte_io = io.BytesIO()
|
||||
img.save(byte_io, format='PNG', compress_level=0)
|
||||
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
|
||||
payload = {
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": int(batch_idx),
|
||||
"image": f"data:image/png;base64,{encoded_image}",
|
||||
"is_last": bool(batch_idx == batch_size - 1),
|
||||
}
|
||||
if payload["is_last"] and encoded_audio is not None:
|
||||
payload["audio"] = encoded_audio
|
||||
|
||||
payloads = []
|
||||
batch_size = 0 if image_batch is None else image_batch.shape[0]
|
||||
if batch_size == 0:
|
||||
if encoded_audio is None:
|
||||
raise ValueError("Worker completion requires image or audio data")
|
||||
payloads.append(
|
||||
{
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": 0,
|
||||
"audio": encoded_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
else:
|
||||
for batch_idx in range(batch_size):
|
||||
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
|
||||
byte_io = io.BytesIO()
|
||||
img.save(byte_io, format='PNG', compress_level=0)
|
||||
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
|
||||
payload = {
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": int(batch_idx),
|
||||
"image": f"data:image/png;base64,{encoded_image}",
|
||||
"is_last": bool(batch_idx == batch_size - 1),
|
||||
}
|
||||
if payload["is_last"] and encoded_audio is not None:
|
||||
payload["audio"] = encoded_audio
|
||||
payloads.append(payload)
|
||||
|
||||
for payload in payloads:
|
||||
timeout_seconds = 60 if "image" in payload else 600
|
||||
try:
|
||||
async with session.post(
|
||||
url,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=timeout_seconds),
|
||||
headers=distributed_auth_headers(load_config()),
|
||||
timeout=aiohttp.ClientTimeout(total=60),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
media_type = "image/audio" if "image" in payload else "audio-only"
|
||||
log(f"Worker - Failed to send canonical {media_type} envelope to master: {e}")
|
||||
log(f"Worker - Failed to send canonical image envelope to master: {e}")
|
||||
debug_log(f"Worker - Full error details: URL={url}")
|
||||
raise
|
||||
raise # Re-raise to handle at caller level
|
||||
|
||||
def _combine_audio(self, master_audio, worker_audio, empty_audio, worker_order=None):
|
||||
"""Combine audio from master and workers into a single audio output.
|
||||
@@ -283,8 +232,8 @@ class DistributedCollectorNode:
|
||||
images_on_cpu,
|
||||
delegate_mode: bool,
|
||||
fallback_images,
|
||||
):
|
||||
"""Assemble final tensor, or return None when the job contains only audio."""
|
||||
) -> torch.Tensor:
|
||||
"""Assemble final tensor: master first, then workers in enabled order."""
|
||||
ordered_tensors = []
|
||||
if not delegate_mode and images_on_cpu is not None:
|
||||
for i in range(master_batch_size):
|
||||
@@ -315,242 +264,270 @@ class DistributedCollectorNode:
|
||||
|
||||
if cpu_tensors:
|
||||
return ensure_contiguous(torch.cat(cpu_tensors, dim=0))
|
||||
if fallback_images is not None:
|
||||
elif fallback_images is not None:
|
||||
return ensure_contiguous(fallback_images)
|
||||
return None
|
||||
|
||||
async def execute(self, images, audio, load_balance=False, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", delegate_only=False):
|
||||
if is_worker:
|
||||
# Worker mode: send images and audio to master in a single batch
|
||||
image_count = 0 if images is None else images.shape[0]
|
||||
debug_log(f"Worker - Job {multi_job_id} complete. Sending {image_count} image(s) to master")
|
||||
await self.send_batch_to_master(images, audio, multi_job_id, master_url, worker_id)
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
else:
|
||||
delegate_mode = delegate_only or is_master_delegate_only()
|
||||
# Master mode: collect images and audio from workers
|
||||
enabled_workers_raw = json.loads(enabled_worker_ids)
|
||||
enabled_workers = []
|
||||
seen_enabled = set()
|
||||
for worker_id in enabled_workers_raw:
|
||||
worker_id_str = str(worker_id)
|
||||
if worker_id_str in seen_enabled:
|
||||
continue
|
||||
seen_enabled.add(worker_id_str)
|
||||
enabled_workers.append(worker_id_str)
|
||||
expected_workers = set(enabled_workers)
|
||||
num_workers = len(expected_workers)
|
||||
if num_workers == 0:
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
raise ValueError("No image data collected from master or workers")
|
||||
|
||||
# Create the queue before any expensive local work to avoid job_complete race.
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_jobs:
|
||||
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
|
||||
debug_log(f"Master - Initialized queue early for job {multi_job_id}")
|
||||
else:
|
||||
existing_size = prompt_server.distributed_pending_jobs[multi_job_id].qsize()
|
||||
debug_log(f"Master - Using existing queue for job {multi_job_id} (current size: {existing_size})")
|
||||
|
||||
if delegate_mode:
|
||||
master_batch_size = 0
|
||||
images_on_cpu = None
|
||||
master_audio = None
|
||||
debug_log(f"Master - Job {multi_job_id}: Delegate-only mode enabled, collecting exclusively from {num_workers} workers")
|
||||
async def _ensure_pending_queue(self, multi_job_id):
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_jobs:
|
||||
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
|
||||
debug_log(f"Master - Initialized queue early for job {multi_job_id}")
|
||||
else:
|
||||
if images is None:
|
||||
images_on_cpu = None
|
||||
master_batch_size = 0
|
||||
else:
|
||||
images_on_cpu = ensure_contiguous(images.cpu())
|
||||
master_batch_size = images.shape[0]
|
||||
master_audio = audio # Keep master's audio for later
|
||||
debug_log(f"Master - Job {multi_job_id}: Master has {master_batch_size} images, collecting from {num_workers} workers...")
|
||||
existing_size = prompt_server.distributed_pending_jobs[multi_job_id].qsize()
|
||||
debug_log(f"Master - Using existing queue for job {multi_job_id} (current size: {existing_size})")
|
||||
|
||||
async def _cleanup_pending_queue(self, multi_job_id):
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_jobs:
|
||||
del prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
if hasattr(prompt_server, "distributed_job_allowed_workers"):
|
||||
prompt_server.distributed_job_allowed_workers.pop(multi_job_id, None)
|
||||
|
||||
# Initialize storage for collected images and audio
|
||||
worker_images = {} # Dict to store images by worker_id and index
|
||||
worker_audio = {} # Dict to store audio by worker_id
|
||||
|
||||
# Collect images until all workers report they're done
|
||||
collected_count = 0
|
||||
workers_done = set()
|
||||
|
||||
# Use unified worker timeout from config/UI with simple sliced waits
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
|
||||
last_activity = time.time()
|
||||
|
||||
|
||||
# Get queue size before starting
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
q = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
initial_size = q.qsize()
|
||||
def _build_master_inputs(self, images, audio, delegate_mode, multi_job_id, num_workers):
|
||||
if delegate_mode:
|
||||
debug_log(
|
||||
f"Master - Job {multi_job_id}: Delegate-only mode enabled, collecting exclusively from {num_workers} workers"
|
||||
)
|
||||
return 0, None, None
|
||||
|
||||
# NEW: Initialize progress bar for workers (total = num_workers)
|
||||
p = ProgressBar(num_workers)
|
||||
images_on_cpu = ensure_contiguous(images.cpu())
|
||||
master_batch_size = images.shape[0]
|
||||
debug_log(
|
||||
f"Master - Job {multi_job_id}: Master has {master_batch_size} images, collecting from {num_workers} workers..."
|
||||
)
|
||||
return master_batch_size, images_on_cpu, audio
|
||||
|
||||
def mark_worker_done(done_worker_id):
|
||||
done_worker_id = str(done_worker_id)
|
||||
if done_worker_id not in expected_workers:
|
||||
async def _probe_missing_workers_busy(self, missing_workers):
|
||||
any_busy = False
|
||||
try:
|
||||
cfg = load_config()
|
||||
cfg_workers = cfg.get('workers', [])
|
||||
for wid in list(missing_workers):
|
||||
wrec = next((w for w in cfg_workers if str(w.get('id')) == str(wid)), None)
|
||||
if not wrec:
|
||||
debug_log(f"Collector probe: worker {wid} not found in config")
|
||||
continue
|
||||
worker_url = build_worker_url(wrec)
|
||||
try:
|
||||
payload = await probe_worker(worker_url, timeout=2.0)
|
||||
queue_remaining = None
|
||||
if payload is not None:
|
||||
queue_remaining = int(payload.get('exec_info', {}).get('queue_remaining', 0))
|
||||
debug_log(
|
||||
f"Master - Ignoring completion from unexpected worker {done_worker_id} for job {multi_job_id}"
|
||||
"Collector probe: worker "
|
||||
f"{wid} online={payload is not None} queue_remaining={queue_remaining}"
|
||||
)
|
||||
return
|
||||
if done_worker_id in workers_done:
|
||||
debug_log(
|
||||
f"Master - Ignoring duplicate completion from worker {done_worker_id} for job {multi_job_id}"
|
||||
)
|
||||
return
|
||||
workers_done.add(done_worker_id)
|
||||
p.update(1) # +1 per completed expected worker
|
||||
|
||||
try:
|
||||
while len(workers_done) < num_workers:
|
||||
# Check for user interruption to abort collection promptly
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
try:
|
||||
# Get the queue again each time to ensure we have the right reference
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
q = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
current_size = q.qsize()
|
||||
|
||||
result = await asyncio.wait_for(q.get(), timeout=slice_timeout)
|
||||
worker_id = result['worker_id']
|
||||
is_last = result.get('is_last', False)
|
||||
count = self._store_worker_result(worker_images, result)
|
||||
collected_count += count
|
||||
debug_log(
|
||||
f"Master - Got canonical result from worker {worker_id}, "
|
||||
f"image {result.get('image_index', 0)}, is_last={is_last}"
|
||||
)
|
||||
|
||||
# Collect audio data if present
|
||||
result_audio = result.get('audio')
|
||||
if result_audio is not None:
|
||||
worker_audio[worker_id] = result_audio
|
||||
debug_log(f"Master - Got audio from worker {worker_id}")
|
||||
|
||||
# Record activity and refresh timeout baseline
|
||||
last_activity = time.time()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
|
||||
if is_last:
|
||||
mark_worker_done(worker_id)
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
# If we still have time, continue polling; otherwise handle timeout
|
||||
if (time.time() - last_activity) < base_timeout:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
continue
|
||||
# Re-check for user interruption after timeout expiry
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
missing_workers = set(str(w) for w in enabled_workers) - workers_done
|
||||
elapsed = time.time() - last_activity
|
||||
for missing_worker_id in sorted(missing_workers):
|
||||
log(
|
||||
"Master - Heartbeat timeout: "
|
||||
f"worker={missing_worker_id}, elapsed={elapsed:.1f}s"
|
||||
)
|
||||
if payload is not None and queue_remaining and queue_remaining > 0:
|
||||
any_busy = True
|
||||
log(
|
||||
f"Master - Heartbeat timeout. Still waiting for workers: {list(missing_workers)} "
|
||||
f"(elapsed={elapsed:.1f}s)"
|
||||
f"Master - Probe grace: worker {wid} appears busy "
|
||||
f"(queue_remaining={queue_remaining}). Continuing to wait."
|
||||
)
|
||||
|
||||
# Probe missing workers' /prompt endpoints to check if they are actively processing
|
||||
any_busy = False
|
||||
try:
|
||||
cfg = load_config()
|
||||
cfg_workers = cfg.get('workers', [])
|
||||
for wid in list(missing_workers):
|
||||
wrec = next((w for w in cfg_workers if str(w.get('id')) == str(wid)), None)
|
||||
if not wrec:
|
||||
debug_log(f"Collector probe: worker {wid} not found in config")
|
||||
continue
|
||||
worker_url = build_worker_url(wrec)
|
||||
try:
|
||||
payload = await probe_worker(worker_url, timeout=2.0)
|
||||
queue_remaining = None
|
||||
if payload is not None:
|
||||
queue_remaining = int(payload.get('exec_info', {}).get('queue_remaining', 0))
|
||||
debug_log(
|
||||
"Collector probe: worker "
|
||||
f"{wid} online={payload is not None} queue_remaining={queue_remaining}"
|
||||
)
|
||||
if payload is not None and queue_remaining and queue_remaining > 0:
|
||||
any_busy = True
|
||||
log(
|
||||
f"Master - Probe grace: worker {wid} appears busy "
|
||||
f"(queue_remaining={queue_remaining}). Continuing to wait."
|
||||
)
|
||||
break
|
||||
except Exception as e:
|
||||
debug_log(f"Collector probe failed for worker {wid}: {e}")
|
||||
except Exception as e:
|
||||
debug_log(f"Collector probe setup error: {e}")
|
||||
|
||||
if any_busy:
|
||||
# Refresh last_activity and continue waiting
|
||||
last_activity = time.time()
|
||||
# Refresh base timeout in case the user changed it in UI
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
continue
|
||||
|
||||
# Check queue size again with lock
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_jobs:
|
||||
final_q = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
final_size = final_q.qsize()
|
||||
|
||||
# Try to drain any remaining items
|
||||
remaining_items = []
|
||||
while not final_q.empty():
|
||||
try:
|
||||
item = final_q.get_nowait()
|
||||
remaining_items.append(item)
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
if remaining_items:
|
||||
# Process them
|
||||
for item in remaining_items:
|
||||
worker_id = item['worker_id']
|
||||
is_last = item.get('is_last', False)
|
||||
|
||||
collected_count += self._store_worker_result(worker_images, item)
|
||||
|
||||
if is_last:
|
||||
mark_worker_done(worker_id)
|
||||
else:
|
||||
log(f"Master - Queue {multi_job_id} no longer exists!")
|
||||
break
|
||||
except comfy.model_management.InterruptProcessingException:
|
||||
# Cleanup queue on interruption and re-raise to abort prompt cleanly
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_jobs:
|
||||
del prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
raise
|
||||
|
||||
total_collected = sum(len(imgs) for imgs in worker_images.values())
|
||||
|
||||
# Clean up job queue
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_jobs:
|
||||
del prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
except Exception as e:
|
||||
debug_log(f"Collector probe failed for worker {wid}: {e}")
|
||||
except Exception as e:
|
||||
debug_log(f"Collector probe setup error: {e}")
|
||||
return any_busy
|
||||
|
||||
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
|
||||
try:
|
||||
combined = self._reorder_and_combine_tensors(
|
||||
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
|
||||
async def _drain_remaining_queue_items(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
worker_images: dict[str, dict[int, torch.Tensor]],
|
||||
mark_worker_done: Callable[[str], None],
|
||||
) -> int:
|
||||
collected_count = 0
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_jobs:
|
||||
log(f"Master - Queue {multi_job_id} no longer exists!")
|
||||
return collected_count
|
||||
|
||||
final_q = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
remaining_items = []
|
||||
while not final_q.empty():
|
||||
try:
|
||||
remaining_items.append(final_q.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
for item in remaining_items:
|
||||
worker_id = item['worker_id']
|
||||
is_last = item.get('is_last', False)
|
||||
collected_count += self._store_worker_result(worker_images, item)
|
||||
if is_last:
|
||||
mark_worker_done(worker_id)
|
||||
return collected_count
|
||||
|
||||
async def _collect_worker_results(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
enabled_workers: list[str],
|
||||
expected_workers: set[str],
|
||||
worker_images: dict[str, dict[int, torch.Tensor]],
|
||||
worker_audio: dict[str, dict[str, Any]],
|
||||
) -> int:
|
||||
num_workers = len(expected_workers)
|
||||
workers_done = set()
|
||||
collected_count = 0
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
|
||||
last_activity = time.time()
|
||||
progress = ProgressBar(num_workers)
|
||||
|
||||
def mark_worker_done(done_worker_id: str) -> None:
|
||||
done_worker_id = str(done_worker_id)
|
||||
if done_worker_id not in expected_workers:
|
||||
debug_log(
|
||||
f"Master - Ignoring completion from unexpected worker {done_worker_id} for job {multi_job_id}"
|
||||
)
|
||||
if combined is None:
|
||||
debug_log(f"Master - Job {multi_job_id} complete with audio only")
|
||||
else:
|
||||
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
|
||||
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
|
||||
return
|
||||
if done_worker_id in workers_done:
|
||||
debug_log(
|
||||
f"Master - Ignoring duplicate completion from worker {done_worker_id} for job {multi_job_id}"
|
||||
)
|
||||
return
|
||||
workers_done.add(done_worker_id)
|
||||
progress.update(1)
|
||||
|
||||
return (combined, combined_audio)
|
||||
except Exception as e:
|
||||
log(f"Master - Error combining images: {e}")
|
||||
# Preserve collected audio even when image assembly fails.
|
||||
return (images, combined_audio)
|
||||
while len(workers_done) < num_workers:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
try:
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
q = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
|
||||
result = await asyncio.wait_for(q.get(), timeout=slice_timeout)
|
||||
worker_id = result['worker_id']
|
||||
is_last = result.get('is_last', False)
|
||||
collected_count += self._store_worker_result(worker_images, result)
|
||||
debug_log(
|
||||
f"Master - Got canonical result from worker {worker_id}, "
|
||||
f"image {result.get('image_index', 0)}, is_last={is_last}"
|
||||
)
|
||||
|
||||
result_audio = result.get('audio')
|
||||
if result_audio is not None:
|
||||
worker_audio[worker_id] = result_audio
|
||||
debug_log(f"Master - Got audio from worker {worker_id}")
|
||||
|
||||
last_activity = time.time()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
if is_last:
|
||||
mark_worker_done(worker_id)
|
||||
except asyncio.TimeoutError:
|
||||
if (time.time() - last_activity) < base_timeout:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
continue
|
||||
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
missing_workers = set(str(w) for w in enabled_workers) - workers_done
|
||||
elapsed = time.time() - last_activity
|
||||
for missing_worker_id in sorted(missing_workers):
|
||||
log(
|
||||
"Master - Heartbeat timeout: "
|
||||
f"worker={missing_worker_id}, elapsed={elapsed:.1f}s"
|
||||
)
|
||||
log(
|
||||
f"Master - Heartbeat timeout. Still waiting for workers: {list(missing_workers)} "
|
||||
f"(elapsed={elapsed:.1f}s)"
|
||||
)
|
||||
|
||||
if await self._probe_missing_workers_busy(missing_workers):
|
||||
last_activity = time.time()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
continue
|
||||
|
||||
collected_count += await self._drain_remaining_queue_items(
|
||||
multi_job_id,
|
||||
worker_images,
|
||||
mark_worker_done,
|
||||
)
|
||||
break
|
||||
|
||||
return collected_count
|
||||
|
||||
async def _execute_master(
|
||||
self,
|
||||
images,
|
||||
audio,
|
||||
multi_job_id,
|
||||
enabled_worker_ids,
|
||||
delegate_only,
|
||||
):
|
||||
delegate_mode = delegate_only or is_master_delegate_only()
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
expected_workers = set(enabled_workers)
|
||||
num_workers = len(expected_workers)
|
||||
if num_workers == 0:
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
|
||||
await self._ensure_pending_queue(multi_job_id)
|
||||
master_batch_size, images_on_cpu, master_audio = self._build_master_inputs(
|
||||
images, audio, delegate_mode, multi_job_id, num_workers
|
||||
)
|
||||
|
||||
worker_images = {}
|
||||
worker_audio = {}
|
||||
try:
|
||||
await self._collect_worker_results(
|
||||
multi_job_id=multi_job_id,
|
||||
enabled_workers=enabled_workers,
|
||||
expected_workers=expected_workers,
|
||||
worker_images=worker_images,
|
||||
worker_audio=worker_audio,
|
||||
)
|
||||
except comfy.model_management.InterruptProcessingException:
|
||||
await self._cleanup_pending_queue(multi_job_id)
|
||||
raise
|
||||
|
||||
await self._cleanup_pending_queue(multi_job_id)
|
||||
try:
|
||||
combined = self._reorder_and_combine_tensors(
|
||||
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
|
||||
)
|
||||
debug_log(
|
||||
f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
|
||||
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})"
|
||||
)
|
||||
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
|
||||
return (combined, combined_audio)
|
||||
except Exception as e:
|
||||
log(f"Master - Error combining images: {e}")
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
|
||||
async def execute(self, *args: Any, **kwargs: Any) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""Compatibility wrapper with a uniform execute() signature across collectors."""
|
||||
return await self._execute_collector(*args, **kwargs)
|
||||
|
||||
async def _execute_collector(
|
||||
self,
|
||||
images: torch.Tensor,
|
||||
audio: dict[str, Any] | None,
|
||||
load_balance: bool = False,
|
||||
context: CollectorRunContext | None = None,
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
run_context = context or CollectorRunContext()
|
||||
_ = load_balance
|
||||
_ = run_context.worker_batch_size
|
||||
if run_context.is_worker:
|
||||
debug_log(
|
||||
"Worker - Job "
|
||||
f"{run_context.multi_job_id} complete. Sending {images.shape[0]} image(s) to master"
|
||||
)
|
||||
await self.send_batch_to_master(
|
||||
images,
|
||||
audio,
|
||||
run_context.multi_job_id,
|
||||
run_context.master_url,
|
||||
run_context.worker_id,
|
||||
)
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
|
||||
return await self._execute_master(
|
||||
images=images,
|
||||
audio=audio,
|
||||
multi_job_id=run_context.multi_job_id,
|
||||
enabled_worker_ids=run_context.enabled_worker_ids,
|
||||
delegate_only=run_context.delegate_only,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
from typing import Any
|
||||
|
||||
from ..utils.parsing import coerce_bool
|
||||
from ..utils.worker_ids import coerce_enabled_worker_ids
|
||||
|
||||
|
||||
def _parse_bool(value: Any, default: bool = False) -> bool:
|
||||
return coerce_bool(value, default=default)
|
||||
|
||||
|
||||
def _parse_enabled_workers(value: Any) -> list[str]:
|
||||
if value is None:
|
||||
return []
|
||||
return coerce_enabled_worker_ids(value)
|
||||
|
||||
|
||||
def parse_distributed_hidden_context(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Parse common distributed hidden inputs from node kwargs."""
|
||||
return {
|
||||
"multi_job_id": str(kwargs.get("multi_job_id", "") or ""),
|
||||
"is_worker": _parse_bool(kwargs.get("is_worker", False)),
|
||||
"master_url": str(kwargs.get("master_url", "") or ""),
|
||||
"enabled_worker_ids": _parse_enabled_workers(kwargs.get("enabled_worker_ids", "[]")),
|
||||
"worker_id": str(kwargs.get("worker_id", "") or ""),
|
||||
"delegate_only": _parse_bool(kwargs.get("delegate_only", False)),
|
||||
}
|
||||
+335
-88
@@ -1,25 +1,28 @@
|
||||
import json
|
||||
import math
|
||||
from functools import wraps
|
||||
from typing import Any, Callable
|
||||
|
||||
import comfy.samplers
|
||||
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.async_helpers import run_async_in_server_loop
|
||||
from ..upscale.job_store import ensure_tile_jobs_initialized
|
||||
from .hidden_inputs import build_distributed_hidden_inputs
|
||||
|
||||
from ..upscale.tile_ops import TileOpsMixin
|
||||
from ..upscale.result_collector import ResultCollectorMixin
|
||||
from ..upscale.worker_comms import WorkerCommsMixin
|
||||
from ..upscale.job_state import JobStateMixin
|
||||
from ..upscale.processing_args import UpscaleCoreArgs
|
||||
from ..upscale.modes.single_gpu import SingleGpuModeMixin
|
||||
from ..upscale.modes.static import StaticModeMixin
|
||||
from ..upscale.modes.dynamic import DynamicModeMixin
|
||||
from ..utils.worker_ids import coerce_enabled_worker_ids
|
||||
|
||||
def sync_wrapper(async_func):
|
||||
def sync_wrapper(async_func: Callable[..., Any]) -> Callable[..., Any]:
|
||||
"""Decorator to wrap async methods for synchronous execution."""
|
||||
@wraps(async_func)
|
||||
def sync_func(self, *args, **kwargs):
|
||||
def sync_func(self, *args: Any, **kwargs: Any) -> Any:
|
||||
# Use run_async_in_server_loop for ComfyUI compatibility
|
||||
return run_async_in_server_loop(
|
||||
async_func(self, *args, **kwargs),
|
||||
@@ -27,21 +30,10 @@ def sync_wrapper(async_func):
|
||||
)
|
||||
return sync_func
|
||||
|
||||
def _parse_enabled_worker_ids(enabled_worker_ids):
|
||||
"""Parse enabled worker IDs from either JSON or list input."""
|
||||
if isinstance(enabled_worker_ids, list):
|
||||
return [str(worker_id) for worker_id in enabled_worker_ids]
|
||||
if not enabled_worker_ids:
|
||||
return []
|
||||
if isinstance(enabled_worker_ids, str):
|
||||
try:
|
||||
parsed = json.loads(enabled_worker_ids)
|
||||
except json.JSONDecodeError:
|
||||
log("USDU Dist: Invalid enabled_worker_ids JSON; defaulting to no workers.")
|
||||
return []
|
||||
if isinstance(parsed, list):
|
||||
return [str(wid) for wid in parsed]
|
||||
return []
|
||||
|
||||
def _parse_enabled_worker_ids(enabled_worker_ids: str | list[str] | None) -> list[str]:
|
||||
"""Backward-compatible alias for enabled-worker normalization."""
|
||||
return coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
|
||||
class UltimateSDUpscaleDistributed(
|
||||
DynamicModeMixin,
|
||||
@@ -56,22 +48,22 @@ class UltimateSDUpscaleDistributed(
|
||||
"""
|
||||
Distributed version of Ultimate SD Upscale (No Upscale).
|
||||
|
||||
Supports two currently selected processing modes:
|
||||
Supports three processing modes:
|
||||
1. Single GPU: No workers available, process everything locally
|
||||
2. Distributed tile queue: Workers pull tiles from a shared queue
|
||||
2. Static Mode: Small batches, distributes tiles across workers (flattened)
|
||||
3. Dynamic Mode: Large batches, assigns whole images to workers dynamically
|
||||
|
||||
Features:
|
||||
- Tile-based batch handling for video/image upscaling
|
||||
- Multi-mode batch handling for efficient video/image upscaling
|
||||
- Tiled VAE support for memory efficiency
|
||||
- Shared work queue so faster workers can process more tiles
|
||||
- Dynamic load balancing for large batches
|
||||
- Backward compatible with single-image workflows
|
||||
|
||||
Environment Variables:
|
||||
- COMFYUI_MAX_BATCH: Chunk size for tile sending (default 20)
|
||||
- COMFYUI_MAX_PAYLOAD_SIZE: Max API payload bytes (default 50MB)
|
||||
|
||||
The hidden dynamic_threshold input is retained for workflow compatibility but
|
||||
does not affect the current mode-selection policy.
|
||||
Threshold: dynamic_threshold input controls mode switch (default 8)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
@@ -81,7 +73,14 @@ class UltimateSDUpscaleDistributed(
|
||||
debug_log("UltimateSDUpscaleDistributed - Node initialized")
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls: type["UltimateSDUpscaleDistributed"]) -> dict[str, Any]:
|
||||
hidden_inputs = build_distributed_hidden_inputs()
|
||||
hidden_inputs.update(
|
||||
{
|
||||
"tile_indices": ("STRING", {"default": ""}), # Unused - kept for compatibility
|
||||
"dynamic_threshold": ("INT", {"default": 8, "min": 1, "max": 64}),
|
||||
}
|
||||
)
|
||||
return {
|
||||
"required": {
|
||||
"upscaled_image": ("IMAGE",),
|
||||
@@ -102,15 +101,7 @@ class UltimateSDUpscaleDistributed(
|
||||
"force_uniform_tiles": ("BOOLEAN", {"default": True}),
|
||||
"tiled_decode": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"master_url": ("STRING", {"default": ""}),
|
||||
"enabled_worker_ids": ("STRING", {"default": "[]"}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
"tile_indices": ("STRING", {"default": ""}), # Unused - kept for compatibility
|
||||
"dynamic_threshold": ("INT", {"default": 8, "min": 1, "max": 64}),
|
||||
},
|
||||
"hidden": hidden_inputs,
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -122,12 +113,47 @@ class UltimateSDUpscaleDistributed(
|
||||
"""Force re-execution."""
|
||||
return float("nan") # Always re-execute
|
||||
|
||||
def run(self, upscaled_image, model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
|
||||
mask_blur, force_uniform_tiles, tiled_decode,
|
||||
multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]",
|
||||
worker_id="", tile_indices="", dynamic_threshold=8):
|
||||
def run(
|
||||
self,
|
||||
upscaled_image: Any,
|
||||
model: Any,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
vae: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
denoise: float,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
tiled_decode: bool,
|
||||
multi_job_id: str = "",
|
||||
is_worker: bool = False,
|
||||
master_url: str = "",
|
||||
enabled_worker_ids: str = "[]",
|
||||
worker_id: str = "",
|
||||
tile_indices: str = "",
|
||||
dynamic_threshold: int = 8,
|
||||
) -> tuple[Any, ...]:
|
||||
"""Entry point - runs SYNCHRONOUSLY like Ultimate SD Upscaler."""
|
||||
core_args = UpscaleCoreArgs(
|
||||
model=model,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
vae=vae,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
denoise=denoise,
|
||||
tiled_decode=tiled_decode,
|
||||
)
|
||||
# Strict WAN/FLOW batching: error if batch is not 4n+1 (except allow 1)
|
||||
try:
|
||||
batch_size = int(getattr(upscaled_image, 'shape', [1])[0])
|
||||
@@ -142,9 +168,15 @@ class UltimateSDUpscaleDistributed(
|
||||
)
|
||||
if not multi_job_id:
|
||||
# No distributed processing, run single GPU version
|
||||
return self.process_single_gpu(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode)
|
||||
return self.process_single_gpu(
|
||||
upscaled_image=upscaled_image,
|
||||
core_args=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mask_blur=mask_blur,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
)
|
||||
|
||||
if is_worker:
|
||||
# Worker mode: process tiles synchronously
|
||||
@@ -155,24 +187,64 @@ class UltimateSDUpscaleDistributed(
|
||||
worker_id, enabled_worker_ids, dynamic_threshold)
|
||||
else:
|
||||
# Master mode: distribute and collect synchronously
|
||||
return self.process_master(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, enabled_worker_ids,
|
||||
dynamic_threshold)
|
||||
return self.process_master(
|
||||
upscaled_image=upscaled_image,
|
||||
core_args=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mask_blur=mask_blur,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
multi_job_id=multi_job_id,
|
||||
enabled_worker_ids=enabled_worker_ids,
|
||||
dynamic_threshold=dynamic_threshold,
|
||||
)
|
||||
|
||||
def process_worker(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
worker_id, enabled_worker_ids, dynamic_threshold):
|
||||
def process_worker(
|
||||
self,
|
||||
upscaled_image: Any,
|
||||
model: Any,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
vae: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
denoise: float,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
tiled_decode: bool,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str,
|
||||
enabled_worker_ids: str,
|
||||
dynamic_threshold: int,
|
||||
) -> tuple[Any, ...]:
|
||||
"""Unified worker processing - handles both static and dynamic modes."""
|
||||
core_args = UpscaleCoreArgs(
|
||||
model=model,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
vae=vae,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
denoise=denoise,
|
||||
tiled_decode=tiled_decode,
|
||||
)
|
||||
# Get batch size to determine mode
|
||||
batch_size = upscaled_image.shape[0]
|
||||
|
||||
# Ensure mode consistency across master/workers via shared threshold
|
||||
# Determine mode (must match master's logic)
|
||||
enabled_workers = json.loads(enabled_worker_ids)
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
num_workers = len(enabled_workers)
|
||||
# Compute number of tiles for this image to decide if tile distribution makes sense
|
||||
_, height, width, _ = upscaled_image.shape
|
||||
@@ -181,31 +253,47 @@ class UltimateSDUpscaleDistributed(
|
||||
|
||||
mode = self._determine_processing_mode(batch_size, num_workers, dynamic_threshold)
|
||||
# For USDU-style processing, we want tile distribution whenever workers are available
|
||||
# and there is more than one tile to process, even if batch == 1.
|
||||
if num_workers > 0 and num_tiles_per_image > 1:
|
||||
# and there is more than one tile to process for single-image runs.
|
||||
if num_workers > 0 and batch_size <= 1 and num_tiles_per_image > 1:
|
||||
mode = "static"
|
||||
|
||||
debug_log(f"USDU Dist Worker - Batch size {batch_size}")
|
||||
|
||||
if mode == "dynamic":
|
||||
return self.process_worker_dynamic(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
worker_id, enabled_worker_ids, dynamic_threshold)
|
||||
return self.process_worker_dynamic(
|
||||
upscaled_image=upscaled_image,
|
||||
core_args=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mask_blur=mask_blur,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
multi_job_id=multi_job_id,
|
||||
master_url=master_url,
|
||||
worker_id=worker_id,
|
||||
enabled_worker_ids=enabled_worker_ids,
|
||||
dynamic_threshold=dynamic_threshold,
|
||||
)
|
||||
|
||||
# Static mode - enhanced with health monitoring and retry logic
|
||||
return self._process_worker_static_sync(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
return self._process_worker_static_sync(upscaled_image, core_args,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
force_uniform_tiles, multi_job_id, master_url,
|
||||
worker_id, enabled_workers)
|
||||
|
||||
def process_master(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, enabled_worker_ids,
|
||||
dynamic_threshold):
|
||||
def process_master(
|
||||
self,
|
||||
upscaled_image: Any,
|
||||
core_args: UpscaleCoreArgs,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
multi_job_id: str,
|
||||
enabled_worker_ids: str,
|
||||
dynamic_threshold: int,
|
||||
) -> tuple[Any, ...]:
|
||||
"""Unified master processing with enhanced monitoring and failure handling."""
|
||||
# Round tile dimensions
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
@@ -224,47 +312,206 @@ class UltimateSDUpscaleDistributed(
|
||||
)
|
||||
|
||||
# Parse enabled workers
|
||||
enabled_workers = json.loads(enabled_worker_ids)
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
num_workers = len(enabled_workers)
|
||||
|
||||
# Determine processing mode
|
||||
mode = self._determine_processing_mode(batch_size, num_workers, dynamic_threshold)
|
||||
# Prefer tile-based static distribution when workers are available and there are multiple tiles,
|
||||
# even for batch == 1, to spread tiles across GPUs like the legacy dynamic tile queue.
|
||||
if num_workers > 0 and num_tiles_per_image > 1:
|
||||
# for single-image jobs to spread tiles across GPUs like the legacy tile queue.
|
||||
if num_workers > 0 and batch_size <= 1 and num_tiles_per_image > 1:
|
||||
mode = "static"
|
||||
|
||||
log(f"USDU Dist: Workers {num_workers} | Mode {mode} | Threshold {dynamic_threshold}")
|
||||
|
||||
if mode == "single_gpu":
|
||||
# No workers, process all tiles locally
|
||||
return self.process_single_gpu(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode)
|
||||
return self.process_single_gpu(
|
||||
upscaled_image=upscaled_image,
|
||||
core_args=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mask_blur=mask_blur,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
)
|
||||
|
||||
elif mode == "dynamic":
|
||||
# Dynamic mode for large batches
|
||||
return self.process_master_dynamic(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers)
|
||||
return self.process_master_dynamic(
|
||||
upscaled_image=upscaled_image,
|
||||
core_args=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mask_blur=mask_blur,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
multi_job_id=multi_job_id,
|
||||
enabled_workers=enabled_workers,
|
||||
)
|
||||
|
||||
# Static mode - enhanced with unified job management
|
||||
return self._process_master_static_sync(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
return self._process_master_static_sync(upscaled_image, core_args,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers,
|
||||
force_uniform_tiles, multi_job_id, enabled_workers,
|
||||
all_tiles, num_tiles_per_image)
|
||||
|
||||
def _determine_processing_mode(self, batch_size: int, num_workers: int, dynamic_threshold: int) -> str:
|
||||
"""Determines processing mode per requested policy:
|
||||
- any workers => prefer static (tile-based) for USDU
|
||||
- no workers => single_gpu
|
||||
"""
|
||||
"""Determine mode from worker availability and configured threshold."""
|
||||
if num_workers == 0:
|
||||
return "single_gpu"
|
||||
# Default to static when distributed; master/worker may still override if special cases arise
|
||||
threshold = max(1, int(dynamic_threshold))
|
||||
if int(batch_size) >= threshold:
|
||||
return "dynamic"
|
||||
return "static"
|
||||
|
||||
# Ensure initialization before registering routes
|
||||
ensure_tile_jobs_initialized()
|
||||
class USDUDelegateCollector:
|
||||
"""Lightweight stand-in used automatically in delegate-only master prompts.
|
||||
|
||||
When the master is in delegate-only mode, the orchestration code swaps the
|
||||
full ``UltimateSDUpscaleDistributed`` class for this one so that no upstream
|
||||
model/image nodes execute on the master. Workers initialise the dynamic job
|
||||
queue via ``/distributed/init_dynamic_job`` and this node simply waits for
|
||||
all images to arrive, then assembles the output tensor.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {},
|
||||
"hidden": {
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
"enabled_worker_ids": ("STRING", {"default": "[]"}),
|
||||
"delegate_only": ("BOOLEAN", {"default": True}),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
"master_url": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "image/upscaling"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs):
|
||||
return float("nan")
|
||||
|
||||
def run(self, multi_job_id="", enabled_worker_ids="[]", **kwargs):
|
||||
import time
|
||||
import torch
|
||||
from ..utils.image import pil_to_tensor
|
||||
from ..upscale.job_models import ImageJobState
|
||||
|
||||
if not multi_job_id:
|
||||
log("USDUDelegateCollector: no multi_job_id, returning empty image")
|
||||
return (torch.zeros(1, 64, 64, 3),)
|
||||
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
num_workers = len(enabled_workers)
|
||||
log(f"USDU delegate-only: waiting for workers to init job {multi_job_id}")
|
||||
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
|
||||
# Wait for workers to create the job via /distributed/init_dynamic_job
|
||||
max_wait = 120.0
|
||||
poll_interval = 1.0
|
||||
start = time.time()
|
||||
job_data = None
|
||||
while time.time() - start < max_wait:
|
||||
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
|
||||
if isinstance(job_data, ImageJobState) and job_data.batch_size > 0:
|
||||
break
|
||||
job_data = None
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if job_data is None:
|
||||
log(f"USDUDelegateCollector: job {multi_job_id} was not initialized by workers within {max_wait}s")
|
||||
return (torch.zeros(1, 64, 64, 3),)
|
||||
|
||||
batch_size = job_data.batch_size
|
||||
log(f"USDU delegate-only: job ready, collecting {batch_size} images from {num_workers} workers")
|
||||
|
||||
# Collect all images using the existing result collector
|
||||
collected = run_async_in_server_loop(
|
||||
self._collect_all_images(multi_job_id, batch_size, num_workers, prompt_server),
|
||||
timeout=600.0,
|
||||
)
|
||||
|
||||
# Assemble tensor from collected PIL images
|
||||
result_images = []
|
||||
for idx in range(batch_size):
|
||||
img = collected.get(idx)
|
||||
if img is not None:
|
||||
result_images.append(pil_to_tensor(img))
|
||||
else:
|
||||
log(f"USDUDelegateCollector: missing image {idx}, using blank")
|
||||
result_images.append(torch.zeros(1, 64, 64, 3))
|
||||
|
||||
result_tensor = torch.cat(result_images, dim=0)
|
||||
log(f"USDU delegate-only: collected all {batch_size} images")
|
||||
return (result_tensor,)
|
||||
|
||||
@staticmethod
|
||||
async def _collect_all_images(multi_job_id, batch_size, num_workers, prompt_server):
|
||||
"""Wait for all images to arrive from workers."""
|
||||
import asyncio
|
||||
from ..upscale.job_models import ImageJobState
|
||||
from ..upscale.job_timeout import check_and_requeue_timed_out_workers
|
||||
|
||||
timeout_seconds = 300.0
|
||||
poll_interval = 2.0
|
||||
start = asyncio.get_event_loop().time()
|
||||
|
||||
while True:
|
||||
# Drain results from the queue
|
||||
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
|
||||
if isinstance(job_data, ImageJobState):
|
||||
# Drain queue items into completed_images
|
||||
while True:
|
||||
try:
|
||||
result = job_data.queue.get_nowait()
|
||||
worker_id = result.get("worker_id")
|
||||
if "image_idx" in result and "image" in result:
|
||||
idx = result["image_idx"]
|
||||
if idx not in job_data.completed_images:
|
||||
job_data.completed_images[idx] = result["image"]
|
||||
debug_log(f"Delegate collected image {idx} from worker {worker_id}")
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
completed_count = len(job_data.completed_images)
|
||||
if completed_count >= batch_size:
|
||||
debug_log(f"Delegate: all {batch_size} images collected")
|
||||
completed = dict(job_data.completed_images)
|
||||
# Cleanup
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
prompt_server.distributed_pending_tile_jobs.pop(multi_job_id, None)
|
||||
return completed
|
||||
|
||||
# Check for timeouts periodically
|
||||
await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
|
||||
|
||||
elapsed = asyncio.get_event_loop().time() - start
|
||||
if elapsed > timeout_seconds:
|
||||
log(f"USDUDelegateCollector: timed out after {elapsed:.0f}s with {len(getattr(job_data, 'completed_images', {}))} of {batch_size} images")
|
||||
if isinstance(job_data, ImageJobState):
|
||||
completed = dict(job_data.completed_images)
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
prompt_server.distributed_pending_tile_jobs.pop(multi_job_id, None)
|
||||
return completed
|
||||
return {}
|
||||
|
||||
await asyncio.sleep(poll_interval)
|
||||
|
||||
|
||||
# Node registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": UltimateSDUpscaleDistributed,
|
||||
"USDUDelegateCollector": USDUDelegateCollector,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": "Ultimate SD Upscale Distributed (No Upscale)",
|
||||
# No display name for USDUDelegateCollector — it's an internal node
|
||||
}
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def build_worker_identity_hidden_inputs() -> dict[str, tuple[str, dict[str, Any]]]:
|
||||
"""Hidden-input contract for worker identity context."""
|
||||
return {
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"enabled_worker_ids": ("STRING", {"default": "[]"}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
}
|
||||
|
||||
|
||||
def build_distributed_hidden_inputs(
|
||||
*,
|
||||
include_assigned_branch: bool = False,
|
||||
assigned_branch_max: int = 9,
|
||||
include_worker_batch_size: bool = False,
|
||||
include_pass_through: bool = False,
|
||||
) -> dict[str, tuple[str, dict[str, Any]]]:
|
||||
hidden_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
"master_url": ("STRING", {"default": ""}),
|
||||
"delegate_only": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
hidden_inputs.update(build_worker_identity_hidden_inputs())
|
||||
if include_assigned_branch:
|
||||
hidden_inputs["assigned_branch"] = (
|
||||
"INT",
|
||||
{"default": -1, "min": -1, "max": assigned_branch_max},
|
||||
)
|
||||
if include_worker_batch_size:
|
||||
hidden_inputs["worker_batch_size"] = (
|
||||
"INT",
|
||||
{"default": 1, "min": 1, "max": 1024},
|
||||
)
|
||||
if include_pass_through:
|
||||
hidden_inputs["pass_through"] = ("BOOLEAN", {"default": False})
|
||||
return hidden_inputs
|
||||
@@ -0,0 +1,46 @@
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from ..utils.config import get_worker_timeout_seconds
|
||||
from ..utils.constants import HEARTBEAT_INTERVAL
|
||||
from ..utils.logging import log
|
||||
|
||||
|
||||
async def collect_worker_queue_results(
|
||||
*,
|
||||
prompt_server: Any,
|
||||
multi_job_id: str,
|
||||
expected_workers: set[str],
|
||||
on_result: Callable[[dict[str, Any]], None],
|
||||
timeout_log_prefix: str,
|
||||
throw_if_interrupted: Callable[[], None],
|
||||
) -> set[str]:
|
||||
workers_done: set[str] = set()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
|
||||
last_activity = time.time()
|
||||
|
||||
while len(workers_done) < len(expected_workers):
|
||||
throw_if_interrupted()
|
||||
try:
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
queue = prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
result = await asyncio.wait_for(queue.get(), timeout=slice_timeout)
|
||||
except asyncio.TimeoutError:
|
||||
if (time.time() - last_activity) < base_timeout:
|
||||
continue
|
||||
missing_workers = sorted(expected_workers - workers_done)
|
||||
log(f"{timeout_log_prefix}{missing_workers}")
|
||||
break
|
||||
|
||||
on_result(result)
|
||||
worker_id_value = str(result.get("worker_id", ""))
|
||||
is_last = bool(result.get("is_last", False))
|
||||
last_activity = time.time()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
if is_last and worker_id_value in expected_workers:
|
||||
workers_done.add(worker_id_value)
|
||||
|
||||
return workers_done
|
||||
@@ -0,0 +1,16 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def get_prompt_server_instance() -> Any:
|
||||
import server as _server
|
||||
|
||||
return _server.PromptServer.instance
|
||||
|
||||
|
||||
def throw_if_processing_interrupted() -> None:
|
||||
try:
|
||||
import comfy.model_management as model_management
|
||||
|
||||
model_management.throw_exception_if_processing_interrupted()
|
||||
except Exception:
|
||||
return
|
||||
+75
-50
@@ -1,7 +1,10 @@
|
||||
import torch
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.worker_ids import coerce_enabled_worker_ids, parse_worker_index, worker_value_key
|
||||
from .hidden_inputs import build_worker_identity_hidden_inputs
|
||||
|
||||
|
||||
def _chunk_bounds(total_items: int, n_splits: int) -> list[tuple[int, int]]:
|
||||
@@ -28,7 +31,7 @@ class DistributedSeed:
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
def INPUT_TYPES(cls: type["DistributedSeed"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"seed": ("INT", {
|
||||
@@ -39,8 +42,7 @@ class DistributedSeed:
|
||||
}),
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
**build_worker_identity_hidden_inputs(),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -49,30 +51,29 @@ class DistributedSeed:
|
||||
FUNCTION = "distribute"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def distribute(self, seed, is_worker=False, worker_id=""):
|
||||
def distribute(
|
||||
self,
|
||||
seed: int,
|
||||
is_worker: bool = False,
|
||||
worker_id: str = "",
|
||||
enabled_worker_ids: str = "[]",
|
||||
) -> tuple[int]:
|
||||
if not is_worker:
|
||||
# Master node: pass through original values
|
||||
debug_log(f"Distributor - Master: seed={seed}")
|
||||
return (seed,)
|
||||
else:
|
||||
# Worker node: apply offset based on worker index
|
||||
# Find worker index from enabled_worker_ids
|
||||
try:
|
||||
# Worker IDs are passed as "worker_0", "worker_1", etc.
|
||||
if worker_id.startswith("worker_"):
|
||||
worker_index = int(worker_id.split("_")[1])
|
||||
else:
|
||||
# Fallback: try to parse as direct index
|
||||
worker_index = int(worker_id)
|
||||
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
worker_index = parse_worker_index(worker_id, enabled_workers)
|
||||
if worker_index is not None:
|
||||
offset = worker_index + 1
|
||||
new_seed = seed + offset
|
||||
debug_log(f"Distributor - Worker {worker_index}: seed={seed} → {new_seed}")
|
||||
return (new_seed,)
|
||||
except (ValueError, IndexError) as e:
|
||||
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': {e}")
|
||||
# Fallback: return original seed
|
||||
return (seed,)
|
||||
|
||||
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': no worker index resolved")
|
||||
# Fallback: return original seed
|
||||
return (seed,)
|
||||
|
||||
|
||||
# Define ByPassTypeTuple for flexible return types
|
||||
@@ -92,15 +93,14 @@ class DistributedValue:
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
def INPUT_TYPES(cls: type["DistributedValue"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"default_value": ("STRING", {"default": ""}),
|
||||
"worker_values": ("STRING", {"default": "{}"}),
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
**build_worker_identity_hidden_inputs(),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -110,7 +110,7 @@ class DistributedValue:
|
||||
CATEGORY = "utils"
|
||||
|
||||
@staticmethod
|
||||
def _coerce(value, value_type):
|
||||
def _coerce(value: Any, value_type: str) -> Any:
|
||||
"""Convert a string value to the requested type."""
|
||||
if value_type == "INT":
|
||||
return int(float(value))
|
||||
@@ -119,51 +119,71 @@ class DistributedValue:
|
||||
return value # STRING and COMBO stay as strings
|
||||
|
||||
@staticmethod
|
||||
def _coerce_safe(value, value_type):
|
||||
def _coerce_safe(value: Any, value_type: str) -> Any:
|
||||
"""Best-effort coercion with graceful fallback to original value."""
|
||||
try:
|
||||
return DistributedValue._coerce(value, value_type)
|
||||
except (TypeError, ValueError):
|
||||
return value
|
||||
|
||||
def distribute(self, default_value, worker_values="{}", is_worker=False, worker_id=""):
|
||||
@staticmethod
|
||||
def _infer_value_type(value: Any) -> str:
|
||||
"""Infer coercion type from the provided default value."""
|
||||
if isinstance(value, bool):
|
||||
return "STRING"
|
||||
if isinstance(value, int):
|
||||
return "INT"
|
||||
if isinstance(value, float):
|
||||
return "FLOAT"
|
||||
return "STRING"
|
||||
|
||||
def distribute(
|
||||
self,
|
||||
default_value: Any,
|
||||
worker_values: str | dict[str, Any] = "{}",
|
||||
is_worker: bool = False,
|
||||
worker_id: str = "",
|
||||
enabled_worker_ids: str = "[]",
|
||||
) -> tuple[Any]:
|
||||
values = {}
|
||||
value_type = "STRING"
|
||||
|
||||
try:
|
||||
values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
|
||||
raw_values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
|
||||
values = dict(raw_values) if isinstance(raw_values, dict) else {}
|
||||
if not isinstance(values, dict):
|
||||
values = {}
|
||||
except json.JSONDecodeError as e:
|
||||
debug_log(f"DistributedValue - Error parsing worker_values: {e}")
|
||||
values = {}
|
||||
|
||||
value_type = values.get("_type", "STRING")
|
||||
inferred_type = self._infer_value_type(default_value)
|
||||
value_type = values.get("_type", inferred_type)
|
||||
if value_type not in {"STRING", "COMBO", "INT", "FLOAT"}:
|
||||
value_type = inferred_type
|
||||
coerced_default = self._coerce_safe(default_value, value_type)
|
||||
|
||||
if not is_worker:
|
||||
debug_log(f"DistributedValue - Master: returning default '{coerced_default}'")
|
||||
return (coerced_default,)
|
||||
|
||||
try:
|
||||
if worker_id.startswith("worker_"):
|
||||
idx = int(worker_id.split("_")[1])
|
||||
else:
|
||||
idx = int(worker_id)
|
||||
key = str(idx + 1) # worker_0 → key "1" (1-indexed)
|
||||
raw = values.get(key, "")
|
||||
if raw:
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
direct_key = str(worker_id).strip()
|
||||
lookup_key = direct_key if direct_key in values else worker_value_key(worker_id, enabled_workers)
|
||||
raw = values.get(lookup_key)
|
||||
if raw is not None and raw != "":
|
||||
try:
|
||||
coerced = self._coerce(raw, value_type)
|
||||
debug_log(f"DistributedValue - Worker {idx}: returning '{coerced}'")
|
||||
debug_log(f"DistributedValue - Worker key {lookup_key}: returning '{coerced}'")
|
||||
return (coerced,)
|
||||
except (ValueError, IndexError) as e:
|
||||
debug_log(f"DistributedValue - Error: {e}")
|
||||
except (TypeError, ValueError) as e:
|
||||
debug_log(f"DistributedValue - Error coercing worker value for key {lookup_key}: {e}")
|
||||
debug_log(f"DistributedValue - Worker fallback: returning default '{coerced_default}'")
|
||||
return (coerced_default,)
|
||||
|
||||
class DistributedModelName:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
def INPUT_TYPES(cls: type["DistributedModelName"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"default": ""}),
|
||||
@@ -180,7 +200,7 @@ class DistributedModelName:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "utils"
|
||||
|
||||
def _stringify(self, value):
|
||||
def _stringify(self, value: Any) -> str:
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, (int, float, bool)):
|
||||
@@ -190,7 +210,7 @@ class DistributedModelName:
|
||||
except Exception:
|
||||
return str(value)
|
||||
|
||||
def _update_workflow(self, extra_pnginfo, unique_id, values):
|
||||
def _update_workflow(self, extra_pnginfo: Any, unique_id: Any, values: list[str]) -> None:
|
||||
if not extra_pnginfo:
|
||||
return
|
||||
info = extra_pnginfo[0] if isinstance(extra_pnginfo, list) else extra_pnginfo
|
||||
@@ -208,7 +228,12 @@ class DistributedModelName:
|
||||
if node:
|
||||
node["widgets_values"] = [values]
|
||||
|
||||
def log_input(self, text, unique_id=None, extra_pnginfo=None):
|
||||
def log_input(
|
||||
self,
|
||||
text: Any,
|
||||
unique_id: Any = None,
|
||||
extra_pnginfo: Any = None,
|
||||
) -> dict[str, Any]:
|
||||
values = []
|
||||
if isinstance(text, list):
|
||||
for val in text:
|
||||
@@ -224,7 +249,7 @@ class DistributedModelName:
|
||||
return {"ui": {"text": values}, "result": (values,)}
|
||||
|
||||
class ByPassTypeTuple(tuple):
|
||||
def __getitem__(self, index):
|
||||
def __getitem__(self, index: int) -> Any:
|
||||
if index > 0:
|
||||
index = 0
|
||||
item = super().__getitem__(index)
|
||||
@@ -234,7 +259,7 @@ class ByPassTypeTuple(tuple):
|
||||
|
||||
class ImageBatchDivider:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls: type["ImageBatchDivider"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -255,7 +280,7 @@ class ImageBatchDivider:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "image"
|
||||
|
||||
def divide_batch(self, images, divide_by):
|
||||
def divide_batch(self, images: torch.Tensor, divide_by: int) -> tuple[torch.Tensor, ...]:
|
||||
total_splits = max(1, min(int(divide_by), 10))
|
||||
total_frames = images.shape[0]
|
||||
empty_tensor = images[:0]
|
||||
@@ -269,10 +294,10 @@ class ImageBatchDivider:
|
||||
|
||||
|
||||
class AudioBatchDivider:
|
||||
"""Divides an audio waveform into sequential segments along the time dimension."""
|
||||
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls: type["AudioBatchDivider"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
@@ -282,7 +307,7 @@ class AudioBatchDivider:
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of sequential time segments to create"
|
||||
"tooltip": "Number of parts to divide the audio into"
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -293,7 +318,7 @@ class AudioBatchDivider:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "audio"
|
||||
|
||||
def divide_audio(self, audio, divide_by):
|
||||
def divide_audio(self, audio: dict[str, Any], divide_by: int) -> tuple[dict[str, Any], ...]:
|
||||
import torch
|
||||
|
||||
waveform = audio.get("waveform")
|
||||
@@ -333,7 +358,7 @@ class DistributedEmptyImage:
|
||||
"""Produces an empty IMAGE batch used when the master delegates all work."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
def INPUT_TYPES(cls: type["DistributedEmptyImage"]) -> dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"height": ("INT", {"default": 64, "min": 1, "max": 4096, "step": 1}),
|
||||
@@ -346,7 +371,7 @@ class DistributedEmptyImage:
|
||||
FUNCTION = "create"
|
||||
CATEGORY = "image"
|
||||
|
||||
def create(self, height, width, channels):
|
||||
def create(self, height: int, width: int, channels: int) -> tuple[torch.Tensor]:
|
||||
import torch
|
||||
|
||||
shape = (0, height, width, channels)
|
||||
|
||||
-181
@@ -1,181 +0,0 @@
|
||||
"""Native V3 schemas with the existing execution algorithms kept intact.
|
||||
|
||||
Runtime objects are private, per-execution helpers, not registered V1 nodes.
|
||||
This avoids sharing mutable instance state through V3's sanitized class clones.
|
||||
"""
|
||||
import comfy.samplers
|
||||
from comfy_api.v0_0_2 import io
|
||||
|
||||
from . import utilities as _utilities
|
||||
from .collector import DistributedCollectorNode as _CollectorRuntime
|
||||
from .distributed_upscale import UltimateSDUpscaleDistributed as _UpscaleRuntime
|
||||
|
||||
|
||||
class DistributedSeed(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedSeed', display_name='Distributed Seed', category='utils',
|
||||
inputs=[io.Int.Input('seed', default=1125899906842, min=0,
|
||||
max=1125899906842624, force_input=False)],
|
||||
outputs=[io.Int.Output(display_name='seed')],
|
||||
accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, seed, is_worker=False, worker_id=''):
|
||||
return io.NodeOutput(*_utilities.DistributedSeed().distribute(seed, is_worker, worker_id))
|
||||
|
||||
|
||||
class DistributedValue(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedValue', display_name='Distributed Value', category='utils',
|
||||
inputs=[io.String.Input('default_value', default=''),
|
||||
io.String.Input('worker_values', default='{}')],
|
||||
outputs=[io.AnyType.Output(display_name='value')],
|
||||
accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, default_value, worker_values='{}', is_worker=False, worker_id=''):
|
||||
return io.NodeOutput(*_utilities.DistributedValue().distribute(
|
||||
default_value, worker_values, is_worker, worker_id))
|
||||
|
||||
|
||||
class DistributedModelName(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedModelName', display_name='Distributed Model Name', category='utils',
|
||||
inputs=[io.String.Input('text', default='')],
|
||||
outputs=[io.AnyType.Output(display_name='output')],
|
||||
hidden=[io.Hidden.unique_id, io.Hidden.extra_pnginfo], is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, text):
|
||||
result = _utilities.DistributedModelName().log_input(
|
||||
text, unique_id=cls.hidden.unique_id, extra_pnginfo=cls.hidden.extra_pnginfo)
|
||||
return io.NodeOutput(*result['result'], ui=result['ui'])
|
||||
|
||||
|
||||
class ImageBatchDivider(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='ImageBatchDivider', display_name='Image Batch Divider', category='image',
|
||||
inputs=[io.Image.Input('images'),
|
||||
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
|
||||
display_mode=io.NumberDisplay.number,
|
||||
tooltip='Number of parts to divide the batch into')],
|
||||
# The existing frontend still displays only divide_by sockets.
|
||||
outputs=[io.Image.Output(display_name=f'batch_{index + 1}') for index in range(10)],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images, divide_by):
|
||||
return io.NodeOutput(*_utilities.ImageBatchDivider().divide_batch(images, divide_by))
|
||||
|
||||
|
||||
class AudioBatchDivider(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='AudioBatchDivider', display_name='Audio Segment Divider', category='audio',
|
||||
inputs=[io.Audio.Input('audio'),
|
||||
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
|
||||
display_mode=io.NumberDisplay.number,
|
||||
tooltip='Number of sequential time segments to create')],
|
||||
outputs=[io.Audio.Output(display_name=f'audio_{index + 1}') for index in range(10)],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, audio, divide_by):
|
||||
return io.NodeOutput(*_utilities.AudioBatchDivider().divide_audio(audio, divide_by))
|
||||
|
||||
|
||||
class DistributedEmptyImage(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedEmptyImage', display_name='Distributed Empty Image', category='image',
|
||||
inputs=[io.Int.Input('height', default=64, min=1, max=4096, step=1),
|
||||
io.Int.Input('width', default=64, min=1, max=4096, step=1),
|
||||
io.Int.Input('channels', default=3, min=1, max=4, step=1)],
|
||||
outputs=[io.Image.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, height, width, channels):
|
||||
return io.NodeOutput(*_utilities.DistributedEmptyImage().create(height, width, channels))
|
||||
|
||||
|
||||
class DistributedCollector(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedCollector', display_name='Distributed Collector', category='image',
|
||||
inputs=[io.Boolean.Input('load_balance', default=False,
|
||||
tooltip='Run this workflow on one least-busy participant (master included when participating).'),
|
||||
io.Image.Input('images', optional=True), io.Audio.Input('audio', optional=True)],
|
||||
outputs=[io.Image.Output(display_name='images'), io.Audio.Output(display_name='audio')],
|
||||
is_input_list=True, accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images=None, load_balance=False, audio=None, multi_job_id='',
|
||||
is_worker=False, master_url='', enabled_worker_ids='[]', worker_batch_size=1,
|
||||
worker_id='', pass_through=False, delegate_only=False):
|
||||
return io.NodeOutput(*_CollectorRuntime().run(
|
||||
images=images, load_balance=load_balance, audio=audio, multi_job_id=multi_job_id,
|
||||
is_worker=is_worker, master_url=master_url, enabled_worker_ids=enabled_worker_ids,
|
||||
worker_batch_size=worker_batch_size, worker_id=worker_id,
|
||||
pass_through=pass_through, delegate_only=delegate_only))
|
||||
|
||||
|
||||
class UltimateSDUpscaleDistributed(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='UltimateSDUpscaleDistributed',
|
||||
display_name='Ultimate SD Upscale Distributed (No Upscale)', category='image/upscaling',
|
||||
inputs=[
|
||||
io.Image.Input('upscaled_image'), io.Model.Input('model'),
|
||||
io.Conditioning.Input('positive'), io.Conditioning.Input('negative'), io.Vae.Input('vae'),
|
||||
io.Int.Input('seed', default=0, min=0, max=0xffffffffffffffff),
|
||||
io.Int.Input('steps', default=20, min=1, max=10000),
|
||||
io.Float.Input('cfg', default=8.0, min=0.0, max=100.0),
|
||||
io.Combo.Input('sampler_name', options=comfy.samplers.KSampler.SAMPLERS),
|
||||
io.Combo.Input('scheduler', options=comfy.samplers.KSampler.SCHEDULERS),
|
||||
io.Float.Input('denoise', default=0.5, min=0.0, max=1.0, step=0.01),
|
||||
io.Int.Input('tile_width', default=512, min=64, max=2048, step=8),
|
||||
io.Int.Input('tile_height', default=512, min=64, max=2048, step=8),
|
||||
io.Int.Input('padding', default=32, min=0, max=256, step=8),
|
||||
io.Int.Input('mask_blur', default=8, min=0, max=256),
|
||||
io.Boolean.Input('force_uniform_tiles', default=True),
|
||||
io.Boolean.Input('tiled_decode', default=False),
|
||||
], outputs=[io.Image.Output()], accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def fingerprint_inputs(cls, **kwargs):
|
||||
return _UpscaleRuntime.IS_CHANGED(**kwargs)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, upscaled_image, model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
|
||||
mask_blur, force_uniform_tiles, tiled_decode, multi_job_id='', is_worker=False,
|
||||
master_url='', enabled_worker_ids='[]', worker_id='', tile_indices='', dynamic_threshold=8):
|
||||
return io.NodeOutput(*_UpscaleRuntime().run(
|
||||
upscaled_image, model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
|
||||
mask_blur, force_uniform_tiles, tiled_decode, multi_job_id, is_worker,
|
||||
master_url, enabled_worker_ids, worker_id, tile_indices, dynamic_threshold))
|
||||
|
||||
|
||||
NODES = [DistributedCollector, DistributedSeed, DistributedModelName, DistributedValue,
|
||||
ImageBatchDivider, AudioBatchDivider, DistributedEmptyImage, UltimateSDUpscaleDistributed]
|
||||
+5
-2
@@ -1,9 +1,12 @@
|
||||
[project]
|
||||
name = "ComfyUI-Distributed"
|
||||
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
|
||||
version = "1.5.0"
|
||||
version = "1.4.1"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = []
|
||||
dependencies = [
|
||||
"aiohttp>=3.9,<4",
|
||||
"Pillow>=10,<12",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/robertvoy/ComfyUI-Distributed"
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""Internal extension lifecycle support (not a node provider)."""
|
||||
@@ -1,35 +0,0 @@
|
||||
"""Initialize routes, distributed state and the existing worker lifecycle."""
|
||||
import atexit
|
||||
import os
|
||||
|
||||
import server
|
||||
|
||||
from ..utils.config import CONFIG_FILE, ensure_config_exists
|
||||
from ..utils.logging import debug_log
|
||||
from ..workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
from ..upscale.job_store import ensure_tile_jobs_initialized
|
||||
|
||||
_initialized = False
|
||||
|
||||
|
||||
def initialize():
|
||||
"""Called by ComfyExtension.on_load; initialize once per loaded package."""
|
||||
global _initialized
|
||||
if _initialized:
|
||||
return
|
||||
|
||||
ensure_config_exists()
|
||||
from .. import api # noqa: F401 - registers the existing @routes.* handlers
|
||||
from ..api.queue_orchestration import ensure_distributed_state
|
||||
|
||||
ensure_distributed_state(server.PromptServer.instance)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
if not os.environ.get('COMFYUI_IS_WORKER'):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
|
||||
_initialized = True
|
||||
debug_log('Loaded Distributed nodes.')
|
||||
debug_log(f'Config file: {CONFIG_FILE}')
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Shared stubs/utilities for API route module unit tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
from typing import Any, Callable
|
||||
|
||||
|
||||
class RoutesStub:
|
||||
"""Simple decorator-style route registrar used in unit tests."""
|
||||
|
||||
def get(self, _path: str):
|
||||
return lambda fn: fn
|
||||
|
||||
def post(self, _path: str):
|
||||
return lambda fn: fn
|
||||
|
||||
|
||||
class ClientTimeoutStub:
|
||||
def __init__(self, total: float | None = None):
|
||||
self.total = total
|
||||
|
||||
|
||||
class WSMsgTypeStub:
|
||||
TEXT = "TEXT"
|
||||
ERROR = "ERROR"
|
||||
CLOSED = "CLOSED"
|
||||
|
||||
|
||||
class WebSocketResponseStub:
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
self.args = args
|
||||
self.kwargs = kwargs
|
||||
|
||||
async def prepare(self, _request: Any) -> None:
|
||||
return None
|
||||
|
||||
async def send_json(self, _payload: Any) -> None:
|
||||
return None
|
||||
|
||||
def __aiter__(self):
|
||||
async def _empty():
|
||||
if False:
|
||||
yield None
|
||||
|
||||
return _empty()
|
||||
|
||||
|
||||
class FormDataStub:
|
||||
def add_field(self, *_args: Any, **_kwargs: Any) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def reset_package_namespace(package_name: str) -> None:
|
||||
"""Clear test package modules from sys.modules."""
|
||||
prefix = f"{package_name}."
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(prefix):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
|
||||
def ensure_namespace_package(module_name: str) -> types.ModuleType:
|
||||
"""Create a namespace-like package module in sys.modules."""
|
||||
module = types.ModuleType(module_name)
|
||||
module.__path__ = []
|
||||
sys.modules[module_name] = module
|
||||
return module
|
||||
|
||||
|
||||
def bootstrap_test_package(
|
||||
package_name: str,
|
||||
*,
|
||||
with_api: bool = True,
|
||||
with_utils: bool = True,
|
||||
with_workers: bool = False,
|
||||
with_upscale: bool = False,
|
||||
with_orchestration: bool = False,
|
||||
) -> None:
|
||||
"""Create a clean package scaffold for module-level import tests."""
|
||||
reset_package_namespace(package_name)
|
||||
ensure_namespace_package(package_name)
|
||||
|
||||
if with_api:
|
||||
ensure_namespace_package(f"{package_name}.api")
|
||||
if with_utils:
|
||||
ensure_namespace_package(f"{package_name}.utils")
|
||||
if with_workers:
|
||||
ensure_namespace_package(f"{package_name}.workers")
|
||||
if with_upscale:
|
||||
ensure_namespace_package(f"{package_name}.upscale")
|
||||
if with_orchestration:
|
||||
ensure_namespace_package(f"{package_name}.api.orchestration")
|
||||
|
||||
|
||||
def install_server_stub(prompt_server_instance: Any | None = None) -> Any:
|
||||
"""Install a minimal `server` module exposing PromptServer.instance.routes."""
|
||||
if prompt_server_instance is None:
|
||||
prompt_server_instance = types.SimpleNamespace()
|
||||
if not hasattr(prompt_server_instance, "routes"):
|
||||
prompt_server_instance.routes = RoutesStub()
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
|
||||
sys.modules["server"] = server_module
|
||||
return prompt_server_instance
|
||||
|
||||
|
||||
def install_aiohttp_stub(
|
||||
json_response_factory: Callable[..., Any],
|
||||
) -> bool:
|
||||
"""Install/augment `aiohttp` with common members used by API tests."""
|
||||
created = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
aiohttp_module = sys.modules["aiohttp"]
|
||||
if not hasattr(aiohttp_module, "ClientTimeout"):
|
||||
aiohttp_module.ClientTimeout = ClientTimeoutStub
|
||||
if not hasattr(aiohttp_module, "WSMsgType"):
|
||||
aiohttp_module.WSMsgType = WSMsgTypeStub
|
||||
if not hasattr(aiohttp_module, "FormData"):
|
||||
aiohttp_module.FormData = FormDataStub
|
||||
|
||||
web_obj = getattr(aiohttp_module, "web", None)
|
||||
if web_obj is None:
|
||||
web_obj = types.SimpleNamespace()
|
||||
aiohttp_module.web = web_obj
|
||||
web_obj.json_response = json_response_factory
|
||||
if not hasattr(web_obj, "WebSocketResponse"):
|
||||
web_obj.WebSocketResponse = WebSocketResponseStub
|
||||
|
||||
return created
|
||||
|
||||
|
||||
def cleanup_optional_module(module_name: str, created: bool) -> None:
|
||||
"""Remove stubbed module only when this test created it."""
|
||||
if created:
|
||||
sys.modules.pop(module_name, None)
|
||||
|
||||
|
||||
def install_request_guards_stub(package_name: str) -> None:
|
||||
"""Install permissive request guard for isolated route tests."""
|
||||
request_guards_module = types.ModuleType(f"{package_name}.api.request_guards")
|
||||
|
||||
async def _authorization_error_or_none(_request):
|
||||
return None
|
||||
|
||||
request_guards_module.authorization_error_or_none = _authorization_error_or_none
|
||||
sys.modules[f"{package_name}.api.request_guards"] = request_guards_module
|
||||
|
||||
|
||||
def install_endpoint_policy_passthrough_stub(package_name: str) -> None:
|
||||
"""Install endpoint policy stub that executes operation directly."""
|
||||
endpoint_policy_module = types.ModuleType(f"{package_name}.api.endpoint_policy")
|
||||
|
||||
async def _run_authorized_endpoint(_request, operation, unexpected_status=500):
|
||||
_ = unexpected_status
|
||||
return await operation()
|
||||
|
||||
endpoint_policy_module.run_authorized_endpoint = _run_authorized_endpoint
|
||||
sys.modules[f"{package_name}.api.endpoint_policy"] = endpoint_policy_module
|
||||
@@ -0,0 +1,423 @@
|
||||
import asyncio
|
||||
import copy
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
AUTH_TOKEN = "integration-secret"
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
self.payload = payload
|
||||
self.status = status
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
def __init__(self, *, headers=None, json_payload=None, post_payload=None, query=None):
|
||||
self.headers = headers or {}
|
||||
self._json_payload = json_payload or {}
|
||||
self._post_payload = post_payload or {}
|
||||
self.query = query or {}
|
||||
|
||||
async def json(self):
|
||||
return self._json_payload
|
||||
|
||||
async def post(self):
|
||||
return self._post_payload
|
||||
|
||||
|
||||
def _load_module(package_name: str, rel_path: str, module_name: str):
|
||||
module_path = Path(__file__).resolve().parents[2] / rel_path
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.{module_name}", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _install_common_utils(package_name: str, config_data: dict):
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: copy.deepcopy(config_data)
|
||||
config_module.save_config = lambda _cfg: True
|
||||
|
||||
@asynccontextmanager
|
||||
async def _config_transaction():
|
||||
yield config_module.load_config()
|
||||
|
||||
config_module.config_transaction = _config_transaction
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def _handle_api_error(_request, error, status=500):
|
||||
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
|
||||
|
||||
network_module.handle_api_error = _handle_api_error
|
||||
network_module.normalize_host = lambda value: value
|
||||
network_module.build_worker_url = lambda worker, path="": f"http://{worker.get('host', '127.0.0.1')}{path}"
|
||||
network_module.get_client_session = lambda: None
|
||||
network_module.probe_worker = lambda *_args, **_kwargs: {"ok": True}
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
|
||||
def _load_real_guard_stack(package_name: str):
|
||||
_load_module(package_name, "utils/auth.py", "utils.auth")
|
||||
_load_module(package_name, "utils/parsing.py", "utils.parsing")
|
||||
_load_module(package_name, "api/schemas.py", "api.schemas")
|
||||
_load_module(package_name, "api/request_guards.py", "api.request_guards")
|
||||
_load_module(package_name, "api/endpoint_policy.py", "api.endpoint_policy")
|
||||
|
||||
|
||||
def _load_request_guards_only_module():
|
||||
package_name = "dist_auth_guard_testpkg"
|
||||
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
request_guards = sys.modules[f"{package_name}.api.request_guards"]
|
||||
schemas = sys.modules[f"{package_name}.api.schemas"]
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
return request_guards, schemas
|
||||
|
||||
|
||||
def _load_config_routes_real_guard_module():
|
||||
package_name = "dist_config_auth_integration_testpkg"
|
||||
config_data = {
|
||||
"workers": [],
|
||||
"master": {"host": "127.0.0.1"},
|
||||
"settings": {"distributed_api_token": AUTH_TOKEN},
|
||||
"tunnel": {},
|
||||
}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
module = _load_module(package_name, "api/config_routes.py", "api.config_routes")
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
return module
|
||||
|
||||
|
||||
def _load_tunnel_routes_real_guard_module():
|
||||
package_name = "dist_tunnel_auth_integration_testpkg"
|
||||
config_data = {
|
||||
"workers": [],
|
||||
"master": {"host": "127.0.0.1"},
|
||||
"settings": {"distributed_api_token": AUTH_TOKEN},
|
||||
"tunnel": {},
|
||||
}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
|
||||
cloudflare_module = types.ModuleType(f"{package_name}.utils.cloudflare")
|
||||
cloudflare_module.cloudflare_tunnel_manager = types.SimpleNamespace(
|
||||
get_status=lambda: {"active": False},
|
||||
start_tunnel=lambda: {"active": True},
|
||||
stop_tunnel=lambda: {"active": False},
|
||||
)
|
||||
sys.modules[f"{package_name}.utils.cloudflare"] = cloudflare_module
|
||||
|
||||
module = _load_module(package_name, "api/tunnel_routes.py", "api.tunnel_routes")
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
return module
|
||||
|
||||
|
||||
def _load_usdu_routes_real_guard_module():
|
||||
package_name = "dist_usdu_auth_integration_testpkg"
|
||||
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_upscale=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
|
||||
usdu_management_module = types.ModuleType(f"{package_name}.utils.usdu_management")
|
||||
usdu_management_module.MAX_PAYLOAD_SIZE = 1024 * 1024
|
||||
sys.modules[f"{package_name}.utils.usdu_management"] = usdu_management_module
|
||||
|
||||
prompt_server_holder = {
|
||||
"value": types.SimpleNamespace(
|
||||
distributed_tile_jobs_lock=asyncio.Lock(),
|
||||
distributed_pending_tile_jobs={},
|
||||
)
|
||||
}
|
||||
|
||||
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
|
||||
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server_holder["value"]
|
||||
job_store_module.init_dynamic_job = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
|
||||
|
||||
class _BaseJobState:
|
||||
pass
|
||||
|
||||
class _ImageJobState(_BaseJobState):
|
||||
def __init__(self):
|
||||
self.mode = "dynamic"
|
||||
self.queue = asyncio.Queue()
|
||||
self.pending_images = asyncio.Queue()
|
||||
self.completed_images = {}
|
||||
self.worker_status = {}
|
||||
self.assigned_to_workers = {}
|
||||
|
||||
class _TileJobState(_BaseJobState):
|
||||
def __init__(self):
|
||||
self.mode = "static"
|
||||
self.queue = asyncio.Queue()
|
||||
self.pending_tasks = asyncio.Queue()
|
||||
self.completed_tasks = {}
|
||||
self.worker_status = {}
|
||||
self.assigned_to_workers = {}
|
||||
self.batched_static = False
|
||||
|
||||
job_models_module.BaseJobState = _BaseJobState
|
||||
job_models_module.ImageJobState = _ImageJobState
|
||||
job_models_module.TileJobState = _TileJobState
|
||||
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
|
||||
|
||||
parsers_module = types.ModuleType(f"{package_name}.upscale.payload_parsers")
|
||||
parsers_module.parse_tiles_from_form = lambda _data: []
|
||||
sys.modules[f"{package_name}.upscale.payload_parsers"] = parsers_module
|
||||
|
||||
module = _load_module(package_name, "api/usdu_routes.py", "api.usdu_routes")
|
||||
module._prompt_server_holder = prompt_server_holder
|
||||
module._ImageJobState = _ImageJobState
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
return module
|
||||
|
||||
|
||||
def _install_torch_stub_if_missing():
|
||||
if "torch" in sys.modules:
|
||||
return False
|
||||
|
||||
torch_module = types.ModuleType("torch")
|
||||
torch_module.cuda = types.SimpleNamespace(
|
||||
is_available=lambda: False,
|
||||
empty_cache=lambda: None,
|
||||
ipc_collect=lambda: None,
|
||||
)
|
||||
sys.modules["torch"] = torch_module
|
||||
return True
|
||||
|
||||
|
||||
def _load_worker_routes_real_guard_module():
|
||||
package_name = "dist_worker_auth_integration_testpkg"
|
||||
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}, "workers": []}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_workers=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
created_torch_stub = _install_torch_stub_if_missing()
|
||||
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.CHUNK_SIZE = 65536
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
workers_module = types.ModuleType(f"{package_name}.workers")
|
||||
workers_module.get_worker_manager = lambda: types.SimpleNamespace(
|
||||
processes={},
|
||||
save_processes=lambda: None,
|
||||
is_process_running=lambda _pid: False,
|
||||
)
|
||||
sys.modules[f"{package_name}.workers"] = workers_module
|
||||
|
||||
detection_module = types.ModuleType(f"{package_name}.workers.detection")
|
||||
detection_module.get_machine_id = lambda: "machine-id"
|
||||
detection_module.is_docker_environment = lambda: False
|
||||
detection_module.is_runpod_environment = lambda: False
|
||||
sys.modules[f"{package_name}.workers.detection"] = detection_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
|
||||
class _PromptValidationError(Exception):
|
||||
def __init__(self, message="validation error"):
|
||||
super().__init__(message)
|
||||
self.validation_error = {}
|
||||
self.node_errors = {}
|
||||
|
||||
async def _queue_prompt_payload(*_args, **_kwargs):
|
||||
return "prompt-id"
|
||||
|
||||
async_helpers_module.PromptValidationError = _PromptValidationError
|
||||
async_helpers_module.queue_prompt_payload = _queue_prompt_payload
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
module = _load_module(package_name, "api/worker_routes.py", "api.worker_routes")
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
return module
|
||||
|
||||
|
||||
def _load_job_routes_real_guard_module():
|
||||
package_name = "dist_job_auth_integration_testpkg"
|
||||
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
install_server_stub()
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
created_torch_stub = _install_torch_stub_if_missing()
|
||||
|
||||
_install_common_utils(package_name, config_data)
|
||||
_load_real_guard_stack(package_name)
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
image_module.pil_to_tensor = lambda _img: None
|
||||
image_module.ensure_contiguous = lambda tensor: tensor
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.JOB_INIT_GRACE_PERIOD = 0.1
|
||||
constants_module.MEMORY_CLEAR_DELAY = 0.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
|
||||
runtime_state = types.SimpleNamespace(
|
||||
distributed_jobs_lock=asyncio.Lock(),
|
||||
distributed_pending_jobs={},
|
||||
)
|
||||
runtime_state_module.ensure_distributed_runtime_state = lambda: runtime_state
|
||||
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
|
||||
|
||||
orchestration_module = types.ModuleType(f"{package_name}.api.queue_orchestration")
|
||||
orchestration_module.orchestrate_distributed_execution = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.api.queue_orchestration"] = orchestration_module
|
||||
|
||||
queue_request_module = types.ModuleType(f"{package_name}.api.queue_request")
|
||||
queue_request_module.parse_queue_request_payload = lambda data: data
|
||||
sys.modules[f"{package_name}.api.queue_request"] = queue_request_module
|
||||
|
||||
module = _load_module(package_name, "api/job_routes.py", "api.job_routes")
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
return module
|
||||
|
||||
|
||||
request_guards_module, schemas_module = _load_request_guards_only_module()
|
||||
config_routes_module = _load_config_routes_real_guard_module()
|
||||
tunnel_routes_module = _load_tunnel_routes_real_guard_module()
|
||||
usdu_routes_module = _load_usdu_routes_real_guard_module()
|
||||
worker_routes_module = _load_worker_routes_real_guard_module()
|
||||
job_routes_module = _load_job_routes_real_guard_module()
|
||||
|
||||
|
||||
class AuthGuardContractTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_is_authorized_request_missing_token_fails(self):
|
||||
request = _FakeRequest(headers={})
|
||||
config = {"settings": {"distributed_api_token": AUTH_TOKEN}}
|
||||
self.assertFalse(schemas_module.is_authorized_request(request, config))
|
||||
|
||||
async def test_is_authorized_request_matching_token_succeeds(self):
|
||||
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN})
|
||||
config = {"settings": {"distributed_api_token": AUTH_TOKEN}}
|
||||
self.assertTrue(schemas_module.is_authorized_request(request, config))
|
||||
|
||||
async def test_authorization_error_or_none_returns_403_for_invalid_token(self):
|
||||
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: "wrong-token"})
|
||||
response = await request_guards_module.authorization_error_or_none(request)
|
||||
self.assertIsNotNone(response)
|
||||
self.assertEqual(response.status, 403)
|
||||
|
||||
async def test_authorization_error_or_none_returns_none_for_valid_token(self):
|
||||
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN})
|
||||
response = await request_guards_module.authorization_error_or_none(request)
|
||||
self.assertIsNone(response)
|
||||
|
||||
|
||||
class RouteAuthIntegrationTests(unittest.IsolatedAsyncioTestCase):
|
||||
def _auth_headers(self):
|
||||
return {schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN}
|
||||
|
||||
async def test_config_endpoint_enforces_real_guard(self):
|
||||
unauthorized = await config_routes_module.get_config_endpoint(_FakeRequest(headers={}))
|
||||
self.assertEqual(unauthorized.status, 403)
|
||||
|
||||
authorized = await config_routes_module.get_config_endpoint(_FakeRequest(headers=self._auth_headers()))
|
||||
self.assertEqual(authorized.status, 200)
|
||||
self.assertIn("settings", authorized.payload)
|
||||
|
||||
async def test_tunnel_endpoint_enforces_real_guard(self):
|
||||
unauthorized = await tunnel_routes_module.tunnel_status_endpoint(_FakeRequest(headers={}))
|
||||
self.assertEqual(unauthorized.status, 403)
|
||||
|
||||
authorized = await tunnel_routes_module.tunnel_status_endpoint(_FakeRequest(headers=self._auth_headers()))
|
||||
self.assertEqual(authorized.status, 200)
|
||||
self.assertEqual(authorized.payload.get("status"), "success")
|
||||
|
||||
async def test_usdu_endpoint_enforces_real_guard(self):
|
||||
unauthorized = await usdu_routes_module.job_status_endpoint(_FakeRequest(headers={}))
|
||||
self.assertEqual(unauthorized.status, 403)
|
||||
|
||||
authorized = await usdu_routes_module.job_status_endpoint(
|
||||
_FakeRequest(headers=self._auth_headers(), query={"multi_job_id": "job-1"})
|
||||
)
|
||||
self.assertEqual(authorized.status, 200)
|
||||
self.assertIn("ready", authorized.payload)
|
||||
|
||||
async def test_worker_endpoint_enforces_real_guard(self):
|
||||
worker_routes_module._collect_network_info_sync = lambda: {
|
||||
"interfaces": [],
|
||||
"recommended_host": "127.0.0.1",
|
||||
}
|
||||
unauthorized = await worker_routes_module.get_network_info_endpoint(_FakeRequest(headers={}))
|
||||
self.assertEqual(unauthorized.status, 403)
|
||||
|
||||
authorized = await worker_routes_module.get_network_info_endpoint(
|
||||
_FakeRequest(headers=self._auth_headers())
|
||||
)
|
||||
self.assertEqual(authorized.status, 200)
|
||||
self.assertEqual(authorized.payload.get("status"), "success")
|
||||
|
||||
async def test_job_endpoint_enforces_real_guard(self):
|
||||
unauthorized = await job_routes_module.prepare_job_endpoint(
|
||||
_FakeRequest(headers={}, json_payload={"multi_job_id": "job-1"})
|
||||
)
|
||||
self.assertEqual(unauthorized.status, 403)
|
||||
|
||||
authorized = await job_routes_module.prepare_job_endpoint(
|
||||
_FakeRequest(headers=self._auth_headers(), json_payload={"multi_job_id": "job-1"})
|
||||
)
|
||||
self.assertEqual(authorized.status, 200)
|
||||
self.assertEqual(authorized.payload.get("status"), "success")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -3,9 +3,19 @@ import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
install_endpoint_policy_passthrough_stub,
|
||||
install_request_guards_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
@@ -25,45 +35,18 @@ def _load_config_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "config_routes.py"
|
||||
package_name = "dist_api_config_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
|
||||
schemas_module.is_authorized_request = lambda _request, _config: True
|
||||
sys.modules[f"{package_name}.api.schemas"] = schemas_module
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
aiohttp_module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
class _Routes:
|
||||
def get(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
def post(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
|
||||
sys.modules["server"] = server_module
|
||||
install_request_guards_stub(package_name)
|
||||
install_endpoint_policy_passthrough_stub(package_name)
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
install_server_stub()
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
@@ -73,7 +56,16 @@ def _load_config_routes_module():
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def _handle_api_error(_request, error, status=500):
|
||||
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
|
||||
if isinstance(error, list):
|
||||
message = "; ".join(str(item) for item in error)
|
||||
error_payload = [str(item) for item in error]
|
||||
else:
|
||||
message = str(error)
|
||||
error_payload = str(error)
|
||||
return _FakeResponse(
|
||||
{"status": "error", "error": error_payload, "message": message},
|
||||
status=status,
|
||||
)
|
||||
|
||||
network_module.handle_api_error = _handle_api_error
|
||||
network_module.normalize_host = lambda value: value
|
||||
@@ -89,6 +81,12 @@ def _load_config_routes_module():
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: copy.deepcopy(default_config)
|
||||
config_module.save_config = lambda _cfg: True
|
||||
|
||||
@asynccontextmanager
|
||||
async def _config_transaction():
|
||||
yield config_module.load_config()
|
||||
|
||||
config_module.config_transaction = _config_transaction
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.api.config_routes", module_path)
|
||||
@@ -96,8 +94,7 @@ def _load_config_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
|
||||
return module
|
||||
|
||||
@@ -118,9 +115,7 @@ class ConfigRoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_update_config_valid_field_persists(self):
|
||||
cfg = {"workers": [], "master": {}, "settings": {"debug": False}, "tunnel": {}}
|
||||
with patch.object(config_routes, "load_config", return_value=cfg), patch.object(
|
||||
config_routes, "save_config", return_value=True
|
||||
):
|
||||
with patch.object(config_routes, "load_config", return_value=cfg):
|
||||
response = await config_routes.update_config_endpoint(_FakeRequest({"debug": True}))
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
|
||||
@@ -11,6 +11,14 @@ from unittest.mock import AsyncMock, patch
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
install_request_guards_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
@@ -19,8 +27,9 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
def __init__(self, payload):
|
||||
def __init__(self, payload, headers=None):
|
||||
self._payload = payload
|
||||
self.headers = headers or {}
|
||||
|
||||
async def json(self):
|
||||
return self._payload
|
||||
@@ -30,53 +39,19 @@ def _load_job_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "job_routes.py"
|
||||
package_name = "dist_api_queue_testpkg"
|
||||
|
||||
# Reset package namespace to avoid stale module state across test runs.
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
|
||||
# aiohttp.web stub
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
aiohttp_module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
|
||||
# server module stub with route decorators
|
||||
class _Routes:
|
||||
def get(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
def post(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
prompt_server_instance = types.SimpleNamespace(
|
||||
routes=_Routes(),
|
||||
distributed_jobs_lock=None,
|
||||
distributed_pending_jobs={},
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
|
||||
sys.modules["server"] = server_module
|
||||
install_server_stub(prompt_server_instance)
|
||||
|
||||
# torch stub (only needed to satisfy import)
|
||||
created_torch_stub = False
|
||||
@@ -157,12 +132,53 @@ def _load_job_routes_module():
|
||||
constants_module.JOB_INIT_GRACE_PERIOD = 10.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"settings": {}}
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
|
||||
|
||||
def _ensure_distributed_runtime_state(server_instance=None):
|
||||
ps = server_instance or prompt_server_instance
|
||||
if not hasattr(ps, "distributed_pending_jobs"):
|
||||
ps.distributed_pending_jobs = {}
|
||||
if not hasattr(ps, "distributed_jobs_lock") or ps.distributed_jobs_lock is None:
|
||||
ps.distributed_jobs_lock = asyncio.Lock()
|
||||
if not hasattr(ps, "distributed_job_allowed_workers"):
|
||||
ps.distributed_job_allowed_workers = {}
|
||||
if not hasattr(ps, "distributed_pending_tile_jobs"):
|
||||
ps.distributed_pending_tile_jobs = {}
|
||||
if not hasattr(ps, "distributed_tile_jobs_lock"):
|
||||
ps.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
return types.SimpleNamespace(
|
||||
distributed_pending_jobs=ps.distributed_pending_jobs,
|
||||
distributed_jobs_lock=ps.distributed_jobs_lock,
|
||||
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
|
||||
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
|
||||
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
|
||||
)
|
||||
|
||||
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
|
||||
runtime_state_module.get_prompt_server_instance = lambda: prompt_server_instance
|
||||
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
|
||||
|
||||
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
|
||||
schemas_module.is_authorized_request = lambda _request, _config: True
|
||||
schemas_module.require_bool_literal = (
|
||||
lambda value, field_name="value": value
|
||||
if isinstance(value, bool)
|
||||
else (_ for _ in ()).throw(ValueError(f"Field '{field_name}' must be a boolean literal."))
|
||||
)
|
||||
sys.modules[f"{package_name}.api.schemas"] = schemas_module
|
||||
|
||||
install_request_guards_stub(package_name)
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.queue_prompt_payload = AsyncMock(return_value="prompt_local")
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
queue_orchestration_module = types.ModuleType(f"{package_name}.api.queue_orchestration")
|
||||
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 7, 1, {}))
|
||||
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 1))
|
||||
sys.modules[f"{package_name}.api.queue_orchestration"] = queue_orchestration_module
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -172,7 +188,6 @@ def _load_job_routes_module():
|
||||
client_id: str
|
||||
delegate_master: object
|
||||
enabled_worker_ids: list
|
||||
auto_prepare: bool
|
||||
trace_execution_id: object
|
||||
|
||||
def _parse_queue_request_payload(data):
|
||||
@@ -193,7 +208,6 @@ def _load_job_routes_module():
|
||||
client_id=client_id,
|
||||
delegate_master=data.get("delegate_master"),
|
||||
enabled_worker_ids=enabled,
|
||||
auto_prepare=bool(data.get("auto_prepare", True)),
|
||||
trace_execution_id=data.get("trace_execution_id"),
|
||||
)
|
||||
|
||||
@@ -206,13 +220,11 @@ def _load_job_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
cleanup_optional_module("torch", created_torch_stub)
|
||||
if created_pil_stub:
|
||||
sys.modules.pop("PIL.Image", None)
|
||||
sys.modules.pop("PIL", None)
|
||||
cleanup_optional_module("PIL.Image", True)
|
||||
cleanup_optional_module("PIL", True)
|
||||
|
||||
return module
|
||||
|
||||
@@ -221,7 +233,7 @@ job_routes = _load_job_routes_module()
|
||||
|
||||
|
||||
class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_metadata(self):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_id(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"prompt": {"1": {"class_type": "Node"}},
|
||||
@@ -233,15 +245,12 @@ class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
with patch.object(
|
||||
job_routes,
|
||||
"orchestrate_distributed_execution",
|
||||
new=AsyncMock(return_value=("prompt_123", 42, 2, {})),
|
||||
new=AsyncMock(return_value=("prompt_123", 2)),
|
||||
):
|
||||
response = await job_routes.distributed_queue_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("prompt_id"), "prompt_123")
|
||||
self.assertEqual(response.payload.get("number"), 42)
|
||||
self.assertEqual(response.payload.get("node_errors"), {})
|
||||
self.assertTrue(response.payload.get("auto_prepare_supported"))
|
||||
|
||||
async def test_distributed_queue_missing_prompt_returns_400(self):
|
||||
request = _FakeRequest(
|
||||
@@ -278,8 +287,10 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_job_complete_accepts_audio_payload(self):
|
||||
queue = asyncio.Queue()
|
||||
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
job_routes.prompt_server.distributed_pending_jobs = {"job-1": queue}
|
||||
prompt_server = job_routes.server.PromptServer.instance
|
||||
prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
prompt_server.distributed_pending_jobs = {"job-1": queue}
|
||||
prompt_server.distributed_job_allowed_workers = {"job-1": {"worker-1"}}
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
@@ -302,76 +313,27 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(queued["audio"]["sample_rate"], 44100)
|
||||
self.assertEqual(tuple(queued["audio"]["waveform"].shape), (1, 2, 4))
|
||||
|
||||
async def test_job_complete_accepts_audio_without_image(self):
|
||||
async def test_job_complete_rejects_worker_not_in_job_allowlist(self):
|
||||
queue = asyncio.Queue()
|
||||
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
job_routes.prompt_server.distributed_pending_jobs = {"audio-only-job": queue}
|
||||
prompt_server = job_routes.server.PromptServer.instance
|
||||
prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
prompt_server.distributed_pending_jobs = {"job-allow": queue}
|
||||
prompt_server.distributed_job_allowed_workers = {"job-allow": {"worker-expected"}}
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "audio-only-job",
|
||||
"worker_id": "worker-1",
|
||||
"job_id": "job-allow",
|
||||
"worker_id": "worker-unexpected",
|
||||
"batch_idx": 0,
|
||||
"audio": self._encoded_audio_payload(),
|
||||
"image": "data:image/png;base64,AAAA",
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
with patch.object(job_routes, "_decode_canonical_png_tensor") as decode_image:
|
||||
with patch.object(job_routes, "_decode_canonical_png_tensor", return_value="tensor-data"):
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
decode_image.assert_not_called()
|
||||
queued = await queue.get()
|
||||
self.assertIsNone(queued["tensor"])
|
||||
self.assertEqual(queued["audio"]["sample_rate"], 44100)
|
||||
|
||||
async def test_job_complete_rejects_payload_without_image_or_audio(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("image or audio", response.payload.get("message", "").lower())
|
||||
|
||||
async def test_job_complete_rejects_invalid_image_even_when_audio_is_present(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"image": 123,
|
||||
"audio": self._encoded_audio_payload(),
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("image", response.payload.get("message", "").lower())
|
||||
|
||||
async def test_job_complete_returns_400_for_malformed_audio(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"audio": {"data": "AAAA", "shape": [1, 2], "dtype": "float32"},
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("audio.shape", response.payload.get("message", "").lower())
|
||||
self.assertEqual(response.status, 403)
|
||||
self.assertIn("unauthorized", response.payload.get("message", "").lower())
|
||||
|
||||
def test_decode_audio_payload_rejects_bad_shape(self):
|
||||
bad = {
|
||||
@@ -4,30 +4,29 @@ import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
)
|
||||
|
||||
|
||||
class _AiohttpResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
self.payload = payload
|
||||
self.status = status
|
||||
|
||||
|
||||
def _load_media_sync_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "orchestration" / "media_sync.py"
|
||||
package_name = "dist_ms_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
|
||||
orch_pkg = types.ModuleType(f"{package_name}.api.orchestration")
|
||||
orch_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api.orchestration"] = orch_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
bootstrap_test_package(
|
||||
package_name,
|
||||
with_api=True,
|
||||
with_utils=True,
|
||||
with_orchestration=True,
|
||||
)
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
@@ -48,22 +47,17 @@ def _load_media_sync_module():
|
||||
trace_module.trace_info = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.trace_logger"] = trace_module
|
||||
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
auth_module = types.ModuleType(f"{package_name}.utils.auth")
|
||||
auth_module.distributed_auth_headers = lambda _config=None: {}
|
||||
sys.modules[f"{package_name}.utils.auth"] = auth_module
|
||||
|
||||
class _ClientTimeout:
|
||||
def __init__(self, total=None):
|
||||
pass
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"settings": {}}
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
class _FormData:
|
||||
def add_field(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
aiohttp_module.ClientTimeout = _ClientTimeout
|
||||
aiohttp_module.FormData = _FormData
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _AiohttpResponse(payload, status=status)
|
||||
)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.api.orchestration.media_sync",
|
||||
@@ -73,8 +67,7 @@ def _load_media_sync_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
|
||||
return module
|
||||
|
||||
@@ -164,11 +157,26 @@ class ConvertPathsForPlatformTests(unittest.TestCase):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class FindMediaReferencesTests(unittest.TestCase):
|
||||
def test_finds_image_input(self):
|
||||
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.png"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("photo.png", refs)
|
||||
|
||||
def test_finds_video_input(self):
|
||||
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"video": "clip.mp4"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("clip.mp4", refs)
|
||||
|
||||
def test_finds_file_input_for_load_video(self):
|
||||
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"file": "1 - Copy.mp4"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("1 - Copy.mp4", refs)
|
||||
|
||||
def test_finds_audio_input(self):
|
||||
prompt = {"1": {"class_type": "LoadAudio", "inputs": {"audio": "track.wav"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("track.wav", refs)
|
||||
|
||||
def test_strips_annotation_suffix(self):
|
||||
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.jpg [abc123]"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
@@ -245,3 +253,7 @@ class RewritePromptMediaInputsTests(unittest.TestCase):
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
class _AiohttpResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
self.payload = payload
|
||||
self.status = status
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
bootstrap_test_package(
|
||||
package_name,
|
||||
with_api=True,
|
||||
with_utils=True,
|
||||
with_orchestration=True,
|
||||
)
|
||||
|
||||
prompt_server_instance = types.SimpleNamespace(address="127.0.0.1", port=8188)
|
||||
install_server_stub(prompt_server_instance)
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
|
||||
async def _queue_prompt_payload(prompt_obj, workflow_meta, client_id):
|
||||
_ = (prompt_obj, workflow_meta, client_id)
|
||||
return "prompt-id"
|
||||
|
||||
async_helpers_module.queue_prompt_payload = _queue_prompt_payload
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"settings": {}, "workers": []}
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.ORCHESTRATION_MEDIA_SYNC_CONCURRENCY = 2
|
||||
constants_module.ORCHESTRATION_MEDIA_SYNC_TIMEOUT = 120.0
|
||||
constants_module.ORCHESTRATION_WORKER_PROBE_CONCURRENCY = 8
|
||||
constants_module.ORCHESTRATION_WORKER_PREP_CONCURRENCY = 4
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.build_master_url = lambda **_kwargs: "http://127.0.0.1:8188"
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
trace_logger_module = types.ModuleType(f"{package_name}.utils.trace_logger")
|
||||
trace_logger_module.trace_debug = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.trace_logger"] = trace_logger_module
|
||||
|
||||
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
|
||||
|
||||
def _ensure_distributed_runtime_state(server_instance=None):
|
||||
ps = server_instance or prompt_server_instance
|
||||
if not hasattr(ps, "distributed_pending_jobs"):
|
||||
ps.distributed_pending_jobs = {}
|
||||
if not hasattr(ps, "distributed_jobs_lock"):
|
||||
ps.distributed_jobs_lock = asyncio.Lock()
|
||||
if not hasattr(ps, "distributed_job_allowed_workers"):
|
||||
ps.distributed_job_allowed_workers = {}
|
||||
if not hasattr(ps, "distributed_pending_tile_jobs"):
|
||||
ps.distributed_pending_tile_jobs = {}
|
||||
if not hasattr(ps, "distributed_tile_jobs_lock"):
|
||||
ps.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
return types.SimpleNamespace(
|
||||
distributed_pending_jobs=ps.distributed_pending_jobs,
|
||||
distributed_jobs_lock=ps.distributed_jobs_lock,
|
||||
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
|
||||
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
|
||||
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
|
||||
)
|
||||
|
||||
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
|
||||
runtime_state_module.get_prompt_server_instance = lambda: prompt_server_instance
|
||||
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
|
||||
|
||||
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
|
||||
schemas_module.coerce_positive_int = (
|
||||
lambda value, default: int(value) if str(value).isdigit() and int(value) > 0 else default
|
||||
)
|
||||
schemas_module.coerce_positive_float = (
|
||||
lambda value, default: float(value) if isinstance(value, (int, float, str)) and float(value) > 0 else default
|
||||
)
|
||||
sys.modules[f"{package_name}.api.schemas"] = schemas_module
|
||||
|
||||
dispatch_module = types.ModuleType(f"{package_name}.api.orchestration.dispatch")
|
||||
dispatch_module.dispatch_worker_prompt = lambda *args, **kwargs: None
|
||||
dispatch_module.rank_workers_by_load = lambda workers, **kwargs: workers
|
||||
dispatch_module.select_active_workers = lambda workers, *_args, **_kwargs: (workers, False)
|
||||
dispatch_module.select_least_busy_worker = lambda workers, **_kwargs: workers[0] if workers else None
|
||||
sys.modules[f"{package_name}.api.orchestration.dispatch"] = dispatch_module
|
||||
|
||||
media_sync_module = types.ModuleType(f"{package_name}.api.orchestration.media_sync")
|
||||
media_sync_module.convert_paths_for_platform = lambda prompt, _sep: prompt
|
||||
media_sync_module.fetch_worker_path_separator = lambda *_args, **_kwargs: "/"
|
||||
media_sync_module.sync_worker_media = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.api.orchestration.media_sync"] = media_sync_module
|
||||
|
||||
prompt_transform_module = types.ModuleType(f"{package_name}.api.orchestration.prompt_transform")
|
||||
prompt_transform_module.PromptIndex = object
|
||||
prompt_transform_module.apply_participant_overrides = lambda prompt, *_args, **_kwargs: prompt
|
||||
prompt_transform_module.find_nodes_by_class = lambda _prompt, _class_name: []
|
||||
prompt_transform_module.generate_job_id_map = lambda _index, _prefix: {}
|
||||
prompt_transform_module.prepare_delegate_master_prompt = lambda prompt, _ids: prompt
|
||||
prompt_transform_module.prune_prompt_for_worker = lambda prompt: prompt
|
||||
sys.modules[f"{package_name}.api.orchestration.prompt_transform"] = prompt_transform_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[2] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_queue_orchestration_module():
|
||||
package_name = "dist_queue_orch_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
return _load_module(package_name, "api/queue_orchestration.py", "api.queue_orchestration")
|
||||
|
||||
|
||||
queue_orchestration = _load_queue_orchestration_module()
|
||||
|
||||
|
||||
class QueueOrchestrationHelpersTests(unittest.TestCase):
|
||||
def test_resolve_enabled_workers_filters_disabled_and_requested(self):
|
||||
config = {
|
||||
"workers": [
|
||||
{"id": "1", "enabled": True, "host": "a", "port": 8188},
|
||||
{"id": "2", "enabled": False, "host": "b", "port": 8189},
|
||||
{"id": "3", "enabled": True, "host": "c", "port": "bad-port"},
|
||||
]
|
||||
}
|
||||
|
||||
enabled_only = queue_orchestration._resolve_enabled_workers(config)
|
||||
requested = queue_orchestration._resolve_enabled_workers(config, requested_ids={"2", "3"})
|
||||
|
||||
self.assertEqual([worker["id"] for worker in enabled_only], ["1", "3"])
|
||||
self.assertEqual([worker["id"] for worker in requested], ["2", "3"])
|
||||
self.assertEqual(requested[1]["port"], 8188)
|
||||
|
||||
def test_resolve_orchestration_limits_uses_defaults_for_invalid_values(self):
|
||||
config = {
|
||||
"settings": {
|
||||
"worker_probe_concurrency": "x",
|
||||
"worker_prep_concurrency": "3",
|
||||
"media_sync_concurrency": "-1",
|
||||
"media_sync_timeout_seconds": "42.5",
|
||||
}
|
||||
}
|
||||
|
||||
limits = queue_orchestration._resolve_orchestration_limits(config)
|
||||
|
||||
self.assertEqual(limits[0], 8) # default
|
||||
self.assertEqual(limits[1], 3) # parsed
|
||||
self.assertEqual(limits[2], 2) # default
|
||||
self.assertEqual(limits[3], 42.5)
|
||||
|
||||
def test_is_load_balance_enabled_accepts_common_forms(self):
|
||||
self.assertTrue(queue_orchestration._is_load_balance_enabled(True))
|
||||
self.assertTrue(queue_orchestration._is_load_balance_enabled(1))
|
||||
self.assertTrue(queue_orchestration._is_load_balance_enabled(" yes "))
|
||||
self.assertFalse(queue_orchestration._is_load_balance_enabled(0))
|
||||
self.assertFalse(queue_orchestration._is_load_balance_enabled("off"))
|
||||
|
||||
def test_prompt_requests_load_balance_detects_collector_flag(self):
|
||||
prompt_index = types.SimpleNamespace(
|
||||
nodes_for_class=lambda class_name: ["10"] if class_name == "DistributedCollector" else [],
|
||||
inputs_by_node={"10": {"load_balance": "true"}},
|
||||
)
|
||||
self.assertTrue(queue_orchestration._prompt_requests_load_balance(prompt_index))
|
||||
|
||||
def test_ensure_distributed_state_initializes_attrs(self):
|
||||
prompt_server = types.SimpleNamespace()
|
||||
queue_orchestration.ensure_distributed_state(prompt_server)
|
||||
self.assertTrue(hasattr(prompt_server, "distributed_pending_jobs"))
|
||||
self.assertTrue(hasattr(prompt_server, "distributed_jobs_lock"))
|
||||
self.assertIsInstance(prompt_server.distributed_jobs_lock, asyncio.Lock)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,95 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
install_aiohttp_stub,
|
||||
install_endpoint_policy_passthrough_stub,
|
||||
install_request_guards_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
self.payload = payload
|
||||
self.status = status
|
||||
|
||||
|
||||
def _load_tunnel_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "tunnel_routes.py"
|
||||
package_name = "dist_tunnel_routes_testpkg"
|
||||
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
install_server_stub()
|
||||
install_aiohttp_stub(lambda payload, status=200: _FakeResponse(payload, status=status))
|
||||
|
||||
cloudflare_module = types.ModuleType(f"{package_name}.utils.cloudflare")
|
||||
|
||||
class _TunnelManager:
|
||||
def get_status(self):
|
||||
return {"status": "running"}
|
||||
|
||||
async def start_tunnel(self):
|
||||
return {"status": "running", "public_url": "https://example.trycloudflare.com"}
|
||||
|
||||
async def stop_tunnel(self):
|
||||
return {"status": "stopped"}
|
||||
|
||||
cloudflare_module.cloudflare_tunnel_manager = _TunnelManager()
|
||||
sys.modules[f"{package_name}.utils.cloudflare"] = cloudflare_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"master": {"host": "https://master.example.com"}}
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def _handle_api_error(_request, error, status=500):
|
||||
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
|
||||
|
||||
network_module.handle_api_error = _handle_api_error
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
install_request_guards_stub(package_name)
|
||||
install_endpoint_policy_passthrough_stub(package_name)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.api.tunnel_routes", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
tunnel_routes = _load_tunnel_routes_module()
|
||||
|
||||
|
||||
class TunnelRoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_status_endpoint_returns_tunnel_and_master_host(self):
|
||||
response = await tunnel_routes.tunnel_status_endpoint(request=None)
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload["status"], "success")
|
||||
self.assertEqual(response.payload["tunnel"]["status"], "running")
|
||||
self.assertEqual(response.payload["master_host"], "https://master.example.com")
|
||||
|
||||
async def test_start_and_stop_endpoints_return_success_payload(self):
|
||||
start_response = await tunnel_routes.tunnel_start_endpoint(request=None)
|
||||
stop_response = await tunnel_routes.tunnel_stop_endpoint(request=None)
|
||||
|
||||
self.assertEqual(start_response.status, 200)
|
||||
self.assertEqual(start_response.payload["tunnel"]["status"], "running")
|
||||
self.assertEqual(stop_response.status, 200)
|
||||
self.assertEqual(stop_response.payload["tunnel"]["status"], "stopped")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -9,6 +9,14 @@ from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
install_request_guards_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
@@ -30,43 +38,30 @@ class _FakeRequest:
|
||||
return self._post_payload
|
||||
|
||||
|
||||
class _Routes:
|
||||
def post(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
|
||||
return _decorator
|
||||
|
||||
def get(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
|
||||
return _decorator
|
||||
|
||||
|
||||
def _load_usdu_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "usdu_routes.py"
|
||||
package_name = "dist_api_usdu_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_upscale=True)
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
def _require_bool_literal(value, field_name: str = "value"):
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized == "true":
|
||||
return True
|
||||
if normalized == "false":
|
||||
return False
|
||||
raise ValueError(f"Field '{field_name}' must be a boolean literal.")
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
schemas_module.require_bool_literal = _require_bool_literal
|
||||
schemas_module.is_authorized_request = lambda _request, _config: True
|
||||
sys.modules[f"{package_name}.api.schemas"] = schemas_module
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
install_request_guards_stub(package_name)
|
||||
|
||||
prompt_server_holder = {
|
||||
"value": types.SimpleNamespace(
|
||||
@@ -75,23 +70,23 @@ def _load_usdu_routes_module():
|
||||
)
|
||||
}
|
||||
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
aiohttp_module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
|
||||
sys.modules["server"] = server_module
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
install_server_stub()
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"settings": {}}
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
usdu_management_module = types.ModuleType(f"{package_name}.utils.usdu_management")
|
||||
usdu_management_module.MAX_PAYLOAD_SIZE = 1024
|
||||
sys.modules[f"{package_name}.utils.usdu_management"] = usdu_management_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def _handle_api_error(_request, error, status=500):
|
||||
@@ -103,6 +98,21 @@ def _load_usdu_routes_module():
|
||||
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
|
||||
job_store_module.MAX_PAYLOAD_SIZE = 1024
|
||||
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server_holder["value"]
|
||||
|
||||
async def _init_dynamic_job(multi_job_id, batch_size, enabled_workers, all_indices=None):
|
||||
prompt_server = prompt_server_holder["value"]
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
|
||||
return
|
||||
job_data = ImageJobState(multi_job_id=multi_job_id)
|
||||
job_data.worker_status = {str(worker_id): 0 for worker_id in (enabled_workers or [])}
|
||||
job_data.assigned_to_workers = {str(worker_id): [] for worker_id in (enabled_workers or [])}
|
||||
indices = all_indices if all_indices is not None else list(range(int(batch_size or 0)))
|
||||
for idx in indices:
|
||||
await job_data.pending_images.put(int(idx))
|
||||
prompt_server.distributed_pending_tile_jobs[multi_job_id] = job_data
|
||||
|
||||
job_store_module.init_dynamic_job = _init_dynamic_job
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
|
||||
@@ -150,6 +160,7 @@ def _load_usdu_routes_module():
|
||||
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
|
||||
|
||||
parsers_module = types.ModuleType(f"{package_name}.upscale.payload_parsers")
|
||||
parsers_module.parse_tiles_from_form = lambda _data: []
|
||||
parsers_module._parse_tiles_from_form = lambda _data: []
|
||||
sys.modules[f"{package_name}.upscale.payload_parsers"] = parsers_module
|
||||
|
||||
@@ -158,8 +169,7 @@ def _load_usdu_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
|
||||
module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
@@ -222,7 +232,8 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
response = await usdu_routes.request_image_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("image_idx"), 7)
|
||||
self.assertEqual(response.payload.get("kind"), "image")
|
||||
self.assertEqual(response.payload.get("task_idx"), 7)
|
||||
self.assertEqual(response.payload.get("estimated_remaining"), 0)
|
||||
self.assertEqual(job_data.assigned_to_workers["worker-a"], [7])
|
||||
self.assertIn("worker-a", job_data.worker_status)
|
||||
@@ -238,7 +249,8 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
response = await usdu_routes.request_image_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("tile_idx"), 4)
|
||||
self.assertEqual(response.payload.get("kind"), "tile")
|
||||
self.assertEqual(response.payload.get("task_idx"), 4)
|
||||
self.assertTrue(response.payload.get("batched_static"))
|
||||
self.assertEqual(job_data.assigned_to_workers["worker-a"], [4])
|
||||
|
||||
|
||||
+34
-111
@@ -8,6 +8,14 @@ from collections import deque
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from tests.api.harness import (
|
||||
bootstrap_test_package,
|
||||
cleanup_optional_module,
|
||||
install_aiohttp_stub,
|
||||
install_request_guards_stub,
|
||||
install_server_stub,
|
||||
)
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
@@ -16,10 +24,11 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
def __init__(self, payload=None, match_info=None, query=None):
|
||||
def __init__(self, payload=None, match_info=None, query=None, headers=None):
|
||||
self._payload = payload
|
||||
self.match_info = match_info or {}
|
||||
self.query = query or {}
|
||||
self.headers = headers or {}
|
||||
|
||||
async def json(self):
|
||||
return self._payload
|
||||
@@ -49,8 +58,8 @@ class _FakeHTTPClientSession:
|
||||
self._status = status
|
||||
self.calls = []
|
||||
|
||||
def get(self, url, params=None, timeout=None):
|
||||
self.calls.append({"url": url, "params": params, "timeout": timeout})
|
||||
def get(self, url, params=None, headers=None, timeout=None):
|
||||
self.calls.append({"url": url, "params": params, "headers": headers, "timeout": timeout})
|
||||
return _FakeHTTPClientResponse(self._payload, status=self._status)
|
||||
|
||||
|
||||
@@ -62,12 +71,12 @@ class _DummyWorkerManager:
|
||||
worker_id = str(worker["id"])
|
||||
self.processes[worker_id] = {
|
||||
"pid": 12345,
|
||||
"log_file": f"/tmp/distributed_worker_{worker_id}.log",
|
||||
"log_file": f"/tmp/distributed_worker_{worker_id}.log", # nosec B108 - deterministic fake test path
|
||||
"process": None,
|
||||
}
|
||||
return 12345
|
||||
|
||||
def _is_process_running(self, _pid):
|
||||
def is_process_running(self, _pid):
|
||||
return False
|
||||
|
||||
def save_processes(self):
|
||||
@@ -89,24 +98,10 @@ def _load_worker_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "worker_routes.py"
|
||||
package_name = "dist_api_worker_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_workers=True)
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = [str(module_path.parents[1] / "workers")]
|
||||
workers_pkg.__path__ = []
|
||||
workers_pkg.get_worker_manager = lambda: _DummyWorkerManager()
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
@@ -119,59 +114,10 @@ def _load_worker_routes_module():
|
||||
detection_module.get_comms_channel = lambda *_args, **_kwargs: "lan"
|
||||
sys.modules[f"{package_name}.workers.detection"] = detection_module
|
||||
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
|
||||
class _ClientTimeout:
|
||||
def __init__(self, total=None):
|
||||
self.total = total
|
||||
|
||||
class _WSMsgType:
|
||||
TEXT = "TEXT"
|
||||
ERROR = "ERROR"
|
||||
CLOSED = "CLOSED"
|
||||
|
||||
class _WebSocketResponse:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.args = args
|
||||
self.kwargs = kwargs
|
||||
|
||||
async def prepare(self, _request):
|
||||
return None
|
||||
|
||||
async def send_json(self, _payload):
|
||||
return None
|
||||
|
||||
def __aiter__(self):
|
||||
async def _empty():
|
||||
if False:
|
||||
yield None
|
||||
return _empty()
|
||||
|
||||
aiohttp_module.ClientTimeout = _ClientTimeout
|
||||
aiohttp_module.WSMsgType = _WSMsgType
|
||||
aiohttp_module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status),
|
||||
WebSocketResponse=_WebSocketResponse,
|
||||
)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
class _Routes:
|
||||
def get(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
def post(self, _path):
|
||||
def _decorator(fn):
|
||||
return fn
|
||||
return _decorator
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
|
||||
sys.modules["server"] = server_module
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
install_server_stub()
|
||||
|
||||
created_torch_stub = False
|
||||
if "torch" not in sys.modules:
|
||||
@@ -213,7 +159,6 @@ def _load_worker_routes_module():
|
||||
raise RuntimeError("not used in these tests")
|
||||
|
||||
network_module.get_client_session = _get_client_session
|
||||
network_module.get_server_port = lambda: 8189
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
@@ -249,19 +194,29 @@ def _load_worker_routes_module():
|
||||
def _validate_worker_id(worker_id, config):
|
||||
return any(str(worker.get("id")) == str(worker_id) for worker in config.get("workers", []))
|
||||
|
||||
def _require_worker_id(worker_id, config, field_name="worker_id"):
|
||||
_ = field_name
|
||||
worker_id_str = str(worker_id).strip()
|
||||
if any(str(worker.get("id")) == worker_id_str for worker in config.get("workers", [])):
|
||||
return worker_id_str
|
||||
raise ValueError(f"Worker {worker_id_str} not found")
|
||||
|
||||
schemas_module.require_fields = _require_fields
|
||||
schemas_module.validate_worker_id = _validate_worker_id
|
||||
schemas_module.require_worker_id = _require_worker_id
|
||||
schemas_module.is_authorized_request = lambda _request, _config: True
|
||||
schemas_module.distributed_auth_headers = lambda _config: {}
|
||||
sys.modules[f"{package_name}.api.schemas"] = schemas_module
|
||||
|
||||
install_request_guards_stub(package_name)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.api.worker_routes", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
cleanup_optional_module("torch", created_torch_stub)
|
||||
|
||||
return module
|
||||
|
||||
@@ -270,38 +225,6 @@ worker_routes = _load_worker_routes_module()
|
||||
|
||||
|
||||
class WorkerRoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_network_info_reports_actual_master_and_available_ports(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 1, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=["127.0.0.1"]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": []}), \
|
||||
patch.object(worker_routes, "get_server_port", return_value=8189), \
|
||||
patch.object(worker_routes, "allocate_worker_ports", return_value=[8191, 8192, 8193]) as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["master_port"], 8189)
|
||||
self.assertEqual(response.payload["local_worker_ports"], [8191, 8192, 8193])
|
||||
self.assertEqual(response.payload["cuda_device_count"], 4)
|
||||
allocate.assert_called_once_with(8189, [], 3)
|
||||
|
||||
async def test_network_info_does_not_reallocate_existing_configuration(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 4, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=[]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": [{"id": "existing", "port": 8189}]}), \
|
||||
patch.object(worker_routes, "allocate_worker_ports") as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["local_worker_ports"], [])
|
||||
allocate.assert_not_called()
|
||||
|
||||
async def test_launch_conflict_is_a_clear_client_error(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8189}]}
|
||||
with patch.object(worker_routes, "get_worker_manager", return_value=manager), \
|
||||
patch.object(worker_routes, "load_config", return_value=config), \
|
||||
patch.object(manager, "launch_worker", side_effect=ValueError("Worker port 8189 conflicts with the master port 8189")):
|
||||
response = await worker_routes.launch_worker_endpoint(_FakeRequest({"worker_id": "worker-a"}))
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("master port 8189", response.payload["message"])
|
||||
self.assertEqual(manager.processes, {})
|
||||
|
||||
async def test_launch_worker_valid_id_returns_200(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8188}]}
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
# This conftest.py marks the tests/ directory as the pytest collection root,
|
||||
# preventing pytest from traversing into the parent package's __init__.py.
|
||||
-691
@@ -1,691 +0,0 @@
|
||||
{
|
||||
"DistributedCollector": {
|
||||
"input": {
|
||||
"required": {
|
||||
"load_balance": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false,
|
||||
"tooltip": "Run this workflow on one least-busy participant (master included when participating)."
|
||||
}
|
||||
]
|
||||
},
|
||||
"optional": {
|
||||
"images": [
|
||||
"IMAGE"
|
||||
],
|
||||
"audio": [
|
||||
"AUDIO"
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"master_url": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"enabled_worker_ids": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "[]"
|
||||
}
|
||||
],
|
||||
"worker_batch_size": [
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 1024
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"pass_through": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"delegate_only": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"load_balance"
|
||||
],
|
||||
"optional": [
|
||||
"images",
|
||||
"audio"
|
||||
],
|
||||
"hidden": [
|
||||
"multi_job_id",
|
||||
"is_worker",
|
||||
"master_url",
|
||||
"enabled_worker_ids",
|
||||
"worker_batch_size",
|
||||
"worker_id",
|
||||
"pass_through",
|
||||
"delegate_only"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE",
|
||||
"AUDIO"
|
||||
],
|
||||
"output_name": [
|
||||
"images",
|
||||
"audio"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": true,
|
||||
"output_node": false,
|
||||
"category": "image",
|
||||
"display_name": "Distributed Collector"
|
||||
},
|
||||
"DistributedSeed": {
|
||||
"input": {
|
||||
"required": {
|
||||
"seed": [
|
||||
"INT",
|
||||
{
|
||||
"default": 1125899906842,
|
||||
"min": 0,
|
||||
"max": 1125899906842624,
|
||||
"forceInput": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"seed"
|
||||
],
|
||||
"hidden": [
|
||||
"is_worker",
|
||||
"worker_id"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"INT"
|
||||
],
|
||||
"output_name": [
|
||||
"seed"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Seed"
|
||||
},
|
||||
"DistributedModelName": {
|
||||
"input": {
|
||||
"required": {
|
||||
"text": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO"
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"text"
|
||||
],
|
||||
"hidden": [
|
||||
"unique_id",
|
||||
"extra_pnginfo"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"output"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Model Name"
|
||||
},
|
||||
"DistributedValue": {
|
||||
"input": {
|
||||
"required": {
|
||||
"default_value": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"worker_values": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "{}"
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"default_value",
|
||||
"worker_values"
|
||||
],
|
||||
"hidden": [
|
||||
"is_worker",
|
||||
"worker_id"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"value"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Value"
|
||||
},
|
||||
"ImageBatchDivider": {
|
||||
"input": {
|
||||
"required": {
|
||||
"images": [
|
||||
"IMAGE"
|
||||
],
|
||||
"divide_by": [
|
||||
"INT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 1,
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of parts to divide the batch into"
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"images",
|
||||
"divide_by"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"batch_1",
|
||||
"batch_2",
|
||||
"batch_3",
|
||||
"batch_4",
|
||||
"batch_5",
|
||||
"batch_6",
|
||||
"batch_7",
|
||||
"batch_8",
|
||||
"batch_9",
|
||||
"batch_10"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "image",
|
||||
"display_name": "Image Batch Divider"
|
||||
},
|
||||
"AudioBatchDivider": {
|
||||
"input": {
|
||||
"required": {
|
||||
"audio": [
|
||||
"AUDIO"
|
||||
],
|
||||
"divide_by": [
|
||||
"INT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 1,
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of sequential time segments to create"
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"audio",
|
||||
"divide_by"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"audio_1",
|
||||
"audio_2",
|
||||
"audio_3",
|
||||
"audio_4",
|
||||
"audio_5",
|
||||
"audio_6",
|
||||
"audio_7",
|
||||
"audio_8",
|
||||
"audio_9",
|
||||
"audio_10"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "audio",
|
||||
"display_name": "Audio Segment Divider"
|
||||
},
|
||||
"DistributedEmptyImage": {
|
||||
"input": {
|
||||
"required": {
|
||||
"height": [
|
||||
"INT",
|
||||
{
|
||||
"default": 64,
|
||||
"min": 1,
|
||||
"max": 4096,
|
||||
"step": 1
|
||||
}
|
||||
],
|
||||
"width": [
|
||||
"INT",
|
||||
{
|
||||
"default": 64,
|
||||
"min": 1,
|
||||
"max": 4096,
|
||||
"step": 1
|
||||
}
|
||||
],
|
||||
"channels": [
|
||||
"INT",
|
||||
{
|
||||
"default": 3,
|
||||
"min": 1,
|
||||
"max": 4,
|
||||
"step": 1
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"height",
|
||||
"width",
|
||||
"channels"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_name": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "image",
|
||||
"display_name": "Distributed Empty Image"
|
||||
},
|
||||
"UltimateSDUpscaleDistributed": {
|
||||
"input": {
|
||||
"required": {
|
||||
"upscaled_image": [
|
||||
"IMAGE"
|
||||
],
|
||||
"model": [
|
||||
"MODEL"
|
||||
],
|
||||
"positive": [
|
||||
"CONDITIONING"
|
||||
],
|
||||
"negative": [
|
||||
"CONDITIONING"
|
||||
],
|
||||
"vae": [
|
||||
"VAE"
|
||||
],
|
||||
"seed": [
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 18446744073709551615
|
||||
}
|
||||
],
|
||||
"steps": [
|
||||
"INT",
|
||||
{
|
||||
"default": 20,
|
||||
"min": 1,
|
||||
"max": 10000
|
||||
}
|
||||
],
|
||||
"cfg": [
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 8.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0
|
||||
}
|
||||
],
|
||||
"sampler_name": [
|
||||
[
|
||||
"euler",
|
||||
"euler_cfg_pp",
|
||||
"euler_ancestral",
|
||||
"euler_ancestral_cfg_pp",
|
||||
"heun",
|
||||
"heunpp2",
|
||||
"exp_heun_2_x0",
|
||||
"exp_heun_2_x0_sde",
|
||||
"dpm_2",
|
||||
"dpm_2_ancestral",
|
||||
"lms",
|
||||
"dpm_fast",
|
||||
"dpm_adaptive",
|
||||
"dpmpp_2s_ancestral",
|
||||
"dpmpp_2s_ancestral_cfg_pp",
|
||||
"dpmpp_sde",
|
||||
"dpmpp_sde_gpu",
|
||||
"dpmpp_2m",
|
||||
"dpmpp_2m_cfg_pp",
|
||||
"dpmpp_2m_sde",
|
||||
"dpmpp_2m_sde_gpu",
|
||||
"dpmpp_2m_sde_heun",
|
||||
"dpmpp_2m_sde_heun_gpu",
|
||||
"dpmpp_3m_sde",
|
||||
"dpmpp_3m_sde_gpu",
|
||||
"ddpm",
|
||||
"lcm",
|
||||
"ipndm",
|
||||
"ipndm_v",
|
||||
"deis",
|
||||
"cfgpp_ud10_ab",
|
||||
"res_multistep",
|
||||
"res_multistep_cfg_pp",
|
||||
"res_multistep_ancestral",
|
||||
"res_multistep_ancestral_cfg_pp",
|
||||
"gradient_estimation",
|
||||
"gradient_estimation_cfg_pp",
|
||||
"er_sde",
|
||||
"seeds_2",
|
||||
"seeds_3",
|
||||
"sa_solver",
|
||||
"sa_solver_pece",
|
||||
"ddim",
|
||||
"uni_pc",
|
||||
"uni_pc_bh2"
|
||||
]
|
||||
],
|
||||
"scheduler": [
|
||||
[
|
||||
"simple",
|
||||
"sgm_uniform",
|
||||
"karras",
|
||||
"exponential",
|
||||
"ddim_uniform",
|
||||
"beta",
|
||||
"normal",
|
||||
"linear_quadratic",
|
||||
"kl_optimal"
|
||||
]
|
||||
],
|
||||
"denoise": [
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.5,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01
|
||||
}
|
||||
],
|
||||
"tile_width": [
|
||||
"INT",
|
||||
{
|
||||
"default": 512,
|
||||
"min": 64,
|
||||
"max": 2048,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"tile_height": [
|
||||
"INT",
|
||||
{
|
||||
"default": 512,
|
||||
"min": 64,
|
||||
"max": 2048,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"padding": [
|
||||
"INT",
|
||||
{
|
||||
"default": 32,
|
||||
"min": 0,
|
||||
"max": 256,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"mask_blur": [
|
||||
"INT",
|
||||
{
|
||||
"default": 8,
|
||||
"min": 0,
|
||||
"max": 256
|
||||
}
|
||||
],
|
||||
"force_uniform_tiles": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": true
|
||||
}
|
||||
],
|
||||
"tiled_decode": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"master_url": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"enabled_worker_ids": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "[]"
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"tile_indices": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"dynamic_threshold": [
|
||||
"INT",
|
||||
{
|
||||
"default": 8,
|
||||
"min": 1,
|
||||
"max": 64
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"upscaled_image",
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
"vae",
|
||||
"seed",
|
||||
"steps",
|
||||
"cfg",
|
||||
"sampler_name",
|
||||
"scheduler",
|
||||
"denoise",
|
||||
"tile_width",
|
||||
"tile_height",
|
||||
"padding",
|
||||
"mask_blur",
|
||||
"force_uniform_tiles",
|
||||
"tiled_decode"
|
||||
],
|
||||
"hidden": [
|
||||
"multi_job_id",
|
||||
"is_worker",
|
||||
"master_url",
|
||||
"enabled_worker_ids",
|
||||
"worker_id",
|
||||
"tile_indices",
|
||||
"dynamic_threshold"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_name": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "image/upscaling",
|
||||
"display_name": "Ultimate SD Upscale Distributed (No Upscale)"
|
||||
}
|
||||
}
|
||||
@@ -1,252 +0,0 @@
|
||||
"""Exercise the actual loader, schemas and executor in a fresh CPU process.
|
||||
|
||||
No HTTP listener, workers, model downloads or live installation changes.
|
||||
The original contracts were captured from 32ac027 using the same core.
|
||||
"""
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import tempfile
|
||||
from unittest.mock import patch
|
||||
|
||||
from PIL import Image
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
COMFY_ROOT = Path(sys.argv[1]).resolve()
|
||||
sys.path.insert(0, str(COMFY_ROOT))
|
||||
os.environ['COMFYUI_IS_WORKER'] = '1'
|
||||
import comfy.cli_args
|
||||
comfy.cli_args.args.cpu = True
|
||||
comfy.cli_args.args.disable_assets = True
|
||||
from app.assets.manager import default_asset_manager
|
||||
from comfy_api.v0_0_2 import io
|
||||
import execution
|
||||
import nodes
|
||||
import server
|
||||
import torch
|
||||
|
||||
BASELINE = json.loads((ROOT / 'tests/fixtures/v1_node_contracts.json').read_text())
|
||||
|
||||
|
||||
def normalized_input(value):
|
||||
kind = value[0]
|
||||
opts = dict(value[1]) if len(value) > 1 else {}
|
||||
if kind == 'STRING':
|
||||
# V3 serializes the same default single-line widget explicitly.
|
||||
opts.setdefault('multiline', False)
|
||||
if kind == 'COMBO':
|
||||
kind = opts['options']
|
||||
opts = {key: val for key, val in opts.items() if key != 'options'}
|
||||
if isinstance(kind, list):
|
||||
# Single-selection is the V1 dropdown default as well.
|
||||
opts.setdefault('multiselect', False)
|
||||
return [kind, opts]
|
||||
|
||||
|
||||
async def map_node(cls, values, extra=None):
|
||||
"""Use the actual input parser and sanitized V3 executor clones."""
|
||||
prepared, missing, hidden = execution.get_input_data(values, cls, unique_id='probe', extra_data=extra or {})
|
||||
assert not missing, missing
|
||||
# Use a worker thread for synchronous nodes: collector/upscale bridge back
|
||||
# to PromptServer's running loop, just as ComfyUI's prompt worker does.
|
||||
def run():
|
||||
return asyncio.run(execution._async_map_node_over_list(
|
||||
'v3-acceptance', 'probe', cls, prepared, cls.FUNCTION, v3_data=hidden))
|
||||
return await asyncio.to_thread(run)
|
||||
|
||||
|
||||
def check_saved_workflows(mapping):
|
||||
checked = 0
|
||||
occurrences = 0
|
||||
for path in sorted((ROOT / 'workflows').glob('*.json')):
|
||||
workflow = json.loads(path.read_text())
|
||||
local_nodes = {node['id']: node for node in workflow['nodes']}
|
||||
for saved in workflow['nodes']:
|
||||
if saved['type'] not in mapping:
|
||||
continue
|
||||
occurrences += 1
|
||||
info = mapping[saved['type']].GET_NODE_INFO_V1()
|
||||
inputs = {**info['input'].get('required', {}), **info['input'].get('optional', {})}
|
||||
for socket in saved.get('inputs', []):
|
||||
assert socket['name'] in inputs, (path.name, saved['id'], socket)
|
||||
assert socket['type'] == inputs[socket['name']][0], (path.name, socket)
|
||||
for index, output in enumerate(saved.get('outputs', [])):
|
||||
assert output['type'] == info['output'][index], (path.name, index)
|
||||
for link in workflow['links']:
|
||||
if link[1] == saved['id']:
|
||||
assert link[2] < len(info['output']), (path.name, link)
|
||||
target = local_nodes[link[3]]['inputs'][link[4]]
|
||||
assert target['type'] == info['output'][link[2]], (path.name, link)
|
||||
checked += 1
|
||||
assert checked == 5 and occurrences == 10, (checked, occurrences)
|
||||
print('SAVED_WORKFLOW_CONTRACTS_OK', checked, occurrences)
|
||||
|
||||
|
||||
async def check_execution(mapping, module, prompt_server, asset_manager):
|
||||
result = await map_node(mapping['DistributedSeed'],
|
||||
{'seed': 123, 'is_worker': True, 'worker_id': 'worker_2'})
|
||||
assert result[0].result == (126,)
|
||||
values = json.dumps({'_type': 'INT', '3': '17'})
|
||||
result = await map_node(mapping['DistributedValue'],
|
||||
{'default_value': '4', 'worker_values': values,
|
||||
'is_worker': True, 'worker_id': 'worker_2'})
|
||||
assert result[0].result == (17,)
|
||||
metadata = {'workflow': {'nodes': [{'id': 'probe', 'widgets_values': []}]}}
|
||||
result = await map_node(mapping['DistributedModelName'], {'text': 'model.ckpt'},
|
||||
{'extra_pnginfo': metadata})
|
||||
assert result[0].result == ('model.ckpt',)
|
||||
assert result[0].ui == {'text': ['model.ckpt']}
|
||||
assert metadata['workflow']['nodes'][0]['widgets_values'] == [['model.ckpt']]
|
||||
result = await map_node(mapping['DistributedEmptyImage'],
|
||||
{'height': 8, 'width': 8, 'channels': 3})
|
||||
empty = result[0].result[0]
|
||||
assert empty.shape == (0, 8, 8, 3) and empty.numel() == 0
|
||||
images = torch.arange(10 * 8 * 8 * 3, dtype=torch.float32).reshape(10, 8, 8, 3)
|
||||
for count in (1, 3, 10):
|
||||
result = await map_node(mapping['ImageBatchDivider'], {'images': images, 'divide_by': count})
|
||||
assert len(result[0].result) == 10
|
||||
assert torch.equal(torch.cat(result[0].result[:count]), images)
|
||||
audio = {'waveform': torch.arange(33, dtype=torch.float32).reshape(1, 1, 33), 'sample_rate': 24000}
|
||||
result = await map_node(mapping['AudioBatchDivider'], {'audio': audio, 'divide_by': 3})
|
||||
assert len(result[0].result) == 10
|
||||
assert torch.equal(torch.cat([item['waveform'] for item in result[0].result[:3]], dim=-1), audio['waveform'])
|
||||
# INPUT_IS_LIST must preserve all images/audio but unwrap transport scalars.
|
||||
collector = mapping['DistributedCollector']
|
||||
prepared, missing, hidden = execution.get_input_data(
|
||||
{'load_balance': False, 'multi_job_id': '', 'pass_through': True}, collector, 'collector-list')
|
||||
assert not missing
|
||||
prepared['images'] = [images[:2], images[2:]]
|
||||
prepared['audio'] = [audio, audio]
|
||||
result = await execution._async_map_node_over_list(
|
||||
'v3-acceptance', 'collector-list', collector, prepared, collector.FUNCTION, v3_data=hidden)
|
||||
assert torch.equal(result[0].result[0], images)
|
||||
assert result[0].result[1]['waveform'].shape[-1] == 66
|
||||
# Exercise real async aggregation without an HTTP endpoint or worker.
|
||||
queue = asyncio.Queue()
|
||||
await queue.put({'worker_id': 'worker_2', 'tensor': images[2:3], 'image_index': 0, 'is_last': True})
|
||||
prompt_server.distributed_pending_jobs['v3-aggregate'] = queue
|
||||
result = await map_node(collector,
|
||||
{'images': images[:2], 'load_balance': False,
|
||||
'multi_job_id': 'v3-aggregate', 'enabled_worker_ids': '["worker_2"]'})
|
||||
assert torch.equal(result[0].result[0], images[:3])
|
||||
assert 'v3-aggregate' not in prompt_server.distributed_pending_jobs
|
||||
|
||||
# Prove V3's per-execution class clones do not leak helper instance state;
|
||||
# replace only the GPU/model boundary, not the input parser or executor.
|
||||
upscale_runtime = sys.modules[module.__name__ + '.nodes.distributed_upscale'].UltimateSDUpscaleDistributed
|
||||
seen = []
|
||||
def fake_upscale(self, *args):
|
||||
assert not hasattr(self, 'acceptance_marker')
|
||||
self.acceptance_marker = True
|
||||
seen.append(args)
|
||||
return (args[0],)
|
||||
inputs = {'upscaled_image': images, 'model': object(), 'positive': [], 'negative': [],
|
||||
'vae': object(), 'seed': 1, 'steps': 1, 'cfg': 1.0,
|
||||
'sampler_name': 'euler', 'scheduler': 'normal', 'denoise': 0.5,
|
||||
'tile_width': 64, 'tile_height': 64, 'padding': 0, 'mask_blur': 0,
|
||||
'force_uniform_tiles': True, 'tiled_decode': False,
|
||||
'multi_job_id': 'tile-job', 'is_worker': True, 'master_url': 'http://master.invalid',
|
||||
'enabled_worker_ids': '["worker_2"]', 'worker_id': 'worker_2',
|
||||
'tile_indices': '[2,3]', 'dynamic_threshold': 9}
|
||||
with patch.object(upscale_runtime, 'run', fake_upscale):
|
||||
for _ in range(2):
|
||||
result = await map_node(mapping['UltimateSDUpscaleDistributed'], inputs)
|
||||
assert result[0].result[0] is images
|
||||
assert len(seen) == 2 and seen[0][-7:] == tuple(inputs[name] for name in (
|
||||
'multi_job_id', 'is_worker', 'master_url', 'enabled_worker_ids', 'worker_id',
|
||||
'tile_indices', 'dynamic_threshold')), seen[0]
|
||||
import math
|
||||
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id='tile-job'))
|
||||
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id=''))
|
||||
|
||||
graph = {
|
||||
'1': {'class_type': 'EmptyImage', 'inputs': {'height': 8, 'width': 8, 'batch_size': 10, 'color': 0}},
|
||||
'2': {'class_type': 'DistributedCollector', 'inputs': {'images': ['1', 0], 'load_balance': False,
|
||||
'pass_through': True}},
|
||||
'3': {'class_type': 'ImageBatchDivider', 'inputs': {'images': ['2', 0], 'divide_by': 10}},
|
||||
'4': {'class_type': 'PreviewImage', 'inputs': {'images': ['3', 9]}},
|
||||
}
|
||||
valid = await execution.validate_prompt('v3-graph', graph, None)
|
||||
assert valid[0], valid
|
||||
invalid = {'1': {'class_type': 'UltimateSDUpscaleDistributed',
|
||||
'inputs': {**inputs, 'sampler_name': 'INVALID_ENUM', 'scheduler': 'INVALID_ENUM'}},
|
||||
'2': {'class_type': 'PreviewImage', 'inputs': {'images': ['1', 0]}}}
|
||||
rejected = await execution.validate_prompt('v3-invalid', invalid, None)
|
||||
assert not rejected[0]
|
||||
errors = [error for entry in rejected[3].values() for error in entry['errors']]
|
||||
enum_errors = [error['extra_info']['input_name'] for error in errors if error['type'] == 'value_not_in_list']
|
||||
assert {'sampler_name', 'scheduler'} <= set(enum_errors), errors
|
||||
import folder_paths
|
||||
scratch = Path(os.environ.get('TMPDIR', Path.home() / '.hermes/cache/scratch'))
|
||||
with tempfile.TemporaryDirectory(prefix='v3-preview-', dir=scratch) as temp:
|
||||
with patch.object(folder_paths, 'temp_directory', temp):
|
||||
executor = execution.PromptExecutor(
|
||||
prompt_server, cache_args={'ram': 0, 'ram_inactive': 0}, asset_manager=asset_manager)
|
||||
await asyncio.to_thread(executor.execute, graph, 'v3-graph', {}, valid[2])
|
||||
assert executor.success, executor.status_messages
|
||||
history = executor.history_result
|
||||
record = history['outputs']['4']['images'][0]
|
||||
preview = Path(temp) / record.get('subfolder', '') / record['filename']
|
||||
with Image.open(preview) as image:
|
||||
assert image.size == (8, 8) and image.mode == 'RGB'
|
||||
assert image.getextrema() == ((0, 0), (0, 0), (0, 0))
|
||||
print('V3_EXECUTION_OK eight nodes; upscale GPU boundary mocked; preview artifact verified')
|
||||
|
||||
|
||||
async def main():
|
||||
asset_manager = default_asset_manager()
|
||||
prompt_server = server.PromptServer(asyncio.get_running_loop(), asset_manager)
|
||||
assert not (ROOT / 'distributed.py').exists(), 'obsolete root bootstrap remains'
|
||||
assert await nodes.load_custom_node(str(ROOT)), 'ComfyUI loader rejected the pack'
|
||||
module = sys.modules[str(ROOT).replace('.', '_x_')]
|
||||
assert not hasattr(module, 'NODE_CLASS_MAPPINGS'), 'V1 map shadows V3 entrypoint'
|
||||
extension = await module.comfy_entrypoint()
|
||||
classes = await extension.get_node_list()
|
||||
mapping = {cls.GET_SCHEMA().node_id: cls for cls in classes}
|
||||
assert len(classes) == len(mapping) == len(BASELINE) == 8
|
||||
assert set(mapping) == set(BASELINE)
|
||||
for node_id, cls in mapping.items():
|
||||
assert issubclass(cls, io.ComfyNode)
|
||||
assert nodes.NODE_CLASS_MAPPINGS[node_id] is cls
|
||||
old = dict(BASELINE[node_id])
|
||||
if node_id in ('ImageBatchDivider', 'AudioBatchDivider'):
|
||||
# V1's ByPassTypeTuple advertises '*' when indexed, while its
|
||||
# underlying tuple and existing frontend declare IMAGE/AUDIO.
|
||||
# Native V3 declares all ten existing typed sockets explicitly.
|
||||
assert old['output'] == ['*'] * 10
|
||||
old['output'] = ['IMAGE' if node_id == 'ImageBatchDivider' else 'AUDIO'] * 10
|
||||
new = json.loads(json.dumps(cls.GET_NODE_INFO_V1()))
|
||||
for group in ('required', 'optional'):
|
||||
old_inputs = old['input'].get(group, {})
|
||||
new_inputs = new['input'].get(group, {})
|
||||
assert list(old_inputs) == list(new_inputs), (node_id, group, 'input order')
|
||||
for name, original in old_inputs.items():
|
||||
assert normalized_input(original) == normalized_input(new_inputs[name]), (node_id, name, original, new_inputs[name])
|
||||
for key in ('output', 'output_name', 'output_is_list', 'is_input_list', 'output_node', 'category', 'display_name'):
|
||||
assert old[key] == new[key], (node_id, key, old[key], new[key])
|
||||
# Standard context lives in cls.hidden; orchestrator metadata remains
|
||||
# accepted by its original kwarg name, without creating new widgets.
|
||||
signature = inspect.signature(cls.execute)
|
||||
for name, field in old['input'].get('hidden', {}).items():
|
||||
if isinstance(field, list):
|
||||
assert cls.GET_SCHEMA().accept_all_inputs, node_id
|
||||
assert name in signature.parameters, (node_id, name)
|
||||
assert signature.parameters[name].default == field[1]['default'], (node_id, name)
|
||||
else:
|
||||
assert name in new['input']['hidden'], (node_id, name)
|
||||
expected_hidden = []
|
||||
if node_id == 'DistributedModelName':
|
||||
expected_hidden.extend(['unique_id', 'extra_pnginfo'])
|
||||
if old['output_node']:
|
||||
expected_hidden.extend(name for name in ['prompt', 'extra_pnginfo'] if name not in expected_hidden)
|
||||
assert list(new['input'].get('hidden', {})) == expected_hidden, (node_id, new['input'].get('hidden'))
|
||||
print('SCHEMA_PARITY_OK', len(mapping))
|
||||
check_saved_workflows(mapping)
|
||||
await check_execution(mapping, module, prompt_server, asset_manager)
|
||||
print('V3_ACCEPTANCE_OK')
|
||||
|
||||
|
||||
asyncio.run(main())
|
||||
@@ -1,364 +0,0 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _load_collector_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "collector.py"
|
||||
package_name = "dist_collector_list_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
class _Routes:
|
||||
def post(self, _path):
|
||||
return lambda fn: fn
|
||||
|
||||
def get(self, _path):
|
||||
return lambda fn: fn
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
routes=_Routes(),
|
||||
distributed_jobs_lock=None,
|
||||
distributed_pending_jobs={},
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
model_management = types.ModuleType("comfy.model_management")
|
||||
|
||||
class InterruptProcessingException(Exception):
|
||||
pass
|
||||
|
||||
model_management.InterruptProcessingException = InterruptProcessingException
|
||||
model_management.throw_exception_if_processing_interrupted = lambda: None
|
||||
comfy_module.model_management = model_management
|
||||
|
||||
comfy_utils = types.ModuleType("comfy.utils")
|
||||
|
||||
class ProgressBar:
|
||||
def __init__(self, _total):
|
||||
self.total = _total
|
||||
self.updates = []
|
||||
|
||||
def update(self, value):
|
||||
self.updates.append(value)
|
||||
|
||||
comfy_utils.ProgressBar = ProgressBar
|
||||
comfy_module.utils = comfy_utils
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.model_management"] = model_management
|
||||
sys.modules["comfy.utils"] = comfy_utils
|
||||
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
aiohttp_module.ClientTimeout = lambda total: types.SimpleNamespace(total=total)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 0.1
|
||||
config_module.load_config = lambda: {"workers": []}
|
||||
config_module.is_master_delegate_only = lambda: False
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.HEARTBEAT_INTERVAL = 1.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
def _ensure_contiguous(tensor):
|
||||
return tensor.contiguous() if hasattr(tensor, "contiguous") else tensor
|
||||
|
||||
image_module.ensure_contiguous = _ensure_contiguous
|
||||
image_module.tensor_to_pil = lambda *_args, **_kwargs: None
|
||||
image_module.pil_to_tensor = lambda value: value
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.build_worker_url = lambda worker: "http://worker"
|
||||
network_module.get_client_session = lambda: None
|
||||
network_module.probe_worker = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
audio_payload_module = types.ModuleType(f"{package_name}.utils.audio_payload")
|
||||
audio_payload_module.encode_audio_payload = lambda audio: audio
|
||||
sys.modules[f"{package_name}.utils.audio_payload"] = audio_payload_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.nodes.collector", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def test_collector_opts_into_comfyui_list_inputs():
|
||||
collector = _load_collector_module().DistributedCollectorNode
|
||||
|
||||
assert collector.INPUT_IS_LIST is True
|
||||
|
||||
|
||||
def test_collector_exposes_images_as_optional_input():
|
||||
input_types = _load_collector_module().DistributedCollectorNode.INPUT_TYPES()
|
||||
|
||||
assert "images" not in input_types["required"]
|
||||
assert input_types["optional"]["images"] == ("IMAGE",)
|
||||
|
||||
|
||||
def test_audio_only_pass_through_returns_no_images_and_preserves_audio():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
|
||||
|
||||
images, returned_audio = collector.run(images=None, audio=[audio])
|
||||
|
||||
assert images is None
|
||||
assert returned_audio is audio
|
||||
|
||||
|
||||
def test_collector_rejects_missing_images_and_audio():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
|
||||
try:
|
||||
collector.run(images=None, audio=None)
|
||||
except ValueError as exc:
|
||||
assert "image or audio" in str(exc).lower()
|
||||
else:
|
||||
raise AssertionError("Expected collector to reject a run with no media input")
|
||||
|
||||
|
||||
def test_delegate_only_master_allows_no_local_media_input():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
|
||||
images, audio = collector.run(
|
||||
images=None,
|
||||
audio=None,
|
||||
multi_job_id=["delegate-audio-job"],
|
||||
delegate_only=[True],
|
||||
enabled_worker_ids=["[]"],
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
|
||||
|
||||
def test_pass_through_collapses_comfyui_image_list_to_batch_and_unwraps_hidden_inputs():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
first = torch.zeros(1, 2, 2, 3)
|
||||
second = torch.ones(1, 2, 2, 3)
|
||||
|
||||
images, audio = collector.run(
|
||||
images=[first, second],
|
||||
load_balance=[False],
|
||||
audio=[None],
|
||||
multi_job_id=[""],
|
||||
is_worker=[False],
|
||||
master_url=[""],
|
||||
enabled_worker_ids=["[]"],
|
||||
worker_batch_size=[1],
|
||||
worker_id=[""],
|
||||
pass_through=[False],
|
||||
delegate_only=[False],
|
||||
)
|
||||
|
||||
assert tuple(images.shape) == (2, 2, 2, 3)
|
||||
assert torch.equal(images[0:1], first)
|
||||
assert torch.equal(images[1:2], second)
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
|
||||
|
||||
def test_worker_list_input_sends_one_completion_sequence_with_last_only_on_final_item():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
first = torch.zeros(1, 2, 2, 3)
|
||||
second = torch.ones(1, 2, 2, 3)
|
||||
posted_payloads = []
|
||||
|
||||
class _FakeImage:
|
||||
def save(self, fp, format=None, compress_level=None):
|
||||
fp.write(b"png-bytes")
|
||||
|
||||
class _FakeResponse:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
class _FakeSession:
|
||||
def post(self, url, json, timeout):
|
||||
posted_payloads.append(json)
|
||||
return _FakeResponse()
|
||||
|
||||
async def _fake_get_client_session():
|
||||
return _FakeSession()
|
||||
|
||||
module.tensor_to_pil = lambda *_args, **_kwargs: _FakeImage()
|
||||
module.get_client_session = _fake_get_client_session
|
||||
module.encode_audio_payload = lambda _audio: None
|
||||
|
||||
images, audio = collector.run(
|
||||
images=[first, second],
|
||||
load_balance=[False],
|
||||
audio=[None],
|
||||
multi_job_id=["job-list-1"],
|
||||
is_worker=[True],
|
||||
master_url=["http://master"],
|
||||
enabled_worker_ids=["[]"],
|
||||
worker_batch_size=[1],
|
||||
worker_id=["worker-a"],
|
||||
pass_through=[False],
|
||||
delegate_only=[False],
|
||||
)
|
||||
|
||||
assert tuple(images.shape) == (2, 2, 2, 3)
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
assert len(posted_payloads) == 2
|
||||
assert [payload["batch_idx"] for payload in posted_payloads] == [0, 1]
|
||||
assert [payload["is_last"] for payload in posted_payloads] == [False, True]
|
||||
assert {payload["job_id"] for payload in posted_payloads} == {"job-list-1"}
|
||||
assert {payload["worker_id"] for payload in posted_payloads} == {"worker-a"}
|
||||
|
||||
|
||||
def test_audio_only_worker_sends_one_completion_without_image():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
|
||||
posted = []
|
||||
|
||||
class _FakeResponse:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
class _FakeSession:
|
||||
def post(self, url, json, timeout):
|
||||
posted.append((url, json, timeout.total))
|
||||
return _FakeResponse()
|
||||
|
||||
async def _fake_get_client_session():
|
||||
return _FakeSession()
|
||||
|
||||
module.get_client_session = _fake_get_client_session
|
||||
module.encode_audio_payload = lambda value: {"encoded": value is audio}
|
||||
|
||||
images, returned_audio = collector.run(
|
||||
images=None,
|
||||
audio=[audio],
|
||||
multi_job_id=["audio-job"],
|
||||
is_worker=[True],
|
||||
master_url=["http://master"],
|
||||
worker_id=["worker-a"],
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert returned_audio is audio
|
||||
assert len(posted) == 1
|
||||
assert posted[0][0] == "http://master/distributed/job_complete"
|
||||
assert posted[0][2] == 600
|
||||
assert posted[0][1] == {
|
||||
"job_id": "audio-job",
|
||||
"worker_id": "worker-a",
|
||||
"batch_idx": 0,
|
||||
"audio": {"encoded": True},
|
||||
"is_last": True,
|
||||
}
|
||||
|
||||
|
||||
def test_audio_only_master_combines_local_and_worker_audio():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
master_audio = {"waveform": torch.ones(1, 2, 2), "sample_rate": 48000}
|
||||
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
|
||||
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
queue = asyncio.Queue()
|
||||
queue.put_nowait(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"image_index": 0,
|
||||
"tensor": None,
|
||||
"audio": worker_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
module.prompt_server.distributed_pending_jobs = {"audio-job": queue}
|
||||
|
||||
images, combined_audio = asyncio.run(
|
||||
collector.execute(
|
||||
images=None,
|
||||
audio=master_audio,
|
||||
multi_job_id="audio-job",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
)
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert combined_audio["sample_rate"] == 48000
|
||||
assert tuple(combined_audio["waveform"].shape) == (1, 2, 5)
|
||||
assert torch.equal(combined_audio["waveform"][..., :2], master_audio["waveform"])
|
||||
assert torch.equal(combined_audio["waveform"][..., 2:], worker_audio["waveform"])
|
||||
|
||||
|
||||
def test_delegate_only_audio_collects_worker_audio_without_placeholder_image():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
|
||||
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
queue = asyncio.Queue()
|
||||
queue.put_nowait(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"image_index": 0,
|
||||
"tensor": None,
|
||||
"audio": worker_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
module.prompt_server.distributed_pending_jobs = {"delegate-audio-job": queue}
|
||||
|
||||
images, combined_audio = asyncio.run(
|
||||
collector.execute(
|
||||
images=None,
|
||||
audio=None,
|
||||
multi_job_id="delegate-audio-job",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
delegate_only=True,
|
||||
)
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert combined_audio["sample_rate"] == 48000
|
||||
assert torch.equal(combined_audio["waveform"], worker_audio["waveform"])
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,21 +0,0 @@
|
||||
"""Real-framework acceptance; opt in with COMFYUI_SOURCE_ROOT."""
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_native_v3_runtime():
|
||||
comfy_root = os.environ.get('COMFYUI_SOURCE_ROOT')
|
||||
if not comfy_root:
|
||||
pytest.skip('Set COMFYUI_SOURCE_ROOT to run native ComfyUI V3 acceptance')
|
||||
helper = Path(__file__).parent / 'helpers' / 'v3_runtime_check.py'
|
||||
result = subprocess.run(
|
||||
[sys.executable, str(helper), str(Path(comfy_root).resolve())],
|
||||
capture_output=True, text=True, timeout=120,
|
||||
env={**os.environ, 'COMFYUI_IS_WORKER': '1'},
|
||||
)
|
||||
assert result.returncode == 0, result.stdout + '\n' + result.stderr
|
||||
assert 'V3_ACCEPTANCE_OK' in result.stdout
|
||||
@@ -1,113 +0,0 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import socket
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def load_ports():
|
||||
path = Path(__file__).resolve().parents[1] / "workers" / "ports.py"
|
||||
spec = importlib.util.spec_from_file_location("worker_ports_test", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
class WorkerPortTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.ports = load_ports()
|
||||
|
||||
def test_allocation_starts_above_actual_master_and_skips_assigned_and_occupied(self):
|
||||
workers = [
|
||||
{"id": "existing", "host": "localhost", "port": 8190, "enabled": False},
|
||||
{"id": "remote", "host": "remote.example", "port": 8192},
|
||||
]
|
||||
with patch.object(self.ports, "is_port_available", side_effect=lambda port: port != 8191):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, workers, 3), [8192, 8193, 8194])
|
||||
|
||||
def test_local_host_forms_reserve_ports(self):
|
||||
for host in ["::1", "[::1]", "http://localhost/", "HTTPS://LOCALHOST", "0.0.0.0", None]:
|
||||
with self.subTest(host=host), patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, [{"id": "local", "host": host, "port": 8190}], 1), [8191])
|
||||
|
||||
def test_exhaustion_is_explicit_not_a_partial_allocation(self):
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
with self.assertRaisesRegex(ValueError, "available.*ports"):
|
||||
self.ports.allocate_worker_ports(65534, [], 2)
|
||||
|
||||
def test_launch_rejects_master_port_without_reassigning_configuration(self):
|
||||
worker = {"id": "manual", "port": 8189}
|
||||
with self.assertRaisesRegex(ValueError, "master.*8189"):
|
||||
self.ports.validate_worker_port(worker, 8189, [])
|
||||
self.assertEqual(worker["port"], 8189)
|
||||
|
||||
def test_launch_rejects_another_local_workers_reserved_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "other", "host": None, "port": 8190, "enabled": False}
|
||||
with self.assertRaisesRegex(ValueError, "other"):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_ignores_itself_and_remote_workers_with_same_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "remote", "host": "example.com", "port": 8190}
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_rejects_an_occupied_port(self):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener:
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen()
|
||||
port = listener.getsockname()[1]
|
||||
self.assertFalse(self.ports.is_port_available(port))
|
||||
with self.assertRaisesRegex(ValueError, "already in use"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
@unittest.skipIf(os.name == "nt", "asyncio does not reuse addresses on Windows")
|
||||
def test_recently_closed_connection_does_not_block_worker_restart(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
port = listener.getsockname()[1]
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=2) as client:
|
||||
accepted, _ = listener.accept()
|
||||
with accepted:
|
||||
accepted.settimeout(2)
|
||||
accepted.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(client.recv(1), b"")
|
||||
client.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(accepted.recv(1), b"")
|
||||
self.assertTrue(self.ports.is_port_available(port))
|
||||
self.ports.validate_worker_port({"id": "restart", "port": port}, 1, [])
|
||||
with socket.socket() as restarted:
|
||||
restarted.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
restarted.bind(("127.0.0.1", port))
|
||||
restarted.listen(1)
|
||||
|
||||
def test_active_listener_with_reuseaddr_is_still_a_conflict(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
@unittest.skipUnless(socket.has_ipv6, "IPv6 unavailable")
|
||||
def test_ipv6_only_listener_is_still_a_conflict(self):
|
||||
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as listener:
|
||||
listener.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
|
||||
try:
|
||||
listener.bind(("::1", 0))
|
||||
except OSError as exc:
|
||||
self.skipTest(f"IPv6 loopback unavailable: {exc}")
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
def test_invalid_ports_fail_clearly(self):
|
||||
for port in [0, 65536, "not-a-port", None, True, 8190.5]:
|
||||
with self.subTest(port=port), self.assertRaisesRegex(ValueError, "port"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,352 +0,0 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
from argparse import Namespace
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
|
||||
|
||||
def _load_process_module(module_filename: str):
|
||||
module_path = Path(__file__).resolve().parents[1] / "workers" / "process" / module_filename
|
||||
package_name = "dist_proc_testpkg"
|
||||
module_name = module_filename[:-3]
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = [str(module_path.parents[1])]
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
process_pkg = types.ModuleType(f"{package_name}.workers.process")
|
||||
process_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.workers.process"] = process_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
process_module = types.ModuleType(f"{package_name}.utils.process")
|
||||
process_module.get_python_executable = lambda: "/usr/bin/test-python"
|
||||
process_module.is_process_alive = lambda _pid: False
|
||||
process_module.terminate_process = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.process"] = process_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"workers": [], "settings": {"stop_workers_on_master_exit": False}}
|
||||
config_module.save_config = lambda _config: None
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.PROCESS_TERMINATION_TIMEOUT = 1
|
||||
constants_module.PROCESS_WAIT_TIMEOUT = 1
|
||||
constants_module.WORKER_CHECK_INTERVAL = 0.01
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.get_server_port = lambda: 8189
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.workers.process.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
root_discovery_module = _load_process_module("root_discovery.py")
|
||||
launch_builder_module = _load_process_module("launch_builder.py")
|
||||
|
||||
|
||||
class ComfyRootDiscoveryTests(unittest.TestCase):
|
||||
def test_prefers_loaded_comfyui_module_path(self):
|
||||
discovery = root_discovery_module.ComfyRootDiscovery()
|
||||
server_module = types.SimpleNamespace(__file__="/opt/ComfyUI/server.py")
|
||||
|
||||
def fake_exists(path):
|
||||
return path == "/opt/ComfyUI/main.py"
|
||||
|
||||
with patch.dict(sys.modules, {"server": server_module}, clear=False), \
|
||||
patch.object(root_discovery_module.os.path, "exists", side_effect=fake_exists), \
|
||||
patch.dict(root_discovery_module.os.environ, {}, clear=True):
|
||||
self.assertEqual(discovery.find_comfy_root(), "/opt/ComfyUI")
|
||||
|
||||
|
||||
class LaunchCommandBuilderTests(unittest.TestCase):
|
||||
def build_command(self, root, worker=None, runtime=None):
|
||||
builder = launch_builder_module.LaunchCommandBuilder()
|
||||
worker = worker or {"id": "worker-a", "port": 8190}
|
||||
runtime = runtime if runtime is not None else Namespace(database_url=None)
|
||||
with patch.object(builder, "_get_runtime_args", return_value=runtime):
|
||||
return builder.build_launch_command(worker, str(root))
|
||||
|
||||
def test_database_is_unique_and_stable_by_worker_id(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
first = self.build_command(root, {"id": "worker-a", "port": 8190})
|
||||
renamed = self.build_command(root, {
|
||||
"id": "worker-a", "port": 9001, "name": "Renamed", "cuda_device": 3,
|
||||
})
|
||||
second = self.build_command(root, {"id": "worker-b", "port": 8191})
|
||||
database = first[first.index("--database-url") + 1]
|
||||
self.assertEqual(database, renamed[renamed.index("--database-url") + 1])
|
||||
self.assertNotEqual(database, second[second.index("--database-url") + 1])
|
||||
self.assertTrue(database.startswith("sqlite:///" + root.as_posix() + "/"))
|
||||
self.assertTrue(Path(database.removeprefix("sqlite:///")).parent.is_dir())
|
||||
self.assertNotIn("--disable-assets", first)
|
||||
|
||||
def test_database_uses_effective_user_or_base_directory(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra_args, expected in [
|
||||
(f'--user-directory "{root / "custom user"}"', root / "custom user"),
|
||||
(f'--base-directory "{root / "custom base"}"', root / "custom base" / "user"),
|
||||
]:
|
||||
with self.subTest(extra_args=extra_args):
|
||||
cmd = self.build_command(root, {
|
||||
"id": "worker-a", "port": 8190, "extra_args": extra_args,
|
||||
})
|
||||
database = cmd[cmd.index("--database-url") + 1]
|
||||
self.assertTrue(database.startswith("sqlite:///" + expected.as_posix() + "/"))
|
||||
cmd = self.build_command(root, runtime=Namespace(
|
||||
database_url="sqlite:///master.db", user_directory=str(root / "runtime user"),
|
||||
))
|
||||
database = cmd[cmd.index("--database-url") + 1]
|
||||
self.assertTrue(database.startswith("sqlite:///" + (root / "runtime user").as_posix() + "/"))
|
||||
self.assertNotEqual(database, "sqlite:///master.db")
|
||||
|
||||
def test_relative_database_directories_use_worker_cwd(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory) / "ComfyUI"
|
||||
root.mkdir()
|
||||
(root / "main.py").touch()
|
||||
launcher = Path(directory) / "launcher"
|
||||
launcher.mkdir()
|
||||
cases = [
|
||||
(Namespace(database_url=None, user_directory="profiles"), "", root / "profiles"),
|
||||
(Namespace(database_url=None, base_directory="data"), "", root / "data/user"),
|
||||
(Namespace(database_url=None), "--user-directory=profiles", root / "profiles"),
|
||||
(Namespace(database_url=None), "--base-directory data", root / "data/user"),
|
||||
]
|
||||
for runtime, extra, expected in cases:
|
||||
with self.subTest(runtime=runtime, extra=extra), \
|
||||
patch.object(os, "getcwd", return_value=str(launcher)):
|
||||
cmd = self.build_command(root, {
|
||||
"id": "worker-a", "port": 8190, "extra_args": extra,
|
||||
}, runtime)
|
||||
database = make_url(cmd[cmd.index("--database-url") + 1]).database
|
||||
self.assertEqual(Path(database).parent, expected / "distributed/workers")
|
||||
self.assertFalse((launcher / "profiles").exists())
|
||||
self.assertFalse((launcher / "data").exists())
|
||||
|
||||
def test_database_url_round_trips_special_characters_without_collisions(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for name in ["user?profile", "user%3Fprofile"]:
|
||||
with self.subTest(name=name):
|
||||
user = root / name
|
||||
runtime = Namespace(database_url=None, user_directory=str(user))
|
||||
# SQLAlchemy 2.0 cannot round-trip '?' in a filename. Fail
|
||||
# clearly on those versions rather than silently sharing a DB.
|
||||
sample = (user / "test.db").as_posix()
|
||||
serialized = URL.create("sqlite", database=sample).render_as_string()
|
||||
if make_url(serialized).database != sample:
|
||||
with self.assertRaisesRegex(ValueError, "SQLAlchemy.*database path"):
|
||||
self.build_command(root, runtime=runtime)
|
||||
self.assertFalse(user.exists())
|
||||
continue
|
||||
databases = []
|
||||
for worker_id in ["worker-a", "worker-b"]:
|
||||
cmd = self.build_command(root, {"id": worker_id, "port": 8190}, runtime)
|
||||
url = cmd[cmd.index("--database-url") + 1]
|
||||
database = make_url(url).database
|
||||
self.assertEqual(Path(database).parent, user / "distributed/workers")
|
||||
engine = create_engine(url)
|
||||
try:
|
||||
with engine.connect() as connection:
|
||||
actual = connection.exec_driver_sql("PRAGMA database_list").one()[2]
|
||||
self.assertEqual(Path(actual), Path(database))
|
||||
finally:
|
||||
engine.dispose()
|
||||
databases.append(database)
|
||||
self.assertNotEqual(*databases)
|
||||
|
||||
def test_explicit_database_and_disabled_assets_are_preserved(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra in ["--database-url sqlite:///explicit.db", "--database-url=sqlite:///explicit.db", "--disable-assets"]:
|
||||
with self.subTest(extra=extra):
|
||||
cmd = self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
|
||||
self.assertEqual(sum(arg.split("=")[0] == "--database-url" for arg in cmd), 0 if extra == "--disable-assets" else 1)
|
||||
self.assertFalse((root / "user").exists())
|
||||
|
||||
def test_older_comfyui_without_database_flag_keeps_existing_command(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
cmd = self.build_command(root, runtime=Namespace())
|
||||
self.assertNotIn("--database-url", cmd)
|
||||
self.assertFalse((root / "user").exists())
|
||||
|
||||
def test_worker_id_cannot_escape_database_directory(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
databases = []
|
||||
for worker_id in ["../outside", "a/b", "a_b"]:
|
||||
cmd = self.build_command(root, {"id": worker_id, "port": 8190})
|
||||
path = Path(cmd[cmd.index("--database-url") + 1].removeprefix("sqlite:///"))
|
||||
self.assertTrue(path.is_relative_to(root / "user"))
|
||||
databases.append(path)
|
||||
self.assertEqual(len(set(databases)), 3)
|
||||
|
||||
def test_extra_args_cannot_silently_override_the_configured_port(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra in ["--port 8189", "--port=8189"]:
|
||||
with self.subTest(extra=extra), self.assertRaisesRegex(ValueError, "configured worker port"):
|
||||
self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
|
||||
|
||||
def test_inherits_runtime_layout_args_for_desktop(self):
|
||||
builder = launch_builder_module.LaunchCommandBuilder()
|
||||
runtime_args = Namespace(
|
||||
listen="127.0.0.1",
|
||||
base_directory="C:/Users/test/ComfyUI",
|
||||
temp_directory=None,
|
||||
input_directory="C:/Users/test/ComfyUI/input",
|
||||
output_directory="C:/Users/test/ComfyUI/output",
|
||||
user_directory="C:/Users/test/ComfyUI/user",
|
||||
front_end_root="C:/Program Files/ComfyUI/web_custom_versions/desktop_app",
|
||||
extra_model_paths_config=[["C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml"]],
|
||||
enable_manager=True,
|
||||
disable_manager_ui=False,
|
||||
enable_manager_legacy_ui=False,
|
||||
windows_standalone_build=True,
|
||||
log_stdout=True,
|
||||
verbose="INFO",
|
||||
enable_cors_header="*",
|
||||
)
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
comfy_cli_args = types.ModuleType("comfy.cli_args")
|
||||
comfy_cli_args.args = runtime_args
|
||||
|
||||
worker_config = {
|
||||
"port": 9001,
|
||||
"extra_args": "--preview-method auto",
|
||||
}
|
||||
|
||||
def fake_exists(path):
|
||||
return path == "/desktop/ComfyUI/main.py"
|
||||
|
||||
with patch.dict(
|
||||
sys.modules,
|
||||
{"comfy": comfy_module, "comfy.cli_args": comfy_cli_args},
|
||||
clear=False,
|
||||
), patch.object(launch_builder_module.os.path, "exists", side_effect=fake_exists):
|
||||
cmd = builder.build_launch_command(worker_config, "/desktop/ComfyUI")
|
||||
|
||||
self.assertEqual(cmd[:2], ["/usr/bin/test-python", "/desktop/ComfyUI/main.py"])
|
||||
self.assertIn("--listen", cmd)
|
||||
self.assertIn("127.0.0.1", cmd)
|
||||
self.assertIn("--base-directory", cmd)
|
||||
self.assertIn("C:/Users/test/ComfyUI", cmd)
|
||||
self.assertIn("--input-directory", cmd)
|
||||
self.assertIn("--output-directory", cmd)
|
||||
self.assertIn("--user-directory", cmd)
|
||||
self.assertIn("--front-end-root", cmd)
|
||||
self.assertIn("--extra-model-paths-config", cmd)
|
||||
self.assertIn("C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml", cmd)
|
||||
self.assertIn("--enable-manager", cmd)
|
||||
self.assertIn("--windows-standalone-build", cmd)
|
||||
self.assertIn("--log-stdout", cmd)
|
||||
self.assertIn("--disable-auto-launch", cmd)
|
||||
self.assertIn("--enable-cors-header", cmd)
|
||||
self.assertIn("*", cmd)
|
||||
self.assertIn("--port", cmd)
|
||||
self.assertIn("9001", cmd)
|
||||
self.assertNotIn("--auto-launch", cmd)
|
||||
|
||||
|
||||
class ProcessLaunchPortTests(unittest.TestCase):
|
||||
def test_launch_before_master_listens_uses_configured_master_port(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
cli_args = types.ModuleType("comfy.cli_args")
|
||||
cli_args.args = Namespace(port=8189)
|
||||
for missing in [AttributeError("PromptServer has no port"), None]:
|
||||
with self.subTest(missing=missing), tempfile.TemporaryDirectory() as directory:
|
||||
manager = types.SimpleNamespace(
|
||||
find_comfy_root=lambda: directory,
|
||||
build_launch_command=lambda _worker, _root: ["python", "main.py"],
|
||||
processes={}, save_processes=lambda: None,
|
||||
)
|
||||
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
|
||||
patch.object(lifecycle_module, "get_server_port", side_effect=missing if isinstance(missing, Exception) else None, return_value=None), \
|
||||
patch.object(lifecycle_module, "validate_worker_port") as validate, \
|
||||
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
|
||||
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
|
||||
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
|
||||
spawn.assert_called_once()
|
||||
self.assertEqual(pid, 1234)
|
||||
|
||||
def test_pre_listen_launch_still_rejects_the_configured_master_port(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
cli_args = types.ModuleType("comfy.cli_args")
|
||||
cli_args.args = Namespace(port=8189)
|
||||
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
|
||||
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
|
||||
patch.object(lifecycle_module, "get_server_port", side_effect=AttributeError("port")), \
|
||||
patch.object(lifecycle_module.subprocess, "Popen") as spawn:
|
||||
with self.assertRaisesRegex(ValueError, "master port 8189"):
|
||||
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
|
||||
spawn.assert_not_called()
|
||||
|
||||
def test_conflict_is_rejected_before_building_or_spawning(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
|
||||
with patch.object(lifecycle_module.subprocess, "Popen") as spawn:
|
||||
with self.assertRaisesRegex(ValueError, "master port 8189"):
|
||||
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
|
||||
spawn.assert_not_called()
|
||||
self.assertEqual(manager.processes, {})
|
||||
|
||||
def test_available_port_reaches_existing_launch_path(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
manager = types.SimpleNamespace(
|
||||
find_comfy_root=lambda: directory,
|
||||
build_launch_command=lambda _worker, _root: ["python", "main.py", "--port", "8190"],
|
||||
processes={}, save_processes=lambda: None,
|
||||
)
|
||||
with patch.object(lifecycle_module, "validate_worker_port") as validate, \
|
||||
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
|
||||
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
|
||||
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
|
||||
spawn.assert_called_once()
|
||||
self.assertEqual(pid, 1234)
|
||||
self.assertEqual(manager.processes["manual"]["pid"], 1234)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests."""
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
|
||||
|
||||
def _load_utilities_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
|
||||
package_name = "dist_divider_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -0,0 +1,80 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_branch_module():
|
||||
package_name = "dist_branch_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "nodes/utilities.py", "nodes.utilities")
|
||||
return _load_module(package_name, "nodes/branch.py", "nodes.branch")
|
||||
|
||||
|
||||
branch_module = _load_branch_module()
|
||||
|
||||
|
||||
class DistributedBranchTests(unittest.TestCase):
|
||||
def test_node_declares_ten_output_slots(self):
|
||||
self.assertEqual(len(branch_module.DistributedBranch.RETURN_TYPES), 10)
|
||||
self.assertEqual(len(branch_module.DistributedBranch.RETURN_NAMES), 10)
|
||||
|
||||
def test_no_distribution_outputs_all_branches(self):
|
||||
node = branch_module.DistributedBranch()
|
||||
token = {"k": "v"}
|
||||
|
||||
outputs = node.branch(token, num_branches=2, assigned_branch=-1)
|
||||
|
||||
self.assertEqual(len(outputs), 10)
|
||||
self.assertTrue(all(output is token for output in outputs))
|
||||
|
||||
def test_assigned_branch_still_passes_input(self):
|
||||
node = branch_module.DistributedBranch()
|
||||
token = "payload"
|
||||
|
||||
outputs = node.branch(token, num_branches=3, is_worker=True, worker_id="worker-a", assigned_branch=1)
|
||||
|
||||
self.assertEqual(outputs[0], "payload")
|
||||
self.assertEqual(outputs[1], "payload")
|
||||
self.assertEqual(outputs[9], "payload")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,282 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import importlib.util
|
||||
import io
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 0.2
|
||||
config_module.is_master_delegate_only = lambda: False
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.HEARTBEAT_INTERVAL = 1.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
|
||||
def ensure_contiguous(tensor):
|
||||
return tensor.contiguous()
|
||||
|
||||
def tensor_to_pil(image_batch, index):
|
||||
tensor = image_batch[index].detach().cpu().clamp(0, 1)
|
||||
array = (tensor.numpy() * 255.0).astype("uint8")
|
||||
return Image.fromarray(array)
|
||||
|
||||
def encode_tensor_png_data_url(image_batch, batch_index=0):
|
||||
image = tensor_to_pil(image_batch, batch_index)
|
||||
byte_io = io.BytesIO()
|
||||
image.save(byte_io, format="PNG", compress_level=0)
|
||||
encoded = base64.b64encode(byte_io.getvalue()).decode("utf-8")
|
||||
return f"data:image/png;base64,{encoded}"
|
||||
|
||||
image_module.ensure_contiguous = ensure_contiguous
|
||||
image_module.tensor_to_pil = tensor_to_pil
|
||||
image_module.encode_tensor_png_data_url = encode_tensor_png_data_url
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def get_client_session():
|
||||
raise RuntimeError("test should patch get_client_session")
|
||||
|
||||
network_module.get_client_session = get_client_session
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_branch_collector_module():
|
||||
package_name = "dist_join_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "nodes/utilities.py", "nodes.utilities")
|
||||
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
|
||||
_load_module(package_name, "nodes/runtime_helpers.py", "nodes.runtime_helpers")
|
||||
_load_module(package_name, "nodes/queue_wait.py", "nodes.queue_wait")
|
||||
_load_module(package_name, "utils/worker_ids.py", "utils.worker_ids")
|
||||
return _load_module(package_name, "nodes/branch_collector.py", "nodes.branch_collector")
|
||||
|
||||
|
||||
branch_collector_module = _load_branch_collector_module()
|
||||
|
||||
|
||||
def _context(**overrides):
|
||||
defaults = {
|
||||
"multi_job_id": "",
|
||||
"is_worker": False,
|
||||
"master_url": "",
|
||||
"enabled_worker_ids": "[]",
|
||||
"worker_id": "",
|
||||
"assigned_branch": -1,
|
||||
"delegate_only": False,
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return branch_collector_module.BranchRunContext(**defaults)
|
||||
|
||||
|
||||
class _FakeResponseCtx:
|
||||
def __init__(self, recorder, payload, url):
|
||||
self._recorder = recorder
|
||||
self._payload = payload
|
||||
self._url = url
|
||||
|
||||
async def __aenter__(self):
|
||||
self._recorder.append({"url": self._url, "payload": self._payload})
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def post(self, url, json=None, timeout=None):
|
||||
return _FakeResponseCtx(self.calls, json, url)
|
||||
|
||||
|
||||
class _FakePromptServer:
|
||||
def __init__(self):
|
||||
self.distributed_pending_jobs = {}
|
||||
self.distributed_jobs_lock = asyncio.Lock()
|
||||
|
||||
|
||||
class DistributedBranchCollectorTests(unittest.IsolatedAsyncioTestCase):
|
||||
def setUp(self):
|
||||
self.node = branch_collector_module.DistributedBranchCollector()
|
||||
|
||||
async def test_worker_mode_sends_result_with_assigned_branch_index(self):
|
||||
fake_session = _FakeSession()
|
||||
|
||||
async def _get_session():
|
||||
return fake_session
|
||||
|
||||
branch_collector_module.get_client_session = _get_session
|
||||
|
||||
tensor = torch.ones((1, 4, 4, 3), dtype=torch.float32)
|
||||
outputs = await self.node.execute(
|
||||
[None, tensor, None, None, None, None, None, None, None, None],
|
||||
num_branches=2,
|
||||
context=_context(
|
||||
multi_job_id="join-job",
|
||||
is_worker=True,
|
||||
master_url="http://master.local:8188",
|
||||
worker_id="worker-a",
|
||||
assigned_branch=1,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertEqual(len(fake_session.calls), 1)
|
||||
payload = fake_session.calls[0]["payload"]
|
||||
self.assertEqual(payload["batch_idx"], 1)
|
||||
self.assertTrue(payload["is_last"])
|
||||
self.assertEqual(payload["worker_id"], "worker-a")
|
||||
self.assertTrue(torch.equal(outputs[1], tensor))
|
||||
|
||||
async def test_master_mode_collects_worker_results_into_branch_slots(self):
|
||||
fake_server = _FakePromptServer()
|
||||
branch_collector_module._get_prompt_server_instance = lambda: fake_server
|
||||
|
||||
queue = asyncio.Queue()
|
||||
worker_1 = torch.full((1, 4, 4, 3), 0.5, dtype=torch.float32)
|
||||
worker_2 = torch.full((1, 4, 4, 3), 0.8, dtype=torch.float32)
|
||||
await queue.put({"tensor": worker_1, "worker_id": "worker-a", "image_index": 1, "is_last": True})
|
||||
await queue.put({"tensor": worker_2, "worker_id": "worker-b", "image_index": 2, "is_last": True})
|
||||
fake_server.distributed_pending_jobs["join-job"] = queue
|
||||
|
||||
master_tensor = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
|
||||
outputs = await self.node.execute(
|
||||
[master_tensor, None, None, None, None, None, None, None, None, None],
|
||||
num_branches=3,
|
||||
context=_context(
|
||||
multi_job_id="join-job",
|
||||
enabled_worker_ids=json.dumps(["worker-a", "worker-b"]),
|
||||
assigned_branch=0,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(outputs[0], master_tensor))
|
||||
self.assertTrue(torch.equal(outputs[1], worker_1))
|
||||
self.assertTrue(torch.equal(outputs[2], worker_2))
|
||||
|
||||
async def test_unassigned_slots_remain_none(self):
|
||||
fake_server = _FakePromptServer()
|
||||
branch_collector_module._get_prompt_server_instance = lambda: fake_server
|
||||
|
||||
queue = asyncio.Queue()
|
||||
worker_1 = torch.full((1, 4, 4, 3), 0.5, dtype=torch.float32)
|
||||
await queue.put({"tensor": worker_1, "worker_id": "worker-a", "image_index": 1, "is_last": True})
|
||||
fake_server.distributed_pending_jobs["join-job"] = queue
|
||||
|
||||
master_tensor = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
|
||||
outputs = await self.node.execute(
|
||||
[master_tensor, None, None, None, None, None, None, None, None, None],
|
||||
num_branches=2,
|
||||
context=_context(
|
||||
multi_job_id="join-job",
|
||||
enabled_worker_ids=json.dumps(["worker-a"]),
|
||||
assigned_branch=0,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertIsNone(outputs[2])
|
||||
self.assertIsNone(outputs[9])
|
||||
|
||||
async def test_missing_worker_results_use_safe_fallback_for_expected_slots(self):
|
||||
fake_server = _FakePromptServer()
|
||||
branch_collector_module._get_prompt_server_instance = lambda: fake_server
|
||||
|
||||
queue = asyncio.Queue()
|
||||
# No worker payloads enqueued; master should timeout and fill expected worker slots.
|
||||
fake_server.distributed_pending_jobs["join-job"] = queue
|
||||
|
||||
master_tensor = torch.full((1, 4, 4, 3), 0.7, dtype=torch.float32)
|
||||
outputs = await self.node.execute(
|
||||
[master_tensor, None, None, None, None, None, None, None, None, None],
|
||||
num_branches=3,
|
||||
context=_context(
|
||||
multi_job_id="join-job",
|
||||
enabled_worker_ids=json.dumps(["worker-a", "worker-b"]),
|
||||
assigned_branch=0,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(outputs[0], master_tensor))
|
||||
self.assertIsInstance(outputs[1], torch.Tensor)
|
||||
self.assertIsInstance(outputs[2], torch.Tensor)
|
||||
self.assertEqual(float(outputs[1].abs().sum().item()), 0.0)
|
||||
self.assertEqual(float(outputs[2].abs().sum().item()), 0.0)
|
||||
|
||||
async def test_assigned_slot_can_be_filled_from_single_mismatched_local_input(self):
|
||||
fake_server = _FakePromptServer()
|
||||
branch_collector_module._get_prompt_server_instance = lambda: fake_server
|
||||
|
||||
# No worker results expected for this check.
|
||||
fake_server.distributed_pending_jobs["join-job"] = asyncio.Queue()
|
||||
|
||||
local_tensor = torch.full((1, 4, 4, 3), 0.33, dtype=torch.float32)
|
||||
outputs = await self.node.execute(
|
||||
# Local value arrives on slot 1 even though participant is assigned slot 0.
|
||||
[None, local_tensor, None, None, None, None, None, None, None, None],
|
||||
num_branches=3,
|
||||
context=_context(
|
||||
multi_job_id="join-job",
|
||||
enabled_worker_ids=json.dumps([]),
|
||||
assigned_branch=0,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertTrue(torch.equal(outputs[0], local_tensor))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,228 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 0.2
|
||||
config_module.load_config = lambda: {"workers": []}
|
||||
config_module.is_master_delegate_only = lambda: False
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.HEARTBEAT_INTERVAL = 1.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
|
||||
def ensure_contiguous(tensor):
|
||||
return tensor.contiguous()
|
||||
|
||||
def tensor_to_pil(image_batch, index):
|
||||
tensor = image_batch[index].detach().cpu().clamp(0, 1)
|
||||
array = (tensor.numpy() * 255.0).astype("uint8")
|
||||
return Image.fromarray(array)
|
||||
|
||||
def pil_to_tensor(image):
|
||||
array = np.array(image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array).unsqueeze(0)
|
||||
|
||||
image_module.ensure_contiguous = ensure_contiguous
|
||||
image_module.tensor_to_pil = tensor_to_pil
|
||||
image_module.pil_to_tensor = pil_to_tensor
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.build_worker_url = lambda _worker: "http://worker.local:8188"
|
||||
network_module.get_client_session = lambda: None
|
||||
|
||||
async def probe_worker(_url, timeout=2.0):
|
||||
_ = timeout
|
||||
return None
|
||||
|
||||
network_module.probe_worker = probe_worker
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
audio_payload_module = types.ModuleType(f"{package_name}.utils.audio_payload")
|
||||
audio_payload_module.encode_audio_payload = lambda payload: payload
|
||||
sys.modules[f"{package_name}.utils.audio_payload"] = audio_payload_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
if "aiohttp" not in sys.modules:
|
||||
try:
|
||||
import aiohttp as _aiohttp # noqa: F401
|
||||
except Exception:
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
|
||||
class _ClientTimeout:
|
||||
def __init__(self, total=None):
|
||||
self.total = total
|
||||
|
||||
aiohttp_module.ClientTimeout = _ClientTimeout
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
class _FakePromptServer:
|
||||
def __init__(self):
|
||||
self.distributed_pending_jobs = {}
|
||||
self.distributed_jobs_lock = asyncio.Lock()
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=_FakePromptServer())
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
model_mgmt = types.ModuleType("comfy.model_management")
|
||||
|
||||
class _InterruptProcessingException(Exception):
|
||||
pass
|
||||
|
||||
model_mgmt.throw_exception_if_processing_interrupted = lambda: None
|
||||
model_mgmt.InterruptProcessingException = _InterruptProcessingException
|
||||
comfy_module.model_management = model_mgmt
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.model_management"] = model_mgmt
|
||||
|
||||
comfy_utils_module = types.ModuleType("comfy.utils")
|
||||
|
||||
class _ProgressBar:
|
||||
def __init__(self, _total):
|
||||
self.total = _total
|
||||
|
||||
def update(self, _step):
|
||||
return None
|
||||
|
||||
comfy_utils_module.ProgressBar = _ProgressBar
|
||||
sys.modules["comfy.utils"] = comfy_utils_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_collector_module():
|
||||
package_name = "dist_collector_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
|
||||
_load_module(package_name, "utils/worker_ids.py", "utils.worker_ids")
|
||||
return _load_module(package_name, "nodes/collector.py", "nodes.collector")
|
||||
|
||||
|
||||
collector_module = _load_collector_module()
|
||||
|
||||
|
||||
class DistributedCollectorTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.node = collector_module.DistributedCollectorNode()
|
||||
|
||||
def test_store_worker_result_tracks_by_worker_and_index(self):
|
||||
worker_images = {}
|
||||
stored = self.node._store_worker_result(
|
||||
worker_images,
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"image_index": 2,
|
||||
"tensor": torch.full((1, 1, 1, 1), 0.25),
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(stored, 1)
|
||||
self.assertIn("worker-a", worker_images)
|
||||
self.assertIn(2, worker_images["worker-a"])
|
||||
|
||||
def test_combine_audio_honors_master_then_worker_order(self):
|
||||
master_audio = {"waveform": torch.ones((1, 2, 1)), "sample_rate": 48000}
|
||||
worker_audio = {
|
||||
"worker-b": {"waveform": torch.full((1, 2, 1), 3.0), "sample_rate": 48000},
|
||||
"worker-a": {"waveform": torch.full((1, 2, 2), 2.0), "sample_rate": 48000},
|
||||
}
|
||||
empty_audio = {"waveform": torch.zeros((1, 2, 1)), "sample_rate": 44100}
|
||||
|
||||
combined = self.node._combine_audio(
|
||||
master_audio=master_audio,
|
||||
worker_audio=worker_audio,
|
||||
empty_audio=empty_audio,
|
||||
worker_order=["worker-a", "worker-b"],
|
||||
)
|
||||
|
||||
expected = torch.cat(
|
||||
[
|
||||
master_audio["waveform"],
|
||||
worker_audio["worker-a"]["waveform"],
|
||||
worker_audio["worker-b"]["waveform"],
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
self.assertEqual(combined["sample_rate"], 48000)
|
||||
self.assertTrue(torch.equal(combined["waveform"], expected))
|
||||
|
||||
def test_reorder_and_combine_tensors_uses_enabled_worker_priority(self):
|
||||
worker_images = {
|
||||
"worker-b": {
|
||||
1: torch.full((1, 1, 1, 1), 4.0),
|
||||
0: torch.full((1, 1, 1, 1), 3.0),
|
||||
},
|
||||
"worker-a": {
|
||||
0: torch.full((1, 1, 1, 1), 2.0),
|
||||
},
|
||||
}
|
||||
master_images = torch.full((1, 1, 1, 1), 1.0)
|
||||
|
||||
combined = self.node._reorder_and_combine_tensors(
|
||||
worker_images=worker_images,
|
||||
worker_order=["worker-a", "worker-b"],
|
||||
master_batch_size=1,
|
||||
images_on_cpu=master_images,
|
||||
delegate_mode=False,
|
||||
fallback_images=None,
|
||||
)
|
||||
|
||||
self.assertEqual(combined.shape, (4, 1, 1, 1))
|
||||
self.assertEqual(combined[0, 0, 0, 0].item(), 1.0)
|
||||
self.assertEqual(combined[1, 0, 0, 0].item(), 2.0)
|
||||
self.assertEqual(combined[2, 0, 0, 0].item(), 3.0)
|
||||
self.assertEqual(combined[3, 0, 0, 0].item(), 4.0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,248 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
|
||||
modes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
worker_ids_module = types.ModuleType(f"{package_name}.utils.worker_ids")
|
||||
worker_ids_module.coerce_enabled_worker_ids = lambda value: (
|
||||
[str(v) for v in value]
|
||||
if isinstance(value, list)
|
||||
else []
|
||||
)
|
||||
sys.modules[f"{package_name}.utils.worker_ids"] = worker_ids_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda coro, timeout=None: (_ for _ in ()).throw(
|
||||
RuntimeError(f"run_async_in_server_loop should not be used in these tests (timeout={timeout})")
|
||||
)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
|
||||
job_store_module.ensure_tile_jobs_initialized = lambda: None
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
tile_ops_module = types.ModuleType(f"{package_name}.upscale.tile_ops")
|
||||
|
||||
class _TileOpsMixin:
|
||||
def round_to_multiple(self, value):
|
||||
return int(value)
|
||||
|
||||
def calculate_tiles(self, width, height, tile_width, tile_height, force_uniform_tiles):
|
||||
_ = force_uniform_tiles
|
||||
tiles = []
|
||||
for y in range(0, int(height), max(1, int(tile_height))):
|
||||
for x in range(0, int(width), max(1, int(tile_width))):
|
||||
tiles.append((x, y))
|
||||
return tiles or [(0, 0)]
|
||||
|
||||
tile_ops_module.TileOpsMixin = _TileOpsMixin
|
||||
sys.modules[f"{package_name}.upscale.tile_ops"] = tile_ops_module
|
||||
|
||||
result_collector_module = types.ModuleType(f"{package_name}.upscale.result_collector")
|
||||
result_collector_module.ResultCollectorMixin = type("ResultCollectorMixin", (), {})
|
||||
sys.modules[f"{package_name}.upscale.result_collector"] = result_collector_module
|
||||
|
||||
worker_comms_module = types.ModuleType(f"{package_name}.upscale.worker_comms")
|
||||
worker_comms_module.WorkerCommsMixin = type("WorkerCommsMixin", (), {})
|
||||
sys.modules[f"{package_name}.upscale.worker_comms"] = worker_comms_module
|
||||
|
||||
job_state_module = types.ModuleType(f"{package_name}.upscale.job_state")
|
||||
job_state_module.JobStateMixin = type("JobStateMixin", (), {})
|
||||
sys.modules[f"{package_name}.upscale.job_state"] = job_state_module
|
||||
|
||||
single_gpu_module = types.ModuleType(f"{package_name}.upscale.modes.single_gpu")
|
||||
|
||||
class _SingleGpuModeMixin:
|
||||
def process_single_gpu(self, *_args, **_kwargs):
|
||||
return ("single_gpu_result",)
|
||||
|
||||
single_gpu_module.SingleGpuModeMixin = _SingleGpuModeMixin
|
||||
sys.modules[f"{package_name}.upscale.modes.single_gpu"] = single_gpu_module
|
||||
|
||||
static_mode_module = types.ModuleType(f"{package_name}.upscale.modes.static")
|
||||
|
||||
class _StaticModeMixin:
|
||||
def _process_worker_static_sync(self, *_args, **_kwargs):
|
||||
return ("worker_static_result",)
|
||||
|
||||
def _process_master_static_sync(self, *_args, **_kwargs):
|
||||
return ("master_static_result",)
|
||||
|
||||
static_mode_module.StaticModeMixin = _StaticModeMixin
|
||||
sys.modules[f"{package_name}.upscale.modes.static"] = static_mode_module
|
||||
|
||||
dynamic_mode_module = types.ModuleType(f"{package_name}.upscale.modes.dynamic")
|
||||
|
||||
class _DynamicModeMixin:
|
||||
def process_worker_dynamic(self, *_args, **_kwargs):
|
||||
return ("worker_dynamic_result",)
|
||||
|
||||
def process_master_dynamic(self, *_args, **_kwargs):
|
||||
return ("master_dynamic_result",)
|
||||
|
||||
dynamic_mode_module.DynamicModeMixin = _DynamicModeMixin
|
||||
sys.modules[f"{package_name}.upscale.modes.dynamic"] = dynamic_mode_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
samplers_module = types.ModuleType("comfy.samplers")
|
||||
|
||||
class _KSampler:
|
||||
SAMPLERS = ("euler",)
|
||||
SCHEDULERS = ("normal",)
|
||||
|
||||
samplers_module.KSampler = _KSampler
|
||||
comfy_module.samplers = samplers_module
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.samplers"] = samplers_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_distributed_upscale_module():
|
||||
package_name = "dist_distributed_upscale_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
|
||||
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
|
||||
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
|
||||
return _load_module(package_name, "nodes/distributed_upscale.py", "nodes.distributed_upscale")
|
||||
|
||||
|
||||
upscale_module = _load_distributed_upscale_module()
|
||||
|
||||
|
||||
class DistributedUpscaleTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.node = upscale_module.UltimateSDUpscaleDistributed()
|
||||
self.base_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32)
|
||||
self.common_kwargs = {
|
||||
"model": object(),
|
||||
"positive": object(),
|
||||
"negative": object(),
|
||||
"vae": object(),
|
||||
"seed": 1,
|
||||
"steps": 10,
|
||||
"cfg": 7.5,
|
||||
"sampler_name": "euler",
|
||||
"scheduler": "normal",
|
||||
"denoise": 0.5,
|
||||
"tile_width": 32,
|
||||
"tile_height": 32,
|
||||
"padding": 16,
|
||||
"mask_blur": 4,
|
||||
"force_uniform_tiles": True,
|
||||
"tiled_decode": False,
|
||||
}
|
||||
|
||||
def test_parse_enabled_worker_ids_supports_json_and_lists(self):
|
||||
self.assertEqual(
|
||||
upscale_module._parse_enabled_worker_ids(["worker-a", 2]),
|
||||
["worker-a", "2"],
|
||||
)
|
||||
self.assertEqual(
|
||||
upscale_module._parse_enabled_worker_ids('["worker-a","worker-b"]'),
|
||||
["worker-a", "worker-b"],
|
||||
)
|
||||
self.assertEqual(upscale_module._parse_enabled_worker_ids("invalid-json"), [])
|
||||
self.assertEqual(upscale_module._parse_enabled_worker_ids(None), [])
|
||||
|
||||
def test_determine_processing_mode(self):
|
||||
self.assertEqual(self.node._determine_processing_mode(batch_size=1, num_workers=0, dynamic_threshold=8), "single_gpu")
|
||||
self.assertEqual(self.node._determine_processing_mode(batch_size=9, num_workers=2, dynamic_threshold=8), "static")
|
||||
|
||||
def test_run_raises_for_non_4n_plus_1_master_batches(self):
|
||||
bad_batch = torch.zeros((2, 64, 64, 3), dtype=torch.float32)
|
||||
with self.assertRaises(ValueError):
|
||||
self.node.run(
|
||||
upscaled_image=bad_batch,
|
||||
multi_job_id="job-1",
|
||||
is_worker=False,
|
||||
**self.common_kwargs,
|
||||
)
|
||||
|
||||
def test_run_dispatches_single_gpu_when_not_distributed(self):
|
||||
self.node.process_single_gpu = lambda *_args, **_kwargs: ("single",)
|
||||
result = self.node.run(
|
||||
upscaled_image=self.base_image,
|
||||
multi_job_id="",
|
||||
is_worker=False,
|
||||
**self.common_kwargs,
|
||||
)
|
||||
self.assertEqual(result, ("single",))
|
||||
|
||||
def test_run_dispatches_worker_or_master_paths(self):
|
||||
self.node.process_worker = lambda *_args, **_kwargs: ("worker",)
|
||||
self.node.process_master = lambda *_args, **_kwargs: ("master",)
|
||||
|
||||
worker_result = self.node.run(
|
||||
upscaled_image=self.base_image,
|
||||
multi_job_id="job-2",
|
||||
is_worker=True,
|
||||
master_url="http://master.local:8188",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
worker_id="worker-a",
|
||||
tile_indices="",
|
||||
dynamic_threshold=8,
|
||||
**self.common_kwargs,
|
||||
)
|
||||
master_result = self.node.run(
|
||||
upscaled_image=self.base_image,
|
||||
multi_job_id="job-2",
|
||||
is_worker=False,
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
worker_id="",
|
||||
tile_indices="",
|
||||
dynamic_threshold=8,
|
||||
**self.common_kwargs,
|
||||
)
|
||||
|
||||
self.assertEqual(worker_result, ("worker",))
|
||||
self.assertEqual(master_result, ("master",))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -13,7 +13,7 @@ class DistributedValueTests(unittest.TestCase):
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
|
||||
pkg_name = "dv_test_pkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -28,6 +28,18 @@ class DistributedValueTests(unittest.TestCase):
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[pkg_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{pkg_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{pkg_name}.nodes"] = nodes_pkg
|
||||
|
||||
hidden_inputs_mod = types.ModuleType(f"{pkg_name}.nodes.hidden_inputs")
|
||||
hidden_inputs_mod.build_worker_identity_hidden_inputs = lambda: {
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"enabled_worker_ids": ("STRING", {"default": "[]"}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
}
|
||||
sys.modules[f"{pkg_name}.nodes.hidden_inputs"] = hidden_inputs_mod
|
||||
|
||||
utils_pkg = types.ModuleType(f"{pkg_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{pkg_name}.utils"] = utils_pkg
|
||||
@@ -37,6 +49,16 @@ class DistributedValueTests(unittest.TestCase):
|
||||
logging_mod.log = lambda *_a, **_k: None
|
||||
sys.modules[f"{pkg_name}.utils.logging"] = logging_mod
|
||||
|
||||
worker_ids_path = Path(__file__).resolve().parents[3] / "utils" / "worker_ids.py"
|
||||
worker_ids_spec = importlib.util.spec_from_file_location(
|
||||
f"{pkg_name}.utils.worker_ids",
|
||||
worker_ids_path,
|
||||
)
|
||||
worker_ids_mod = importlib.util.module_from_spec(worker_ids_spec)
|
||||
assert worker_ids_spec is not None and worker_ids_spec.loader is not None
|
||||
worker_ids_spec.loader.exec_module(worker_ids_mod)
|
||||
sys.modules[f"{pkg_name}.utils.worker_ids"] = worker_ids_mod
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{pkg_name}.nodes.utilities", module_path
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests."""
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
@@ -5,7 +6,7 @@ import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class _PromptQueue:
|
||||
class _FakePromptQueue:
|
||||
def __init__(self):
|
||||
self.items = []
|
||||
|
||||
@@ -13,10 +14,16 @@ class _PromptQueue:
|
||||
self.items.append(item)
|
||||
|
||||
|
||||
def _load_async_helpers_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "async_helpers.py"
|
||||
package_name = "dist_async_helpers_testpkg"
|
||||
class _FakePromptServer:
|
||||
def __init__(self):
|
||||
self.number = 0
|
||||
self.prompt_queue = _FakePromptQueue()
|
||||
|
||||
def trigger_on_prompt(self, payload):
|
||||
return payload
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
@@ -29,62 +36,68 @@ def _load_async_helpers_module():
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
execution_module = types.ModuleType("execution")
|
||||
|
||||
async def _validate_prompt(prompt_id, prompt, partial_execution_targets):
|
||||
return (True, None, ["9"], {})
|
||||
|
||||
execution_module.validate_prompt = _validate_prompt
|
||||
execution_module.SENSITIVE_EXTRA_DATA_KEYS = []
|
||||
sys.modules["execution"] = execution_module
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
trigger_on_prompt=lambda payload: payload,
|
||||
number=12,
|
||||
prompt_queue=_PromptQueue(),
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.get_server_loop = lambda: None
|
||||
network_module.get_server_loop = asyncio.get_event_loop
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.utils.async_helpers", module_path)
|
||||
fake_prompt_server = _FakePromptServer()
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=fake_prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
execution_module = types.ModuleType("execution")
|
||||
|
||||
async def _validate_prompt(_prompt_id, _prompt, _partial_targets):
|
||||
return (True, None, ["1"], {})
|
||||
|
||||
execution_module.validate_prompt = _validate_prompt
|
||||
execution_module.SENSITIVE_EXTRA_DATA_KEYS = ()
|
||||
sys.modules["execution"] = execution_module
|
||||
|
||||
return fake_prompt_server
|
||||
|
||||
|
||||
def _load_async_helpers_module():
|
||||
package_name = "dist_async_helpers_testpkg"
|
||||
fake_prompt_server = _bootstrap_package(package_name)
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils/async_helpers.py"
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.utils.async_helpers",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module, prompt_server
|
||||
return module, fake_prompt_server
|
||||
|
||||
|
||||
async_helpers, prompt_server = _load_async_helpers_module()
|
||||
class AsyncHelpersQueuePromptPayloadTests(unittest.TestCase):
|
||||
def test_queue_prompt_payload_includes_create_time_metadata(self):
|
||||
async_helpers_module, fake_prompt_server = _load_async_helpers_module()
|
||||
|
||||
|
||||
class QueuePromptPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_queue_prompt_payload_includes_create_time_and_client_metadata(self):
|
||||
result = await async_helpers.queue_prompt_payload(
|
||||
{"1": {"class_type": "Node"}},
|
||||
workflow_meta={"id": "workflow-1"},
|
||||
client_id="client-1",
|
||||
include_queue_metadata=True,
|
||||
prompt_id = asyncio.run(
|
||||
async_helpers_module.queue_prompt_payload(
|
||||
{"1": {"class_type": "KSampler", "inputs": {}}},
|
||||
workflow_meta={"id": "workflow-1"},
|
||||
client_id="client-1",
|
||||
)
|
||||
)
|
||||
|
||||
self.assertIsInstance(result["prompt_id"], str)
|
||||
self.assertTrue(result["prompt_id"])
|
||||
self.assertEqual(result["number"], 12)
|
||||
self.assertEqual(result["node_errors"], {})
|
||||
self.assertIsInstance(prompt_id, str)
|
||||
self.assertTrue(fake_prompt_server.prompt_queue.items)
|
||||
|
||||
queued_item = fake_prompt_server.prompt_queue.items[-1]
|
||||
self.assertEqual(len(queued_item), 6)
|
||||
self.assertEqual(queued_item[1], prompt_id)
|
||||
|
||||
self.assertEqual(prompt_server.number, 13)
|
||||
self.assertEqual(len(prompt_server.prompt_queue.items), 1)
|
||||
queued_item = prompt_server.prompt_queue.items[0]
|
||||
self.assertEqual(queued_item[0], 12)
|
||||
extra_data = queued_item[3]
|
||||
self.assertEqual(extra_data["client_id"], "client-1")
|
||||
self.assertEqual(extra_data["extra_pnginfo"]["workflow"]["id"], "workflow-1")
|
||||
self.assertIn("create_time", extra_data)
|
||||
self.assertIsInstance(extra_data["create_time"], int)
|
||||
self.assertGreater(extra_data["create_time"], 0)
|
||||
self.assertEqual(extra_data["extra_pnginfo"]["workflow"], {"id": "workflow-1"})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -10,7 +10,7 @@ from unittest.mock import patch
|
||||
|
||||
|
||||
def _load_config_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "config.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "config.py"
|
||||
package_name = "dist_cfg_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -21,13 +21,13 @@ def _load_config_module():
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.logging")
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.logging"] = logging_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.constants")
|
||||
constants_module.GPU_CONFIG_FILE = "gpu_config.json"
|
||||
constants_module.HEARTBEAT_TIMEOUT = 30
|
||||
constants_module.ORCHESTRATION_WORKER_PROBE_CONCURRENCY = 8
|
||||
constants_module.ORCHESTRATION_WORKER_PREP_CONCURRENCY = 4
|
||||
constants_module.ORCHESTRATION_MEDIA_SYNC_CONCURRENCY = 2
|
||||
constants_module.ORCHESTRATION_MEDIA_SYNC_TIMEOUT = 120
|
||||
sys.modules[f"{package_name}.constants"] = constants_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.config", module_path)
|
||||
@@ -213,7 +213,7 @@ class SaveConfigTests(unittest.TestCase):
|
||||
data = config.get_default_config()
|
||||
config.save_config(data)
|
||||
# Cache is now None; load_config should re-read
|
||||
self.assertIsNone(config._config_cache)
|
||||
self.assertIsNone(config._config_state().cache)
|
||||
|
||||
def test_written_file_is_valid_json(self):
|
||||
data = config.get_default_config()
|
||||
@@ -7,7 +7,7 @@ from unittest.mock import patch
|
||||
|
||||
|
||||
def _load_detection_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "workers" / "detection.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "workers" / "detection.py"
|
||||
package_name = "dist_det_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -32,6 +32,7 @@ def _load_detection_module():
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.normalize_host = lambda value: value
|
||||
network_module.build_worker_url = lambda worker, endpoint="": f"http://{worker.get('host', 'localhost')}{endpoint}"
|
||||
|
||||
async def _fake_session():
|
||||
raise RuntimeError("network calls not used in these tests")
|
||||
@@ -156,7 +157,7 @@ class IsLocalWorkerTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertTrue(result)
|
||||
|
||||
async def test_true_for_0_0_0_0(self):
|
||||
result = await detection.is_local_worker({"host": "0.0.0.0", "port": 8188})
|
||||
result = await detection.is_local_worker({"host": "0.0.0.0", "port": 8188}) # nosec B104 - explicit wildcard host test case
|
||||
self.assertTrue(result)
|
||||
|
||||
async def test_true_when_type_is_local(self):
|
||||
@@ -8,7 +8,7 @@ from unittest.mock import patch
|
||||
|
||||
|
||||
def _load_dispatch_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "api" / "orchestration" / "dispatch.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "dispatch.py"
|
||||
|
||||
package_name = "dist_dispatch_testpkg"
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
@@ -249,6 +249,50 @@ class DispatchSelectionTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
self.assertIsNone(selected)
|
||||
|
||||
async def test_rank_workers_by_load_returns_sorted_by_queue_depth(self):
|
||||
workers = [
|
||||
{"id": "w1", "name": "Worker 1"},
|
||||
{"id": "w2", "name": "Worker 2"},
|
||||
{"id": "w3", "name": "Worker 3"},
|
||||
]
|
||||
queue_map = {"w1": 4, "w2": 0, "w3": 2}
|
||||
|
||||
async def fake_probe(worker_url, timeout=3.0):
|
||||
worker_id = worker_url.rsplit("/", 1)[-1]
|
||||
return {"exec_info": {"queue_remaining": queue_map[worker_id]}}
|
||||
|
||||
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
|
||||
dispatch,
|
||||
"probe_worker",
|
||||
side_effect=fake_probe,
|
||||
):
|
||||
ranked = await dispatch.rank_workers_by_load(workers, probe_concurrency=3)
|
||||
|
||||
self.assertEqual([worker["id"] for worker in ranked], ["w2", "w3", "w1"])
|
||||
|
||||
async def test_rank_workers_by_load_handles_unreachable_workers(self):
|
||||
workers = [
|
||||
{"id": "w1", "name": "Worker 1"},
|
||||
{"id": "w2", "name": "Worker 2"},
|
||||
{"id": "w3", "name": "Worker 3"},
|
||||
]
|
||||
|
||||
async def fake_probe(worker_url, timeout=3.0):
|
||||
worker_id = worker_url.rsplit("/", 1)[-1]
|
||||
if worker_id == "w2":
|
||||
return None
|
||||
queue_map = {"w1": 1, "w3": 0}
|
||||
return {"exec_info": {"queue_remaining": queue_map[worker_id]}}
|
||||
|
||||
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
|
||||
dispatch,
|
||||
"probe_worker",
|
||||
side_effect=fake_probe,
|
||||
):
|
||||
ranked = await dispatch.rank_workers_by_load(workers, probe_concurrency=3)
|
||||
|
||||
self.assertEqual([worker["id"] for worker in ranked], ["w3", "w1", "w2"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,71 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _load_distributed_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "distributed.py"
|
||||
package_name = "dist_distributed_entry_testpkg"
|
||||
calls = {
|
||||
"build_node_mappings": 0,
|
||||
"initialize_runtime": 0,
|
||||
"initialize_runtime_args": [],
|
||||
}
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
bootstrap_pkg = types.ModuleType(f"{package_name}.bootstrap")
|
||||
bootstrap_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.bootstrap"] = bootstrap_pkg
|
||||
|
||||
entrypoint_module = types.ModuleType(f"{package_name}.bootstrap.entrypoint")
|
||||
|
||||
def _build_node_mappings():
|
||||
calls["build_node_mappings"] += 1
|
||||
return ({"MockNode": object}, {"MockNode": "Mock Node"})
|
||||
|
||||
def _initialize_runtime(prompt_server=None):
|
||||
calls["initialize_runtime"] += 1
|
||||
calls["initialize_runtime_args"].append(prompt_server)
|
||||
|
||||
entrypoint_module.build_node_mappings = _build_node_mappings
|
||||
entrypoint_module.initialize_runtime = _initialize_runtime
|
||||
sys.modules[f"{package_name}.bootstrap.entrypoint"] = entrypoint_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.distributed", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module, calls
|
||||
|
||||
|
||||
class DistributedEntryTests(unittest.TestCase):
|
||||
def test_import_builds_mappings_without_runtime_bootstrap(self):
|
||||
module, calls = _load_distributed_module()
|
||||
|
||||
self.assertIn("MockNode", module.NODE_CLASS_MAPPINGS)
|
||||
self.assertEqual(module.NODE_DISPLAY_NAME_MAPPINGS["MockNode"], "Mock Node")
|
||||
self.assertEqual(calls["build_node_mappings"], 1)
|
||||
self.assertEqual(calls["initialize_runtime"], 0)
|
||||
|
||||
def test_initialize_runtime_delegates_to_bootstrap_entrypoint(self):
|
||||
module, calls = _load_distributed_module()
|
||||
sentinel = object()
|
||||
|
||||
module.initialize_runtime(sentinel)
|
||||
|
||||
self.assertEqual(calls["initialize_runtime"], 1)
|
||||
self.assertEqual(calls["initialize_runtime_args"], [sentinel])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,40 @@
|
||||
import importlib.util
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _load_exceptions_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "exceptions.py"
|
||||
spec = importlib.util.spec_from_file_location("dist_test_exceptions", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
exceptions = _load_exceptions_module()
|
||||
|
||||
|
||||
class ExceptionsModuleTests(unittest.TestCase):
|
||||
def test_worker_error_tracks_worker_id(self):
|
||||
err = exceptions.WorkerError("worker failed", worker_id="w-1")
|
||||
self.assertIsInstance(err, exceptions.DistributedError)
|
||||
self.assertEqual(str(err), "worker failed")
|
||||
self.assertEqual(err.worker_id, "w-1")
|
||||
|
||||
def test_process_error_tracks_pid_and_worker_id(self):
|
||||
err = exceptions.ProcessError("process failed", pid=1001, worker_id="worker-x")
|
||||
self.assertIsInstance(err, exceptions.DistributedError)
|
||||
self.assertEqual(err.pid, 1001)
|
||||
self.assertEqual(err.worker_id, "worker-x")
|
||||
|
||||
def test_specialized_errors_inherit_expected_base_types(self):
|
||||
self.assertTrue(issubclass(exceptions.WorkerTimeoutError, exceptions.WorkerError))
|
||||
self.assertTrue(issubclass(exceptions.WorkerNotAvailableError, exceptions.WorkerError))
|
||||
self.assertTrue(issubclass(exceptions.JobQueueError, exceptions.DistributedError))
|
||||
self.assertTrue(issubclass(exceptions.TileCollectionError, exceptions.DistributedError))
|
||||
self.assertTrue(issubclass(exceptions.TunnelError, exceptions.DistributedError))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,228 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
process_pkg = types.ModuleType(f"{package_name}.workers.process")
|
||||
process_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.workers.process"] = process_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_lifecycle_module():
|
||||
package_name = "dist_lifecycle_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
|
||||
state = {
|
||||
"config": {"settings": {"stop_workers_on_master_exit": True}, "managed_processes": {}},
|
||||
"saved_configs": [],
|
||||
"launch_calls": [],
|
||||
"terminated": [],
|
||||
"alive_pids": set(),
|
||||
}
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: state["config"]
|
||||
config_module.save_config = lambda cfg: state["saved_configs"].append(dict(cfg))
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.PROCESS_TERMINATION_TIMEOUT = 5.0
|
||||
constants_module.PROCESS_WAIT_TIMEOUT = 1.0
|
||||
constants_module.WORKER_CHECK_INTERVAL = 0.01
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
process_module = types.ModuleType(f"{package_name}.utils.process")
|
||||
|
||||
class _FakeProcess:
|
||||
def __init__(self, pid=1234, poll_result=None):
|
||||
self.pid = pid
|
||||
self._poll_result = poll_result
|
||||
self.returncode = poll_result
|
||||
|
||||
def poll(self):
|
||||
return self._poll_result
|
||||
|
||||
def _launch_process_with_timeout(command, timeout_seconds=10.0, **kwargs):
|
||||
state["launch_calls"].append((list(command), float(timeout_seconds), dict(kwargs)))
|
||||
return _FakeProcess(pid=4321, poll_result=None)
|
||||
|
||||
def _terminate_process(process, timeout=5):
|
||||
state["terminated"].append((process.pid, timeout))
|
||||
process._poll_result = 0
|
||||
process.returncode = 0
|
||||
|
||||
process_module.launch_process_with_timeout = _launch_process_with_timeout
|
||||
process_module.terminate_process = _terminate_process
|
||||
process_module.is_process_alive = lambda pid: int(pid) in state["alive_pids"]
|
||||
process_module.get_python_executable = lambda: "/usr/bin/python3"
|
||||
process_module._FakeProcess = _FakeProcess
|
||||
sys.modules[f"{package_name}.utils.process"] = process_module
|
||||
|
||||
lifecycle_module = _load_module(package_name, "workers/process/lifecycle.py", "workers.process.lifecycle")
|
||||
return lifecycle_module, state, process_module
|
||||
|
||||
|
||||
lifecycle_module, test_state, process_stub_module = _load_lifecycle_module()
|
||||
|
||||
|
||||
class _FakeManager:
|
||||
def __init__(self, comfy_root):
|
||||
self._comfy_root = comfy_root
|
||||
self.processes = {}
|
||||
self.save_calls = 0
|
||||
|
||||
def find_comfy_root(self):
|
||||
return self._comfy_root
|
||||
|
||||
def build_launch_command(self, worker_config, _comfy_root):
|
||||
return ["python", "main.py", "--port", str(worker_config["port"])]
|
||||
|
||||
def save_processes(self):
|
||||
self.save_calls += 1
|
||||
|
||||
|
||||
class ProcessLifecycleTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
test_state["config"] = {"settings": {"stop_workers_on_master_exit": True}, "managed_processes": {}}
|
||||
test_state["saved_configs"].clear()
|
||||
test_state["launch_calls"].clear()
|
||||
test_state["terminated"].clear()
|
||||
test_state["alive_pids"].clear()
|
||||
|
||||
def test_launch_worker_uses_monitor_wrapper_when_enabled(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
manager = _FakeManager(tmpdir)
|
||||
lifecycle = lifecycle_module.ProcessLifecycle(manager)
|
||||
|
||||
pid = lifecycle.launch_worker(
|
||||
{
|
||||
"id": "1",
|
||||
"name": "Worker A",
|
||||
"port": 8189,
|
||||
"cuda_device": 2,
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(pid, 4321)
|
||||
self.assertEqual(manager.save_calls, 1)
|
||||
self.assertIn("1", manager.processes)
|
||||
self.assertTrue(manager.processes["1"]["is_monitor"])
|
||||
self.assertIn("logs/workers", manager.processes["1"]["log_file"])
|
||||
self.assertEqual(len(test_state["launch_calls"]), 1)
|
||||
launched_command, timeout_seconds, launch_kwargs = test_state["launch_calls"][0]
|
||||
self.assertEqual(timeout_seconds, 1.0)
|
||||
self.assertEqual(launched_command[0], "/usr/bin/python3")
|
||||
self.assertTrue(launched_command[1].endswith("workers/worker_monitor.py"))
|
||||
self.assertEqual(launch_kwargs["env"]["CUDA_VISIBLE_DEVICES"], "2")
|
||||
|
||||
def test_launch_worker_skips_monitor_when_disabled(self):
|
||||
test_state["config"] = {"settings": {"stop_workers_on_master_exit": False}, "managed_processes": {}}
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
manager = _FakeManager(tmpdir)
|
||||
lifecycle = lifecycle_module.ProcessLifecycle(manager)
|
||||
|
||||
lifecycle.launch_worker(
|
||||
{
|
||||
"id": "2",
|
||||
"name": "Worker-B",
|
||||
"port": 8190,
|
||||
"cuda_device": 0,
|
||||
}
|
||||
)
|
||||
|
||||
launched_command, _, _ = test_state["launch_calls"][0]
|
||||
self.assertEqual(launched_command[0], "python")
|
||||
self.assertFalse(manager.processes["2"]["is_monitor"])
|
||||
|
||||
def test_stop_worker_uses_fallback_terminate_when_tree_kill_fails(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
manager = _FakeManager(tmpdir)
|
||||
lifecycle = lifecycle_module.ProcessLifecycle(manager)
|
||||
running_process = process_stub_module._FakeProcess(pid=88, poll_result=None)
|
||||
manager.processes["7"] = {
|
||||
"pid": 88,
|
||||
"process": running_process,
|
||||
"started_at": 0.0,
|
||||
"config": {},
|
||||
}
|
||||
lifecycle._kill_process_tree = lambda _pid: False
|
||||
|
||||
ok, message = lifecycle.stop_worker("7")
|
||||
|
||||
self.assertTrue(ok)
|
||||
self.assertIn("fallback", message.lower())
|
||||
self.assertEqual(test_state["terminated"], [(88, 5.0)])
|
||||
self.assertNotIn("7", manager.processes)
|
||||
|
||||
def test_stop_worker_without_process_uses_tree_kill(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
manager = _FakeManager(tmpdir)
|
||||
lifecycle = lifecycle_module.ProcessLifecycle(manager)
|
||||
manager.processes["9"] = {
|
||||
"pid": 99,
|
||||
"process": None,
|
||||
"started_at": 0.0,
|
||||
"config": {},
|
||||
}
|
||||
lifecycle._kill_process_tree = lambda _pid: True
|
||||
|
||||
ok, message = lifecycle.stop_worker("9")
|
||||
|
||||
self.assertTrue(ok)
|
||||
self.assertEqual(message, "Worker stopped")
|
||||
self.assertNotIn("9", manager.processes)
|
||||
|
||||
def test_check_worker_process_falls_back_to_pid_probe(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
manager = _FakeManager(tmpdir)
|
||||
lifecycle = lifecycle_module.ProcessLifecycle(manager)
|
||||
lifecycle._is_process_running = lambda pid: int(pid) == 42
|
||||
|
||||
running, from_subprocess = lifecycle._check_worker_process(
|
||||
"any",
|
||||
{"pid": 42, "process": None},
|
||||
)
|
||||
self.assertTrue(running)
|
||||
self.assertFalse(from_subprocess)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,32 @@
|
||||
import io
|
||||
import unittest
|
||||
from contextlib import redirect_stdout
|
||||
from unittest.mock import patch
|
||||
|
||||
from utils import logging as logging_utils
|
||||
|
||||
|
||||
class LoggingUtilsTests(unittest.TestCase):
|
||||
def test_is_debug_enabled_delegates_to_config(self):
|
||||
with patch.object(logging_utils, "get_debug_enabled", return_value=True) as mocked:
|
||||
self.assertTrue(logging_utils.is_debug_enabled())
|
||||
mocked.assert_called_once_with(default=False)
|
||||
|
||||
def test_is_debug_enabled_handles_config_errors(self):
|
||||
with patch.object(logging_utils, "get_debug_enabled", side_effect=RuntimeError("boom")):
|
||||
self.assertFalse(logging_utils.is_debug_enabled())
|
||||
|
||||
def test_debug_log_emits_only_when_enabled(self):
|
||||
output = io.StringIO()
|
||||
with patch.object(logging_utils, "is_debug_enabled", return_value=True), redirect_stdout(output):
|
||||
logging_utils.debug_log("hello")
|
||||
self.assertIn("[Distributed] hello", output.getvalue())
|
||||
|
||||
output = io.StringIO()
|
||||
with patch.object(logging_utils, "is_debug_enabled", return_value=False), redirect_stdout(output):
|
||||
logging_utils.debug_log("hidden")
|
||||
self.assertEqual(output.getvalue(), "")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -6,7 +6,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_network_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "network.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "network.py"
|
||||
|
||||
package_name = "dist_utils_testpkg"
|
||||
package_module = types.ModuleType(package_name)
|
||||
@@ -17,6 +17,10 @@ def _load_network_module():
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.logging"] = logging_module
|
||||
|
||||
exceptions_module = types.ModuleType(f"{package_name}.exceptions")
|
||||
exceptions_module.DistributedError = type("DistributedError", (Exception,), {})
|
||||
sys.modules[f"{package_name}.exceptions"] = exceptions_module
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(
|
||||
instance=types.SimpleNamespace(address="127.0.0.1", port=8188, loop=None)
|
||||
@@ -90,48 +94,14 @@ class NetworkHelpersTests(unittest.TestCase):
|
||||
"https://master.example.com",
|
||||
)
|
||||
|
||||
def test_build_master_url_ignores_stale_saved_port_and_uses_runtime_port(self):
|
||||
cfg = {"master": {"host": "192.168.68.56", "port": 8001}}
|
||||
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
|
||||
self.assertEqual(
|
||||
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
|
||||
"http://192.168.68.56:8188",
|
||||
)
|
||||
|
||||
def test_build_master_url_keeps_explicit_port_in_host(self):
|
||||
cfg = {"master": {"host": "192.168.68.56:8001"}}
|
||||
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
|
||||
self.assertEqual(
|
||||
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
|
||||
"http://192.168.68.56:8001",
|
||||
)
|
||||
|
||||
def test_build_master_url_falls_back_to_server_address(self):
|
||||
cfg = {"master": {"host": "", "port": 8001}}
|
||||
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190)
|
||||
cfg = {"master": {"host": ""}}
|
||||
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190) # nosec B104 - wildcard bind normalization test
|
||||
self.assertEqual(
|
||||
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
|
||||
"http://127.0.0.1:8190",
|
||||
)
|
||||
|
||||
def test_build_master_callback_url_uses_loopback_for_local_worker(self):
|
||||
cfg = {"master": {"host": "192.168.68.56"}}
|
||||
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
|
||||
worker = {"id": "w1", "type": "local", "host": "localhost", "port": 8189}
|
||||
self.assertEqual(
|
||||
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
|
||||
"http://127.0.0.1:8001",
|
||||
)
|
||||
|
||||
def test_build_master_callback_url_keeps_public_master_url_for_remote_worker(self):
|
||||
cfg = {"master": {"host": "192.168.68.56"}}
|
||||
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
|
||||
worker = {"id": "w2", "type": "remote", "host": "192.168.68.99", "port": 8189}
|
||||
self.assertEqual(
|
||||
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
|
||||
"http://192.168.68.56:8001",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,705 @@
|
||||
import importlib.util
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _load_prompt_transform_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "prompt_transform.py"
|
||||
package_name = "dist_pt_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_pkg
|
||||
|
||||
orch_pkg = types.ModuleType(f"{package_name}.api.orchestration")
|
||||
orch_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api.orchestration"] = orch_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.api.orchestration.prompt_transform",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
pt = _load_prompt_transform_module()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _linear_prompt():
|
||||
"""1 → 2 → 3 → 4(DistributedCollector) → 5(SaveImage)"""
|
||||
return {
|
||||
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"2": {"class_type": "CLIPTextEncode", "inputs": {"clip": ["1", 1]}},
|
||||
"3": {"class_type": "KSampler", "inputs": {"model": ["1", 0], "positive": ["2", 0]}},
|
||||
"4": {"class_type": "DistributedCollector", "inputs": {"images": ["3", 0]}},
|
||||
"5": {"class_type": "SaveImage", "inputs": {"images": ["4", 0]}},
|
||||
}
|
||||
|
||||
|
||||
def _collector_only_prompt():
|
||||
"""1(Checkpoint) → 2(DistributedCollector) [no downstream from 2]"""
|
||||
return {
|
||||
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {"images": ["1", 0]}},
|
||||
}
|
||||
|
||||
|
||||
def _delegate_prompt():
|
||||
"""1 → 2 → 3(DistributedCollector) → 4(SaveImage)"""
|
||||
return {
|
||||
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"2": {"class_type": "KSampler", "inputs": {"model": ["1", 0]}},
|
||||
"3": {"class_type": "DistributedCollector", "inputs": {"images": ["2", 0]}},
|
||||
"4": {"class_type": "SaveImage", "inputs": {"images": ["3", 0]}},
|
||||
}
|
||||
|
||||
|
||||
def _branch_prompt():
|
||||
"""1 → 2(DistributedBranch) -> slot0:3->4, slot1:5->6."""
|
||||
return {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
|
||||
"4": {"class_type": "SaveImage", "inputs": {"images": ["3", 0]}},
|
||||
"5": {"class_type": "Sharpen", "inputs": {"image": ["2", 1]}},
|
||||
"6": {"class_type": "SaveImage", "inputs": {"images": ["5", 0]}},
|
||||
}
|
||||
|
||||
|
||||
def _branch_collector_prompt():
|
||||
"""1 → 2(DistributedBranch) -> 3(branch0) and 4(branch1) -> 5(DistributedBranchCollector) -> 6(SaveImage)."""
|
||||
return {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
|
||||
"4": {"class_type": "Sharpen", "inputs": {"image": ["2", 1]}},
|
||||
"5": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["3", 0], "num_branches": 2}},
|
||||
"6": {"class_type": "SaveImage", "inputs": {"images": ["5", 0]}},
|
||||
}
|
||||
|
||||
|
||||
def _apply(prompt, participant_id, enabled_worker_ids=None, delegate_master=False):
|
||||
if enabled_worker_ids is None:
|
||||
enabled_worker_ids = ["worker-a", "worker-b"]
|
||||
prompt_copy = json.loads(json.dumps(prompt))
|
||||
idx = pt.PromptIndex(prompt_copy)
|
||||
job_id_map = pt.generate_job_id_map(idx, "run")
|
||||
return pt.apply_participant_overrides(
|
||||
prompt_copy,
|
||||
participant_id=participant_id,
|
||||
enabled_worker_ids=enabled_worker_ids,
|
||||
job_id_map=job_id_map,
|
||||
master_url="http://master.example.com",
|
||||
delegate_master=delegate_master,
|
||||
prompt_index=idx,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# PromptIndex
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class PromptIndexTests(unittest.TestCase):
|
||||
def test_nodes_by_class_groups_correctly(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"3": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
self.assertCountEqual(idx.nodes_for_class("DistributedCollector"), ["2", "3"])
|
||||
self.assertEqual(idx.nodes_for_class("CheckpointLoaderSimple"), ["1"])
|
||||
|
||||
def test_nodes_for_class_unknown_returns_empty(self):
|
||||
idx = pt.PromptIndex({"1": {"class_type": "KSampler", "inputs": {}}})
|
||||
self.assertEqual(idx.nodes_for_class("Nonexistent"), [])
|
||||
|
||||
def test_nodes_without_class_type_are_indexed_under_none(self):
|
||||
prompt = {"1": {"inputs": {}}}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
# Should not raise; nodes_for_class with None key or missing class_type
|
||||
self.assertEqual(idx.nodes_for_class("KSampler"), [])
|
||||
|
||||
def test_copy_prompt_is_a_deep_copy(self):
|
||||
prompt = {"1": {"class_type": "KSampler", "inputs": {"seed": 42}}}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
copy = idx.copy_prompt()
|
||||
copy["1"]["inputs"]["seed"] = 999
|
||||
self.assertEqual(prompt["1"]["inputs"]["seed"], 42)
|
||||
|
||||
def test_has_upstream_direct_connection(self):
|
||||
"""Node 4 reads directly from node 3 (KSampler)."""
|
||||
idx = pt.PromptIndex(_linear_prompt())
|
||||
self.assertTrue(idx.has_upstream("4", "KSampler"))
|
||||
|
||||
def test_has_upstream_transitive_connection(self):
|
||||
"""Node 4 → 3 → 2 → 1 (CheckpointLoaderSimple)."""
|
||||
idx = pt.PromptIndex(_linear_prompt())
|
||||
self.assertTrue(idx.has_upstream("4", "CheckpointLoaderSimple"))
|
||||
|
||||
def test_has_upstream_returns_false_when_no_path(self):
|
||||
idx = pt.PromptIndex(_linear_prompt())
|
||||
# CheckpointLoaderSimple has no upstream nodes
|
||||
self.assertFalse(idx.has_upstream("1", "DistributedCollector"))
|
||||
|
||||
def test_has_upstream_result_is_cached(self):
|
||||
idx = pt.PromptIndex(_linear_prompt())
|
||||
r1 = idx.has_upstream("4", "KSampler")
|
||||
r2 = idx.has_upstream("4", "KSampler")
|
||||
self.assertEqual(r1, r2)
|
||||
self.assertIn(("4", "KSampler"), idx._upstream_cache)
|
||||
|
||||
def test_has_upstream_does_not_infinite_loop_on_cycle(self):
|
||||
"""Cyclic references in inputs should not cause infinite recursion."""
|
||||
prompt = {
|
||||
"1": {"class_type": "A", "inputs": {"x": ["2", 0]}},
|
||||
"2": {"class_type": "B", "inputs": {"x": ["1", 0]}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
# Should terminate without error
|
||||
result = idx.has_upstream("1", "NonExistent")
|
||||
self.assertFalse(result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# find_nodes_by_class
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class FindNodesByClassTests(unittest.TestCase):
|
||||
def test_finds_matching_nodes(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
}
|
||||
result = pt.find_nodes_by_class(prompt, "KSampler")
|
||||
self.assertEqual(result, ["1"])
|
||||
|
||||
def test_returns_empty_when_no_match(self):
|
||||
prompt = {"1": {"class_type": "KSampler", "inputs": {}}}
|
||||
self.assertEqual(pt.find_nodes_by_class(prompt, "DistributedCollector"), [])
|
||||
|
||||
def test_skips_non_dict_nodes(self):
|
||||
prompt = {"1": "not a dict", "2": {"class_type": "KSampler", "inputs": {}}}
|
||||
result = pt.find_nodes_by_class(prompt, "KSampler")
|
||||
self.assertEqual(result, ["2"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# prune_prompt_for_worker
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class PrunePromptForWorkerTests(unittest.TestCase):
|
||||
def test_no_distributed_nodes_returns_prompt_unchanged(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
|
||||
"2": {"class_type": "SaveImage", "inputs": {"images": ["1", 0]}},
|
||||
}
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertCountEqual(result.keys(), ["1", "2"])
|
||||
|
||||
def test_keeps_collector_and_upstream(self):
|
||||
prompt = _linear_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
for node_id in ("1", "2", "3", "4"):
|
||||
self.assertIn(node_id, result)
|
||||
|
||||
def test_removes_downstream_of_collector(self):
|
||||
prompt = _linear_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertNotIn("5", result)
|
||||
|
||||
def test_injects_preview_image_when_downstream_exists(self):
|
||||
prompt = _linear_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
|
||||
self.assertEqual(len(preview_nodes), 1)
|
||||
self.assertEqual(preview_nodes[0]["inputs"]["images"], ["4", 0])
|
||||
|
||||
def test_no_preview_image_when_no_downstream(self):
|
||||
result = pt.prune_prompt_for_worker(_collector_only_prompt())
|
||||
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
|
||||
self.assertEqual(len(preview_nodes), 0)
|
||||
|
||||
def test_unrelated_nodes_are_pruned(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"2": {"class_type": "UnrelatedNode", "inputs": {}}, # no connection to 1
|
||||
}
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertIn("1", result)
|
||||
self.assertNotIn("2", result)
|
||||
|
||||
def test_result_is_a_copy_not_same_object(self):
|
||||
prompt = _linear_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
# Mutating the result should not affect the original
|
||||
original_keys = set(prompt.keys())
|
||||
result["NEW"] = {"class_type": "Test", "inputs": {}}
|
||||
self.assertEqual(set(prompt.keys()), original_keys)
|
||||
|
||||
def test_upscale_node_is_treated_as_distributed(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"image": ["1", 0]}},
|
||||
"3": {"class_type": "SaveImage", "inputs": {"images": ["2", 0]}},
|
||||
}
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertIn("1", result)
|
||||
self.assertIn("2", result)
|
||||
self.assertNotIn("3", result)
|
||||
|
||||
def test_branch_anchor_keeps_downstream_for_later_branch_pruning(self):
|
||||
prompt = _branch_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertIn("2", result)
|
||||
self.assertIn("3", result)
|
||||
self.assertIn("4", result)
|
||||
self.assertIn("5", result)
|
||||
self.assertIn("6", result)
|
||||
|
||||
def test_branch_collector_anchor_keeps_collector_but_prunes_downstream(self):
|
||||
prompt = _branch_collector_prompt()
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertIn("5", result)
|
||||
self.assertNotIn("6", result)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# prepare_delegate_master_prompt
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class PrepareDelegateMasterPromptTests(unittest.TestCase):
|
||||
def test_keeps_collector_and_downstream(self):
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
self.assertIn("3", result)
|
||||
self.assertIn("4", result)
|
||||
self.assertNotIn("1", result)
|
||||
self.assertNotIn("2", result)
|
||||
|
||||
def test_removes_dangling_upstream_refs(self):
|
||||
"""Collector must not retain dangling refs to pruned upstream nodes."""
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
collector_inputs = result["3"].get("inputs", {})
|
||||
# Original "images" pointed at node 2, which is pruned.
|
||||
# It should now point at a newly injected placeholder node.
|
||||
self.assertIn("images", collector_inputs)
|
||||
source_id = str(collector_inputs["images"][0])
|
||||
self.assertNotEqual(source_id, "2")
|
||||
self.assertIn(source_id, result)
|
||||
self.assertEqual(result[source_id].get("class_type"), "DistributedEmptyImage")
|
||||
|
||||
def test_injects_empty_image_placeholder(self):
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
empty_nodes = [(nid, n) for nid, n in result.items() if n.get("class_type") == "DistributedEmptyImage"]
|
||||
self.assertEqual(len(empty_nodes), 1)
|
||||
placeholder_id = empty_nodes[0][0]
|
||||
self.assertEqual(result["3"]["inputs"]["images"], [placeholder_id, 0])
|
||||
|
||||
def test_one_placeholder_per_collector(self):
|
||||
"""Two collectors → two placeholders."""
|
||||
prompt = {
|
||||
"1": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"3": {"class_type": "SaveImage", "inputs": {"images": ["1", 0]}},
|
||||
}
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["1", "2"])
|
||||
empty_nodes = [n for n in result.values() if n.get("class_type") == "DistributedEmptyImage"]
|
||||
self.assertEqual(len(empty_nodes), 2)
|
||||
|
||||
def test_result_is_independent_copy(self):
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
result["3"]["inputs"]["NEW"] = "injected"
|
||||
# Original should be untouched
|
||||
self.assertNotIn("NEW", prompt["3"].get("inputs", {}))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# generate_job_id_map
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class GenerateJobIdMapTests(unittest.TestCase):
|
||||
def test_maps_collector_nodes(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"2": {"class_type": "KSampler", "inputs": {}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_map = pt.generate_job_id_map(idx, "prefix")
|
||||
self.assertEqual(job_map["1"], "prefix_1")
|
||||
self.assertNotIn("2", job_map)
|
||||
|
||||
def test_maps_upscale_nodes(self):
|
||||
prompt = {
|
||||
"5": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_map = pt.generate_job_id_map(idx, "run")
|
||||
self.assertEqual(job_map["5"], "run_5")
|
||||
|
||||
def test_maps_branch_collector_nodes(self):
|
||||
prompt = {
|
||||
"10": {"class_type": "DistributedBranchCollector", "inputs": {}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_map = pt.generate_job_id_map(idx, "run")
|
||||
self.assertEqual(job_map["10"], "run_10")
|
||||
|
||||
def test_maps_branch_nodes(self):
|
||||
prompt = {
|
||||
"9": {"class_type": "DistributedBranch", "inputs": {"num_branches": 2}},
|
||||
}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_map = pt.generate_job_id_map(idx, "run")
|
||||
self.assertEqual(job_map["9"], "run_9")
|
||||
|
||||
def test_empty_prompt_returns_empty_map(self):
|
||||
idx = pt.PromptIndex({})
|
||||
self.assertEqual(pt.generate_job_id_map(idx, "prefix"), {})
|
||||
|
||||
def test_stable_ids_across_calls(self):
|
||||
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {}}}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
m1 = pt.generate_job_id_map(idx, "run")
|
||||
m2 = pt.generate_job_id_map(idx, "run")
|
||||
self.assertEqual(m1, m2)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apply_participant_overrides – DistributedCollector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ApplyOverridesCollectorTests(unittest.TestCase):
|
||||
def _collector_prompt(self):
|
||||
return {"1": {"class_type": "DistributedCollector", "inputs": {}}}
|
||||
|
||||
def test_worker_sets_is_worker_true(self):
|
||||
result = _apply(self._collector_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_sets_master_url(self):
|
||||
result = _apply(self._collector_prompt(), "worker-a")
|
||||
self.assertEqual(result["1"]["inputs"]["master_url"], "http://master.example.com")
|
||||
|
||||
def test_worker_sets_worker_id(self):
|
||||
result = _apply(self._collector_prompt(), "worker-a")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-a")
|
||||
|
||||
def test_worker_sets_delegate_only_false(self):
|
||||
result = _apply(self._collector_prompt(), "worker-a")
|
||||
self.assertFalse(result["1"]["inputs"]["delegate_only"])
|
||||
|
||||
def test_master_sets_is_worker_false(self):
|
||||
result = _apply(self._collector_prompt(), "master")
|
||||
self.assertFalse(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_master_clears_stale_master_url(self):
|
||||
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {"master_url": "stale"}}}
|
||||
result = _apply(prompt, "master")
|
||||
self.assertNotIn("master_url", result["1"]["inputs"])
|
||||
|
||||
def test_master_clears_stale_worker_id(self):
|
||||
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {"worker_id": "stale"}}}
|
||||
result = _apply(prompt, "master")
|
||||
self.assertNotIn("worker_id", result["1"]["inputs"])
|
||||
|
||||
def test_master_with_delegate_master_sets_delegate_only_true(self):
|
||||
result = _apply(self._collector_prompt(), "master", delegate_master=True)
|
||||
self.assertTrue(result["1"]["inputs"]["delegate_only"])
|
||||
|
||||
def test_master_without_delegate_master_sets_delegate_only_false(self):
|
||||
result = _apply(self._collector_prompt(), "master", delegate_master=False)
|
||||
self.assertFalse(result["1"]["inputs"]["delegate_only"])
|
||||
|
||||
def test_enabled_worker_ids_serialized_as_json(self):
|
||||
enabled = ["worker-a", "worker-b"]
|
||||
result = _apply(self._collector_prompt(), "master", enabled_worker_ids=enabled)
|
||||
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
|
||||
|
||||
def test_multi_job_id_is_set_from_job_map(self):
|
||||
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {}}}
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_id_map = {"1": "run_abc_1"}
|
||||
result = pt.apply_participant_overrides(
|
||||
prompt,
|
||||
participant_id="worker-a",
|
||||
enabled_worker_ids=["worker-a"],
|
||||
job_id_map=job_id_map,
|
||||
master_url="http://master",
|
||||
delegate_master=False,
|
||||
prompt_index=idx,
|
||||
)
|
||||
self.assertEqual(result["1"]["inputs"]["multi_job_id"], "run_abc_1")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apply_participant_overrides – DistributedSeed
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ApplyOverridesSeedTests(unittest.TestCase):
|
||||
def _seed_prompt(self):
|
||||
return {"1": {"class_type": "DistributedSeed", "inputs": {}}}
|
||||
|
||||
def test_worker_sets_is_worker_true(self):
|
||||
result = _apply(self._seed_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_id_uses_canonical_worker_id(self):
|
||||
result = _apply(self._seed_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
|
||||
|
||||
def test_master_sets_is_worker_false(self):
|
||||
result = _apply(self._seed_prompt(), "master")
|
||||
self.assertFalse(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_master_sets_empty_worker_id(self):
|
||||
result = _apply(self._seed_prompt(), "master")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "")
|
||||
|
||||
def test_enabled_worker_ids_is_set(self):
|
||||
enabled = ["worker-a", "worker-b"]
|
||||
result = _apply(self._seed_prompt(), "worker-a", enabled_worker_ids=enabled)
|
||||
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apply_participant_overrides – UltimateSDUpscaleDistributed
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ApplyOverridesUpscaleTests(unittest.TestCase):
|
||||
def _upscale_prompt(self):
|
||||
return {"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}}}
|
||||
|
||||
def test_worker_sets_is_worker_true(self):
|
||||
result = _apply(self._upscale_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_sets_master_url_and_worker_id(self):
|
||||
result = _apply(self._upscale_prompt(), "worker-a")
|
||||
self.assertEqual(result["1"]["inputs"]["master_url"], "http://master.example.com")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-a")
|
||||
|
||||
def test_master_clears_master_url_and_worker_id(self):
|
||||
prompt = {"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"master_url": "x", "worker_id": "y"}}}
|
||||
result = _apply(prompt, "master")
|
||||
self.assertNotIn("master_url", result["1"]["inputs"])
|
||||
self.assertNotIn("worker_id", result["1"]["inputs"])
|
||||
|
||||
def test_collector_downstream_of_upscale_gets_pass_through(self):
|
||||
"""A DistributedCollector that is downstream of UltimateSDUpscaleDistributed → pass_through=True."""
|
||||
prompt = {
|
||||
"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {"images": ["1", 0]}},
|
||||
}
|
||||
result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
|
||||
self.assertTrue(result["2"]["inputs"].get("pass_through"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apply_participant_overrides – DistributedValue
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ApplyOverridesValueTests(unittest.TestCase):
|
||||
def _value_prompt(self):
|
||||
return {"1": {"class_type": "DistributedValue", "inputs": {}}}
|
||||
|
||||
def test_worker_sets_is_worker_true(self):
|
||||
result = _apply(self._value_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_id_uses_canonical_worker_id(self):
|
||||
result = _apply(self._value_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
|
||||
|
||||
def test_master_sets_is_worker_false(self):
|
||||
result = _apply(self._value_prompt(), "master")
|
||||
self.assertFalse(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_master_sets_empty_worker_id(self):
|
||||
result = _apply(self._value_prompt(), "master")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "")
|
||||
|
||||
def test_enabled_worker_ids_is_set(self):
|
||||
enabled = ["worker-a", "worker-b"]
|
||||
result = _apply(self._value_prompt(), "worker-a", enabled_worker_ids=enabled)
|
||||
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
|
||||
|
||||
|
||||
class ApplyOverridesBranchTests(unittest.TestCase):
|
||||
def test_master_gets_branch_zero_worker_gets_branch_one(self):
|
||||
prompt = _branch_prompt()
|
||||
master_result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
|
||||
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
|
||||
|
||||
self.assertEqual(master_result["2"]["inputs"]["assigned_branch"], 0)
|
||||
self.assertEqual(worker_result["2"]["inputs"]["assigned_branch"], 1)
|
||||
self.assertIn("3", master_result)
|
||||
self.assertIn("4", master_result)
|
||||
self.assertNotIn("5", master_result)
|
||||
self.assertNotIn("6", master_result)
|
||||
self.assertIn("5", worker_result)
|
||||
self.assertIn("6", worker_result)
|
||||
self.assertNotIn("3", worker_result)
|
||||
self.assertNotIn("4", worker_result)
|
||||
|
||||
def test_delegate_mode_assigns_worker_a_to_branch_zero(self):
|
||||
prompt = _branch_prompt()
|
||||
worker_result = _apply(
|
||||
prompt,
|
||||
"worker-a",
|
||||
enabled_worker_ids=["worker-a", "worker-b"],
|
||||
delegate_master=True,
|
||||
)
|
||||
self.assertEqual(worker_result["2"]["inputs"]["assigned_branch"], 0)
|
||||
|
||||
def test_pruning_keeps_shared_nodes_between_branches(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "NodeA", "inputs": {"image": ["2", 0]}},
|
||||
"4": {"class_type": "NodeB", "inputs": {"image": ["2", 1]}},
|
||||
"5": {"class_type": "SaveImage", "inputs": {"images": ["3", 0], "aux": ["4", 0]}},
|
||||
}
|
||||
worker_result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
|
||||
self.assertIn("3", worker_result)
|
||||
self.assertNotIn("4", worker_result)
|
||||
self.assertIn("5", worker_result)
|
||||
self.assertNotIn("aux", worker_result["5"]["inputs"])
|
||||
|
||||
def test_more_participants_than_branches_sets_unassigned_to_minus_one(self):
|
||||
prompt = _branch_prompt()
|
||||
worker_b_result = _apply(
|
||||
prompt,
|
||||
"worker-b",
|
||||
enabled_worker_ids=["worker-a", "worker-b", "worker-c"],
|
||||
)
|
||||
self.assertEqual(worker_b_result["2"]["inputs"]["assigned_branch"], -1)
|
||||
class_types = {node.get("class_type") for node in worker_b_result.values()}
|
||||
self.assertNotIn("Blur", class_types)
|
||||
self.assertNotIn("Sharpen", class_types)
|
||||
|
||||
def test_worker_with_pruned_outputs_gets_auto_preview_for_assigned_branch(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 3}},
|
||||
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
|
||||
"4": {"class_type": "PreviewImage", "inputs": {"images": ["3", 0]}},
|
||||
}
|
||||
|
||||
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a", "worker-b"])
|
||||
preview_nodes = [node for node in worker_result.values() if node.get("class_type") == "PreviewImage"]
|
||||
|
||||
self.assertTrue(preview_nodes)
|
||||
self.assertIn(["2", 1], [node.get("inputs", {}).get("images") for node in preview_nodes])
|
||||
|
||||
def test_unassigned_worker_gets_idle_fallback_output(self):
|
||||
prompt = _branch_prompt()
|
||||
worker_result = _apply(
|
||||
prompt,
|
||||
"worker-c",
|
||||
enabled_worker_ids=["worker-a", "worker-b", "worker-c"],
|
||||
)
|
||||
|
||||
preview_nodes = [node for node in worker_result.values() if node.get("class_type") == "PreviewImage"]
|
||||
empty_nodes = [node for node in worker_result.values() if node.get("class_type") == "DistributedEmptyImage"]
|
||||
|
||||
self.assertTrue(preview_nodes)
|
||||
self.assertTrue(empty_nodes)
|
||||
|
||||
|
||||
class ApplyOverridesBranchCollectorTests(unittest.TestCase):
|
||||
def test_branch_collector_inherits_assigned_branch_from_upstream_branch_node(self):
|
||||
master_prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0], "num_branches": 2}},
|
||||
}
|
||||
worker_prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 1], "num_branches": 2}},
|
||||
}
|
||||
|
||||
master_result = _apply(master_prompt, "master", enabled_worker_ids=["worker-a"])
|
||||
worker_result = _apply(worker_prompt, "worker-a", enabled_worker_ids=["worker-a"])
|
||||
|
||||
self.assertEqual(master_result["3"]["inputs"]["assigned_branch"], 0)
|
||||
self.assertEqual(worker_result["3"]["inputs"]["assigned_branch"], 1)
|
||||
|
||||
def test_branch_collector_sets_multi_job_id(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0]}},
|
||||
"2": {"class_type": "KSampler", "inputs": {}},
|
||||
}
|
||||
result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
|
||||
self.assertEqual(result["1"]["inputs"]["multi_job_id"], "run_1")
|
||||
|
||||
def test_branch_collector_uses_upstream_branch_job_id_for_grouped_convergence(self):
|
||||
master_prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0], "num_branches": 2}},
|
||||
}
|
||||
worker_prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 1], "num_branches": 2}},
|
||||
}
|
||||
|
||||
master_result = _apply(master_prompt, "master", enabled_worker_ids=["worker-a"])
|
||||
worker_result = _apply(worker_prompt, "worker-a", enabled_worker_ids=["worker-a"])
|
||||
|
||||
self.assertEqual(master_result["3"]["inputs"]["multi_job_id"], "run_2")
|
||||
self.assertEqual(worker_result["3"]["inputs"]["multi_job_id"], "run_2")
|
||||
|
||||
def test_worker_prunes_nodes_downstream_of_branch_collector(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
|
||||
"3": {"class_type": "Blur", "inputs": {"image": ["2", 1]}},
|
||||
"4": {
|
||||
"class_type": "DistributedBranchCollector",
|
||||
"inputs": {"branch_2": ["3", 0], "num_branches": 2},
|
||||
},
|
||||
"5": {"class_type": "SaveImage", "inputs": {"images": ["4", 0]}},
|
||||
}
|
||||
|
||||
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
|
||||
|
||||
class_types = {node.get("class_type") for node in worker_result.values()}
|
||||
self.assertIn("DistributedBranchCollector", class_types)
|
||||
self.assertNotIn("SaveImage", class_types)
|
||||
self.assertIn("PreviewImage", class_types)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -4,7 +4,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_queue_request_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "api" / "queue_request.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "api" / "queue_request.py"
|
||||
spec = importlib.util.spec_from_file_location("queue_request", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
@@ -26,27 +26,26 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
|
||||
def test_normalizes_enabled_worker_ids(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data["enabled_worker_ids"] = ["a", 2, 3]
|
||||
payload_data["enabled_worker_ids"] = [" worker-a ", "worker-b", "worker-a"]
|
||||
payload_data["delegate_master"] = True
|
||||
payload = parse_queue_request_payload(
|
||||
payload_data
|
||||
)
|
||||
self.assertEqual(payload.enabled_worker_ids, ["a", "2", "3"])
|
||||
self.assertEqual(payload.enabled_worker_ids, ["worker-a", "worker-b"])
|
||||
self.assertTrue(payload.delegate_master)
|
||||
|
||||
def test_supports_legacy_workers_field(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data.pop("enabled_worker_ids", None)
|
||||
payload_data["workers"] = [{"id": "w1"}, "w2", {"id": 3}, {"name": "no-id"}]
|
||||
payload_data["workers"] = [{"id": "w1"}, "w2", {"id": "w3"}, {"name": "no-id"}]
|
||||
payload = parse_queue_request_payload(
|
||||
payload_data
|
||||
)
|
||||
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "3"])
|
||||
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "w3"])
|
||||
|
||||
def test_supports_auto_prepare_prompt_fallback(self):
|
||||
def test_supports_workflow_prompt_fallback(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data.pop("prompt", None)
|
||||
payload_data["auto_prepare"] = True
|
||||
payload_data["workflow"] = {
|
||||
"prompt": {
|
||||
"10": {"class_type": "DistributedCollector"},
|
||||
@@ -56,7 +55,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
payload_data
|
||||
)
|
||||
self.assertIn("10", payload.prompt)
|
||||
self.assertTrue(payload.auto_prepare)
|
||||
|
||||
def test_normalizes_trace_execution_id(self):
|
||||
payload_data = self._base_payload()
|
||||
@@ -74,10 +72,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
)
|
||||
self.assertIsNone(payload.trace_execution_id)
|
||||
|
||||
def test_auto_prepare_defaults_true(self):
|
||||
payload = parse_queue_request_payload(self._base_payload())
|
||||
self.assertTrue(payload.auto_prepare)
|
||||
|
||||
def test_workers_field_must_be_list(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data.pop("enabled_worker_ids", None)
|
||||
@@ -91,7 +85,7 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
with self.assertRaisesRegex(ValueError, "trace_execution_id must be a string"):
|
||||
parse_queue_request_payload(payload_data)
|
||||
|
||||
def test_auto_prepare_false_still_falls_back_to_workflow_prompt(self):
|
||||
def test_auto_prepare_is_ignored_for_backward_compat(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data.pop("prompt", None)
|
||||
payload_data["auto_prepare"] = False
|
||||
@@ -100,13 +94,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
}
|
||||
payload = parse_queue_request_payload(payload_data)
|
||||
self.assertIn("10", payload.prompt)
|
||||
self.assertFalse(payload.auto_prepare)
|
||||
|
||||
def test_auto_prepare_must_be_boolean(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data["auto_prepare"] = "true"
|
||||
with self.assertRaisesRegex(ValueError, "auto_prepare must be a boolean"):
|
||||
parse_queue_request_payload(payload_data)
|
||||
|
||||
def test_invalid_delegate_master_type_raises(self):
|
||||
payload_data = self._base_payload()
|
||||
@@ -120,6 +107,12 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
with self.assertRaisesRegex(ValueError, "enabled_worker_ids must be a list"):
|
||||
parse_queue_request_payload(payload_data)
|
||||
|
||||
def test_legacy_index_worker_tokens_are_rejected(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data["enabled_worker_ids"] = ["worker-a", "0", "worker_1"]
|
||||
with self.assertRaisesRegex(ValueError, "legacy index token"):
|
||||
parse_queue_request_payload(payload_data)
|
||||
|
||||
def test_invalid_top_level_payload_raises(self):
|
||||
with self.assertRaisesRegex(ValueError, "Expected a JSON object body"):
|
||||
parse_queue_request_payload(["not", "an", "object"])
|
||||
@@ -0,0 +1,154 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _load_worker_monitor_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "workers" / "worker_monitor.py"
|
||||
spec = importlib.util.spec_from_file_location("dist_test_worker_monitor", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
worker_monitor = _load_worker_monitor_module()
|
||||
|
||||
|
||||
class _FakeWorkerProcess:
|
||||
def __init__(self, pid=9001, poll_sequence=None):
|
||||
self.pid = pid
|
||||
self._poll_sequence = list(poll_sequence or [None])
|
||||
self._poll_calls = 0
|
||||
self.returncode = None
|
||||
self.terminate_calls = 0
|
||||
self.kill_calls = 0
|
||||
self.wait_calls = 0
|
||||
|
||||
def poll(self):
|
||||
if self._poll_calls < len(self._poll_sequence):
|
||||
value = self._poll_sequence[self._poll_calls]
|
||||
else:
|
||||
value = self._poll_sequence[-1]
|
||||
self._poll_calls += 1
|
||||
if value is not None:
|
||||
self.returncode = value
|
||||
return value
|
||||
|
||||
def terminate(self):
|
||||
self.terminate_calls += 1
|
||||
self.returncode = 0
|
||||
|
||||
def kill(self):
|
||||
self.kill_calls += 1
|
||||
self.returncode = -9
|
||||
|
||||
def wait(self, timeout=None):
|
||||
_ = timeout
|
||||
self.wait_calls += 1
|
||||
if self.returncode is None:
|
||||
self.returncode = 0
|
||||
return self.returncode
|
||||
|
||||
|
||||
class WorkerMonitorTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self._orig_launch = worker_monitor.launch_process_with_timeout
|
||||
self._orig_is_alive = worker_monitor.is_process_alive
|
||||
self._orig_terminate = getattr(worker_monitor, "terminate_process", None)
|
||||
self._orig_sleep = worker_monitor.time.sleep
|
||||
self._orig_signal = worker_monitor.signal.signal
|
||||
self._env_backup = dict(os.environ)
|
||||
|
||||
worker_monitor.time.sleep = lambda _seconds: None
|
||||
worker_monitor.signal.signal = lambda *_args, **_kwargs: None
|
||||
|
||||
def tearDown(self):
|
||||
worker_monitor.launch_process_with_timeout = self._orig_launch
|
||||
worker_monitor.is_process_alive = self._orig_is_alive
|
||||
if self._orig_terminate is not None:
|
||||
worker_monitor.terminate_process = self._orig_terminate
|
||||
worker_monitor.time.sleep = self._orig_sleep
|
||||
worker_monitor.signal.signal = self._orig_signal
|
||||
os.environ.clear()
|
||||
os.environ.update(self._env_backup)
|
||||
|
||||
def test_main_validates_required_env_and_args(self):
|
||||
os.environ.pop("COMFYUI_MASTER_PID", None)
|
||||
self.assertEqual(worker_monitor.main(["python", "main.py"]), 1)
|
||||
|
||||
os.environ["COMFYUI_MASTER_PID"] = "not-an-int"
|
||||
self.assertEqual(worker_monitor.main(["python", "main.py"]), 1)
|
||||
|
||||
os.environ["COMFYUI_MASTER_PID"] = "101"
|
||||
self.assertEqual(worker_monitor.main([]), 1)
|
||||
|
||||
def test_main_delegates_to_monitor_and_run(self):
|
||||
calls = []
|
||||
|
||||
def _monitor(master_pid, command):
|
||||
calls.append((master_pid, list(command)))
|
||||
return 7
|
||||
|
||||
os.environ["COMFYUI_MASTER_PID"] = "202"
|
||||
original_monitor = worker_monitor.monitor_and_run
|
||||
worker_monitor.monitor_and_run = _monitor
|
||||
try:
|
||||
exit_code = worker_monitor.main(["python", "worker.py"])
|
||||
finally:
|
||||
worker_monitor.monitor_and_run = original_monitor
|
||||
|
||||
self.assertEqual(exit_code, 7)
|
||||
self.assertEqual(calls, [(202, ["python", "worker.py"])])
|
||||
|
||||
def test_monitor_and_run_returns_worker_exit_code(self):
|
||||
fake_process = _FakeWorkerProcess(poll_sequence=[None, 5])
|
||||
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
|
||||
worker_monitor.is_process_alive = lambda _pid: True
|
||||
worker_monitor.terminate_process = lambda *_args, **_kwargs: None
|
||||
|
||||
exit_code = worker_monitor.monitor_and_run(100, ["python", "worker.py"])
|
||||
|
||||
self.assertEqual(exit_code, 5)
|
||||
|
||||
def test_monitor_and_run_terminates_worker_when_master_dies(self):
|
||||
fake_process = _FakeWorkerProcess(poll_sequence=[None, None])
|
||||
term_calls = []
|
||||
|
||||
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
|
||||
worker_monitor.is_process_alive = lambda _pid: False
|
||||
|
||||
def _terminate(process, timeout):
|
||||
term_calls.append((process.pid, timeout))
|
||||
process.returncode = 0
|
||||
|
||||
worker_monitor.terminate_process = _terminate
|
||||
|
||||
exit_code = worker_monitor.monitor_and_run(333, ["python", "worker.py"])
|
||||
|
||||
self.assertEqual(exit_code, 0)
|
||||
self.assertEqual(term_calls, [(9001, worker_monitor.PROCESS_TERMINATION_TIMEOUT)])
|
||||
|
||||
def test_monitor_and_run_writes_pid_file_when_requested(self):
|
||||
fake_process = _FakeWorkerProcess(poll_sequence=[0])
|
||||
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
|
||||
worker_monitor.is_process_alive = lambda _pid: True
|
||||
worker_monitor.terminate_process = lambda *_args, **_kwargs: None
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
pid_file = Path(tmpdir) / "worker.pid"
|
||||
os.environ["WORKER_PID_FILE"] = str(pid_file)
|
||||
exit_code = worker_monitor.monitor_and_run(444, ["python", "worker.py"])
|
||||
|
||||
self.assertEqual(exit_code, 0)
|
||||
contents = pid_file.read_text(encoding="utf-8")
|
||||
monitor_pid, worker_pid = contents.split(",")
|
||||
self.assertTrue(monitor_pid.isdigit())
|
||||
self.assertEqual(worker_pid, "9001")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests."""
|
||||
@@ -0,0 +1,53 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from upscale.conditioning import clone_conditioning, clone_control_chain
|
||||
|
||||
|
||||
class _FakeControl:
|
||||
def __init__(self, hint, previous=None):
|
||||
self.cond_hint_original = hint
|
||||
self.previous_controlnet = previous
|
||||
|
||||
|
||||
class ConditioningUtilsTests(unittest.TestCase):
|
||||
def test_clone_control_chain_clones_hints_and_links(self):
|
||||
tail = _FakeControl(torch.ones((1, 2)))
|
||||
head = _FakeControl(torch.zeros((1, 2)), previous=tail)
|
||||
|
||||
cloned = clone_control_chain(head, clone_hint=True)
|
||||
|
||||
self.assertIsNot(cloned, head)
|
||||
self.assertIsNot(cloned.previous_controlnet, tail)
|
||||
self.assertTrue(torch.equal(cloned.cond_hint_original, head.cond_hint_original))
|
||||
self.assertTrue(torch.equal(cloned.previous_controlnet.cond_hint_original, tail.cond_hint_original))
|
||||
self.assertIsNot(cloned.cond_hint_original, head.cond_hint_original)
|
||||
|
||||
def test_clone_conditioning_clones_emb_and_known_fields(self):
|
||||
control = _FakeControl(torch.ones((1, 3)))
|
||||
cond = [
|
||||
[
|
||||
torch.tensor([[1.0, 2.0, 3.0]]),
|
||||
{
|
||||
"control": control,
|
||||
"mask": torch.ones((1, 1, 1)),
|
||||
"pooled_output": torch.zeros((1, 4)),
|
||||
"area": [0, 1, 2, 3],
|
||||
},
|
||||
]
|
||||
]
|
||||
|
||||
cloned = clone_conditioning(cond, clone_hints=True)
|
||||
|
||||
self.assertTrue(torch.equal(cloned[0][0], cond[0][0]))
|
||||
self.assertIsNot(cloned[0][0], cond[0][0])
|
||||
self.assertIsNot(cloned[0][1]["control"], cond[0][1]["control"])
|
||||
self.assertIsNot(cloned[0][1]["mask"], cond[0][1]["mask"])
|
||||
self.assertIsNot(cloned[0][1]["pooled_output"], cond[0][1]["pooled_output"])
|
||||
self.assertEqual(cloned[0][1]["area"], cond[0][1]["area"])
|
||||
self.assertIsNot(cloned[0][1]["area"], cond[0][1]["area"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,247 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError: # pragma: no cover - optional dependency in CI/runtime
|
||||
torch = None
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
|
||||
modes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
|
||||
def tensor_to_pil(image_batch, index):
|
||||
tensor = image_batch[index].detach().cpu().clamp(0, 1)
|
||||
array = (tensor.numpy() * 255.0).astype("uint8")
|
||||
return Image.fromarray(array)
|
||||
|
||||
def pil_to_tensor(image):
|
||||
if torch is None:
|
||||
raise RuntimeError("torch is required for this test helper")
|
||||
array = np.array(image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array).unsqueeze(0)
|
||||
|
||||
image_module.tensor_to_pil = tensor_to_pil
|
||||
image_module.pil_to_tensor = pil_to_tensor
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda _coro, timeout=None: (_ for _ in ()).throw(
|
||||
RuntimeError(f"run_async_in_server_loop unexpectedly called with timeout={timeout}")
|
||||
)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 5.0
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.JOB_POLL_INTERVAL = 0.0
|
||||
constants_module.JOB_POLL_MAX_ATTEMPTS = 3
|
||||
constants_module.TILE_WAIT_TIMEOUT = 5.0
|
||||
constants_module.TILE_SEND_TIMEOUT = 5.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
if torch is None:
|
||||
torch_module = types.ModuleType("torch")
|
||||
torch_module.cat = lambda *_args, **_kwargs: None
|
||||
sys.modules["torch"] = torch_module
|
||||
|
||||
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
|
||||
job_store_module.ensure_tile_jobs_initialized = lambda: None
|
||||
|
||||
async def init_dynamic_job(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
job_store_module.init_dynamic_job = init_dynamic_job
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
model_mgmt = types.ModuleType("comfy.model_management")
|
||||
model_mgmt.throw_exception_if_processing_interrupted = lambda: None
|
||||
comfy_module.model_management = model_mgmt
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.model_management"] = model_mgmt
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_dynamic_mode_module():
|
||||
package_name = "dist_dynamic_mode_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
|
||||
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
|
||||
return _load_module(package_name, "upscale/modes/dynamic.py", "upscale.modes.dynamic")
|
||||
|
||||
|
||||
dynamic_module = _load_dynamic_mode_module()
|
||||
|
||||
|
||||
def _run_sync(value, timeout=None):
|
||||
_ = timeout
|
||||
if asyncio.iscoroutine(value):
|
||||
return asyncio.run(value)
|
||||
return value
|
||||
|
||||
|
||||
dynamic_module.run_async_in_server_loop = _run_sync
|
||||
|
||||
|
||||
class _DummyDynamicNode(dynamic_module.DynamicModeMixin):
|
||||
def __init__(self, *, poll_ready=True, assignments=None):
|
||||
self._poll_ready = poll_ready
|
||||
self._assignments = list(assignments or [])
|
||||
self.sent_images = []
|
||||
self.heartbeat_count = 0
|
||||
self.completion_sent = False
|
||||
|
||||
def round_to_multiple(self, value):
|
||||
return int(value)
|
||||
|
||||
def calculate_tiles(self, _width, _height, _tile_width, _tile_height, _force_uniform_tiles):
|
||||
return [(0, 0)]
|
||||
|
||||
def _poll_job_ready(self, *_args, **_kwargs):
|
||||
return self._poll_ready
|
||||
|
||||
async def check_job_status(self, *_args, **_kwargs):
|
||||
return self._poll_ready
|
||||
|
||||
async def request_assignment(self, *_args, **_kwargs):
|
||||
if self._assignments:
|
||||
image_idx, estimated_remaining = self._assignments.pop(0)
|
||||
return SimpleNamespace(
|
||||
kind="image" if image_idx is not None else "none",
|
||||
task_idx=image_idx,
|
||||
estimated_remaining=estimated_remaining,
|
||||
batched_static=False,
|
||||
)
|
||||
return SimpleNamespace(kind="none", task_idx=None, estimated_remaining=0, batched_static=False)
|
||||
|
||||
def slice_conditioning(self, positive, negative, _image_idx):
|
||||
return positive, negative
|
||||
|
||||
def process_and_blend_tile(self, _tile_idx, _pos, _source_tensor, local_image, *_args, **_kwargs):
|
||||
return local_image
|
||||
|
||||
async def send_heartbeat(self, *_args, **_kwargs):
|
||||
self.heartbeat_count += 1
|
||||
|
||||
async def send_full_image(self, local_image, image_idx, _multi_job_id, _master_url, _worker_id, is_last):
|
||||
self.sent_images.append((image_idx, is_last, local_image.size))
|
||||
|
||||
async def send_worker_complete_signal(self, *_args, **_kwargs):
|
||||
self.completion_sent = True
|
||||
|
||||
|
||||
@unittest.skipIf(torch is None, "torch is not installed")
|
||||
class DynamicModeTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.image_batch = torch.zeros((1, 8, 8, 3), dtype=torch.float32)
|
||||
self.core_args = dynamic_module.UpscaleCoreArgs(
|
||||
model=None,
|
||||
positive=None,
|
||||
negative=None,
|
||||
vae=None,
|
||||
seed=1,
|
||||
steps=10,
|
||||
cfg=7.0,
|
||||
sampler_name="euler",
|
||||
scheduler="normal",
|
||||
denoise=0.5,
|
||||
tiled_decode=False,
|
||||
)
|
||||
|
||||
def test_worker_dynamic_returns_early_if_job_not_ready(self):
|
||||
node = _DummyDynamicNode(poll_ready=False)
|
||||
|
||||
result = node.process_worker_dynamic(
|
||||
upscaled_image=self.image_batch,
|
||||
core_args=self.core_args,
|
||||
tile_width=8,
|
||||
tile_height=8,
|
||||
padding=4,
|
||||
mask_blur=2,
|
||||
force_uniform_tiles=True,
|
||||
multi_job_id="job-1",
|
||||
master_url="http://master.local:8188",
|
||||
worker_id="worker-a",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
dynamic_threshold=8,
|
||||
)
|
||||
|
||||
self.assertEqual(result[0].shape, self.image_batch.shape)
|
||||
self.assertEqual(node.sent_images, [])
|
||||
self.assertFalse(node.completion_sent)
|
||||
|
||||
def test_worker_dynamic_processes_assigned_image_and_sends_completion(self):
|
||||
node = _DummyDynamicNode(poll_ready=True, assignments=[(0, 0)])
|
||||
|
||||
result = node.process_worker_dynamic(
|
||||
upscaled_image=self.image_batch,
|
||||
core_args=self.core_args,
|
||||
tile_width=8,
|
||||
tile_height=8,
|
||||
padding=4,
|
||||
mask_blur=2,
|
||||
force_uniform_tiles=True,
|
||||
multi_job_id="job-2",
|
||||
master_url="http://master.local:8188",
|
||||
worker_id="worker-a",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
dynamic_threshold=8,
|
||||
)
|
||||
|
||||
self.assertEqual(result[0].shape, self.image_batch.shape)
|
||||
self.assertEqual(len(node.sent_images), 1)
|
||||
self.assertEqual(node.sent_images[0][0], 0)
|
||||
self.assertTrue(node.sent_images[0][1])
|
||||
self.assertTrue(node.completion_sent)
|
||||
self.assertGreaterEqual(node.heartbeat_count, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,39 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from utils.image import ensure_contiguous, pil_to_tensor, tensor_to_pil
|
||||
|
||||
|
||||
class ImageUtilsTests(unittest.TestCase):
|
||||
def test_tensor_to_pil_converts_batch_item(self):
|
||||
tensor = torch.tensor(
|
||||
[[[[0.0, 0.5, 1.0], [1.0, 0.0, 0.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
image = tensor_to_pil(tensor, batch_index=0)
|
||||
|
||||
self.assertEqual(image.size, (2, 1))
|
||||
self.assertEqual(image.getpixel((0, 0)), (0, 127, 255))
|
||||
self.assertEqual(image.getpixel((1, 0)), (255, 0, 0))
|
||||
|
||||
def test_pil_to_tensor_adds_channel_for_grayscale(self):
|
||||
grayscale = Image.new("L", (3, 2), color=128)
|
||||
tensor = pil_to_tensor(grayscale)
|
||||
|
||||
self.assertEqual(tuple(tensor.shape), (1, 2, 3, 1))
|
||||
self.assertAlmostEqual(float(tensor[0, 0, 0, 0]), 128 / 255.0, places=6)
|
||||
|
||||
def test_ensure_contiguous_makes_non_contiguous_tensor_contiguous(self):
|
||||
base = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
|
||||
non_contiguous = base.transpose(1, 2)
|
||||
self.assertFalse(non_contiguous.is_contiguous())
|
||||
|
||||
contiguous = ensure_contiguous(non_contiguous)
|
||||
self.assertTrue(contiguous.is_contiguous())
|
||||
self.assertTrue(torch.equal(contiguous, non_contiguous))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,131 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_job_state_modules():
|
||||
package_name = "dist_job_state_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
distributed_pending_tile_jobs={},
|
||||
distributed_tile_jobs_lock=asyncio.Lock(),
|
||||
)
|
||||
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
|
||||
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
call_log = {"args": None}
|
||||
job_timeout_module = types.ModuleType(f"{package_name}.upscale.job_timeout")
|
||||
|
||||
async def _requeue(multi_job_id, batch_size):
|
||||
call_log["args"] = (multi_job_id, batch_size)
|
||||
return 3
|
||||
|
||||
job_timeout_module._check_and_requeue_timed_out_workers = _requeue
|
||||
sys.modules[f"{package_name}.upscale.job_timeout"] = job_timeout_module
|
||||
|
||||
job_state = _load_module(package_name, "upscale/job_state.py", "upscale.job_state")
|
||||
return job_models, job_state, prompt_server, call_log
|
||||
|
||||
|
||||
job_models_module, job_state_module, prompt_server, timeout_call_log = _load_job_state_modules()
|
||||
|
||||
|
||||
class _Node(job_state_module.JobStateMixin):
|
||||
pass
|
||||
|
||||
|
||||
class JobStateMixinTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
prompt_server.distributed_pending_tile_jobs = {}
|
||||
prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
self.node = _Node()
|
||||
|
||||
async def test_get_all_completed_tasks_supports_tile_and_image_modes(self):
|
||||
tile_state = job_models_module.TileJobState("tile-job")
|
||||
tile_state.completed_tasks[4] = "tile-4"
|
||||
image_state = job_models_module.ImageJobState("image-job")
|
||||
image_state.completed_images[2] = "image-2"
|
||||
prompt_server.distributed_pending_tile_jobs = {
|
||||
"tile-job": tile_state,
|
||||
"image-job": image_state,
|
||||
}
|
||||
|
||||
self.assertEqual(await self.node._get_all_completed_tasks("tile-job"), {4: "tile-4"})
|
||||
self.assertEqual(await self.node._get_all_completed_tasks("image-job"), {2: "image-2"})
|
||||
|
||||
async def test_next_index_and_pending_count_from_queues(self):
|
||||
image_state = job_models_module.ImageJobState("image-job")
|
||||
await image_state.pending_images.put(7)
|
||||
tile_state = job_models_module.TileJobState("tile-job")
|
||||
await tile_state.pending_tasks.put(9)
|
||||
prompt_server.distributed_pending_tile_jobs = {
|
||||
"image-job": image_state,
|
||||
"tile-job": tile_state,
|
||||
}
|
||||
|
||||
self.assertEqual(await self.node._get_next_image_index("image-job"), 7)
|
||||
self.assertEqual(await self.node._get_next_tile_index("tile-job"), 9)
|
||||
self.assertEqual(await self.node._get_pending_count("image-job"), 0)
|
||||
self.assertEqual(await self.node._get_pending_count("tile-job"), 0)
|
||||
|
||||
async def test_drain_worker_results_queue_collects_completed_images(self):
|
||||
image_state = job_models_module.ImageJobState("image-job")
|
||||
await image_state.queue.put({"worker_id": "w1", "image_idx": 1, "image": "img1"})
|
||||
await image_state.queue.put({"worker_id": "w2", "image_idx": 2, "image": "img2"})
|
||||
prompt_server.distributed_pending_tile_jobs = {"image-job": image_state}
|
||||
|
||||
drained = await self.node._drain_worker_results_queue("image-job")
|
||||
|
||||
self.assertEqual(drained, 2)
|
||||
self.assertEqual(image_state.completed_images[1], "img1")
|
||||
self.assertEqual(image_state.completed_images[2], "img2")
|
||||
|
||||
async def test_check_and_requeue_delegates_to_timeout_helper(self):
|
||||
count = await self.node._check_and_requeue_timed_out_workers("job-abc", 5)
|
||||
self.assertEqual(count, 3)
|
||||
self.assertEqual(timeout_call_log["args"], ("job-abc", 5))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,164 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
prompt_server = types.SimpleNamespace()
|
||||
|
||||
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
|
||||
|
||||
def _ensure_distributed_runtime_state(server_instance=None):
|
||||
ps = server_instance or prompt_server
|
||||
if not hasattr(ps, "distributed_pending_jobs"):
|
||||
ps.distributed_pending_jobs = {}
|
||||
if not hasattr(ps, "distributed_jobs_lock"):
|
||||
ps.distributed_jobs_lock = asyncio.Lock()
|
||||
if not hasattr(ps, "distributed_job_allowed_workers"):
|
||||
ps.distributed_job_allowed_workers = {}
|
||||
if not hasattr(ps, "distributed_pending_tile_jobs"):
|
||||
ps.distributed_pending_tile_jobs = {}
|
||||
if not hasattr(ps, "distributed_tile_jobs_lock"):
|
||||
ps.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
return types.SimpleNamespace(
|
||||
distributed_pending_jobs=ps.distributed_pending_jobs,
|
||||
distributed_jobs_lock=ps.distributed_jobs_lock,
|
||||
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
|
||||
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
|
||||
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
|
||||
)
|
||||
|
||||
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
|
||||
runtime_state_module.get_prompt_server_instance = lambda: prompt_server
|
||||
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_job_store_modules():
|
||||
package_name = "dist_job_store_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
|
||||
job_store = _load_module(package_name, "upscale/job_store.py", "upscale.job_store")
|
||||
return job_models, job_store
|
||||
|
||||
|
||||
job_models_module, job_store_module = _load_job_store_modules()
|
||||
|
||||
|
||||
class JobStoreTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
self.prompt_server = job_store_module.get_prompt_server_instance()
|
||||
self.prompt_server.distributed_pending_tile_jobs = {}
|
||||
self.prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
|
||||
async def test_ensure_tile_jobs_initialized_prunes_invalid_entries(self):
|
||||
self.prompt_server.distributed_pending_tile_jobs["valid"] = job_models_module.TileJobState("job-valid")
|
||||
self.prompt_server.distributed_pending_tile_jobs["invalid"] = {"not": "state"}
|
||||
|
||||
job_store_module.ensure_tile_jobs_initialized()
|
||||
|
||||
self.assertIn("valid", self.prompt_server.distributed_pending_tile_jobs)
|
||||
self.assertNotIn("invalid", self.prompt_server.distributed_pending_tile_jobs)
|
||||
|
||||
async def test_init_dynamic_job_populates_pending_images(self):
|
||||
await job_store_module.init_dynamic_job(
|
||||
multi_job_id="job-dynamic",
|
||||
batch_size=3,
|
||||
enabled_workers=["worker-a"],
|
||||
all_indices=[2, 0, 1],
|
||||
)
|
||||
|
||||
job_data = self.prompt_server.distributed_pending_tile_jobs["job-dynamic"]
|
||||
self.assertIsInstance(job_data, job_models_module.ImageJobState)
|
||||
self.assertEqual(job_data.batch_size, 3)
|
||||
self.assertIn("worker-a", job_data.worker_status)
|
||||
|
||||
pulled = [job_data.pending_images.get_nowait() for _ in range(3)]
|
||||
self.assertEqual(pulled, [2, 0, 1])
|
||||
|
||||
async def test_init_static_job_batched_queues_tile_ids(self):
|
||||
await job_store_module.init_static_job_batched(
|
||||
multi_job_id="job-static",
|
||||
batch_size=2,
|
||||
num_tiles_per_image=4,
|
||||
enabled_workers=["worker-a", "worker-b"],
|
||||
)
|
||||
|
||||
job_data = self.prompt_server.distributed_pending_tile_jobs["job-static"]
|
||||
self.assertIsInstance(job_data, job_models_module.TileJobState)
|
||||
self.assertTrue(job_data.batched_static)
|
||||
|
||||
pulled = [job_data.pending_tasks.get_nowait() for _ in range(4)]
|
||||
self.assertEqual(pulled, [0, 1, 2, 3])
|
||||
|
||||
async def test_drain_results_queue_collects_images_tiles_and_completion(self):
|
||||
job_data = job_models_module.TileJobState("job-drain")
|
||||
job_data.worker_status = {"worker-a": 1.0, "worker-b": 1.0}
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-drain"] = job_data
|
||||
|
||||
await job_data.queue.put(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"is_last": True,
|
||||
"image_idx": 5,
|
||||
"image": "img-5",
|
||||
}
|
||||
)
|
||||
await job_data.queue.put(
|
||||
{
|
||||
"worker_id": "worker-b",
|
||||
"is_last": False,
|
||||
"tiles": [
|
||||
{"tile_idx": 0, "global_idx": 0, "tensor": "tile-0"},
|
||||
{"tile_idx": 1, "global_idx": 1, "tensor": "tile-1"},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
drained = await job_store_module.drain_results_queue("job-drain")
|
||||
|
||||
self.assertEqual(drained, 3)
|
||||
self.assertEqual(job_data.completed_tasks[5], "img-5")
|
||||
self.assertEqual(job_data.completed_tasks[0]["tensor"], "tile-0")
|
||||
self.assertEqual(job_data.completed_tasks[1]["tensor"], "tile-1")
|
||||
self.assertNotIn("worker-a", job_data.worker_status)
|
||||
self.assertIn("worker-b", job_data.worker_status)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -9,7 +9,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_job_timeout_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "upscale" / "job_timeout.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "upscale" / "job_timeout.py"
|
||||
package_name = "dist_job_timeout_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -15,7 +15,7 @@ except ImportError:
|
||||
|
||||
def _load_payload_parsers_module():
|
||||
# payload_parsers.py has no relative imports; only stdlib + PIL
|
||||
module_path = Path(__file__).resolve().parents[1] / "upscale" / "payload_parsers.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "upscale" / "payload_parsers.py"
|
||||
spec = importlib.util.spec_from_file_location("upscale_payload_parsers", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
@@ -86,25 +86,27 @@ class ParseTilesFromFormTests(unittest.TestCase):
|
||||
|
||||
# --- happy paths ---
|
||||
|
||||
def test_single_tile_returns_image_and_metadata(self):
|
||||
def test_single_tile_returns_one_entry(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
self.assertEqual(len(tiles), 1)
|
||||
|
||||
def test_multiple_tiles_all_returned(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(len(tiles), 3)
|
||||
|
||||
def test_tile_image_is_pil_image(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
self.assertIsInstance(tiles[0]["image"], PILImage.Image)
|
||||
|
||||
def test_tile_metadata_fields_are_parsed(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
tile = tiles[0]
|
||||
self.assertIsInstance(tile["image"], PILImage.Image)
|
||||
self.assertEqual(tile["tile_idx"], 0)
|
||||
self.assertEqual(tile["x"], 0)
|
||||
self.assertEqual(tile["y"], 0)
|
||||
self.assertEqual(tile["extracted_width"], 64)
|
||||
self.assertEqual(tile["extracted_height"], 64)
|
||||
|
||||
def test_multiple_tiles_preserve_count_order_and_coordinates(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(len(tiles), 3)
|
||||
for i, tile in enumerate(tiles):
|
||||
self.assertEqual(tile["tile_idx"], i)
|
||||
self.assertEqual(tiles[1]["x"], 64)
|
||||
self.assertEqual(tiles[2]["x"], 128)
|
||||
|
||||
def test_padding_is_parsed_from_form(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1, padding=16))
|
||||
self.assertEqual(tiles[0]["padding"], 16)
|
||||
@@ -136,6 +138,16 @@ class ParseTilesFromFormTests(unittest.TestCase):
|
||||
self.assertNotIn("batch_idx", tiles[0])
|
||||
self.assertNotIn("global_idx", tiles[0])
|
||||
|
||||
def test_tile_indices_match_metadata_order(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
for i, tile in enumerate(tiles):
|
||||
self.assertEqual(tile["tile_idx"], i)
|
||||
|
||||
def test_x_coordinates_reflect_metadata(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(tiles[1]["x"], 64)
|
||||
self.assertEqual(tiles[2]["x"], 128)
|
||||
|
||||
# --- error cases ---
|
||||
|
||||
def test_missing_tiles_metadata_raises_value_error(self):
|
||||
@@ -0,0 +1,248 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import time
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 0.05
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.DYNAMIC_MODE_MAX_POLL_TIMEOUT = 10.0
|
||||
constants_module.HEARTBEAT_INTERVAL = 0.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
model_mgmt = types.ModuleType("comfy.model_management")
|
||||
|
||||
class _InterruptProcessingException(Exception):
|
||||
pass
|
||||
|
||||
model_mgmt.processing_interrupted = lambda: False
|
||||
model_mgmt.InterruptProcessingException = _InterruptProcessingException
|
||||
comfy_module.model_management = model_mgmt
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.model_management"] = model_mgmt
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
distributed_pending_tile_jobs={},
|
||||
distributed_tile_jobs_lock=asyncio.Lock(),
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_result_collector_modules():
|
||||
package_name = "dist_result_collector_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
|
||||
job_store = _load_module(package_name, "upscale/job_store.py", "upscale.job_store")
|
||||
|
||||
job_timeout = types.ModuleType(f"{package_name}.upscale.job_timeout")
|
||||
|
||||
async def _noop_requeue(_multi_job_id, _batch_size):
|
||||
return 0
|
||||
|
||||
job_timeout._check_and_requeue_timed_out_workers = _noop_requeue
|
||||
job_timeout.check_and_requeue_timed_out_workers = _noop_requeue
|
||||
sys.modules[f"{package_name}.upscale.job_timeout"] = job_timeout
|
||||
|
||||
result_collector = _load_module(
|
||||
package_name,
|
||||
"upscale/result_collector.py",
|
||||
"upscale.result_collector",
|
||||
)
|
||||
return job_models, job_store, result_collector
|
||||
|
||||
|
||||
job_models_module, job_store_module, result_collector_module = _load_result_collector_modules()
|
||||
|
||||
|
||||
class _Node(result_collector_module.ResultCollectorMixin):
|
||||
def __init__(self):
|
||||
self.requeue_calls = []
|
||||
|
||||
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
|
||||
self.requeue_calls.append((multi_job_id, batch_size))
|
||||
return 1
|
||||
|
||||
|
||||
class ResultCollectorMixinTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
self.prompt_server = job_store_module.server.PromptServer.instance
|
||||
self.prompt_server.distributed_pending_tile_jobs = {}
|
||||
self.prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
self.node = _Node()
|
||||
|
||||
async def test_static_collection_collects_tiles_and_cleans_up_job(self):
|
||||
job_data = job_models_module.TileJobState("job-static")
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-static"] = job_data
|
||||
await job_data.queue.put(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"is_last": True,
|
||||
"tiles": [
|
||||
{"tile_idx": 3},
|
||||
{
|
||||
"tile_idx": 2,
|
||||
"x": 10,
|
||||
"y": 20,
|
||||
"extracted_width": 64,
|
||||
"extracted_height": 64,
|
||||
"padding": 8,
|
||||
"batch_idx": 1,
|
||||
"global_idx": 9,
|
||||
"tensor": "tile-9",
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
collected = await self.node._async_collect_worker_tiles(
|
||||
"job-static",
|
||||
num_workers=1,
|
||||
)
|
||||
|
||||
self.assertEqual(list(collected.keys()), [9])
|
||||
self.assertEqual(collected[9]["tensor"], "tile-9")
|
||||
self.assertEqual(collected[9]["worker_id"], "worker-a")
|
||||
self.assertNotIn("job-static", self.prompt_server.distributed_pending_tile_jobs)
|
||||
|
||||
async def test_dynamic_collection_stops_at_remaining_target_and_cleans_up(self):
|
||||
job_data = job_models_module.ImageJobState("job-dynamic")
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-dynamic"] = job_data
|
||||
await job_data.queue.put(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"is_last": False,
|
||||
"image_idx": 2,
|
||||
"image": "img-2",
|
||||
}
|
||||
)
|
||||
|
||||
completed = await self.node.collect_dynamic_images(
|
||||
"job-dynamic",
|
||||
remaining_to_collect=1,
|
||||
num_workers=2,
|
||||
batch_size=3,
|
||||
master_processed_count=0,
|
||||
)
|
||||
|
||||
self.assertEqual(completed[2], "img-2")
|
||||
self.assertNotIn("job-dynamic", self.prompt_server.distributed_pending_tile_jobs)
|
||||
|
||||
async def test_dynamic_timeout_calls_requeue_helper(self):
|
||||
job_data = job_models_module.ImageJobState("job-timeout")
|
||||
job_data.worker_status = {"worker-a": time.time() - 5}
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-timeout"] = job_data
|
||||
|
||||
old_timeout = result_collector_module.get_worker_timeout_seconds
|
||||
result_collector_module.get_worker_timeout_seconds = lambda: 0.01
|
||||
try:
|
||||
completed = await self.node.collect_dynamic_images(
|
||||
"job-timeout",
|
||||
remaining_to_collect=None,
|
||||
num_workers=1,
|
||||
batch_size=4,
|
||||
master_processed_count=0,
|
||||
)
|
||||
finally:
|
||||
result_collector_module.get_worker_timeout_seconds = old_timeout
|
||||
|
||||
self.assertEqual(completed, {})
|
||||
self.assertEqual(self.node.requeue_calls, [("job-timeout", 4)])
|
||||
self.assertNotIn("job-timeout", self.prompt_server.distributed_pending_tile_jobs)
|
||||
|
||||
async def test_log_worker_timeout_status_handles_non_job_state(self):
|
||||
worker_ids = self.node._log_worker_timeout_status(
|
||||
{"worker_status": {"worker-a": time.time()}},
|
||||
current_time=time.time(),
|
||||
multi_job_id="job-invalid",
|
||||
)
|
||||
self.assertEqual(worker_ids, [])
|
||||
|
||||
async def test_log_worker_timeout_status_returns_worker_ids_for_job_state(self):
|
||||
job_data = job_models_module.ImageJobState("job-state")
|
||||
job_data.worker_status = {"worker-a": time.time() - 1}
|
||||
|
||||
worker_ids = self.node._log_worker_timeout_status(
|
||||
job_data,
|
||||
current_time=time.time(),
|
||||
multi_job_id="job-state",
|
||||
)
|
||||
self.assertEqual(worker_ids, ["worker-a"])
|
||||
|
||||
async def test_static_mode_mismatch_raises(self):
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-mismatch"] = (
|
||||
job_models_module.ImageJobState("job-mismatch")
|
||||
)
|
||||
with self.assertRaises(RuntimeError):
|
||||
await self.node._async_collect_worker_tiles(
|
||||
"job-mismatch",
|
||||
num_workers=1,
|
||||
)
|
||||
|
||||
async def test_dynamic_mode_mismatch_raises(self):
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-mismatch"] = (
|
||||
job_models_module.TileJobState("job-mismatch")
|
||||
)
|
||||
with self.assertRaises(RuntimeError):
|
||||
await self.node.collect_dynamic_images(
|
||||
"job-mismatch",
|
||||
remaining_to_collect=1,
|
||||
num_workers=1,
|
||||
batch_size=1,
|
||||
master_processed_count=0,
|
||||
)
|
||||
|
||||
async def test_mark_image_completed_updates_job_state(self):
|
||||
job_data = job_models_module.ImageJobState("job-mark")
|
||||
self.prompt_server.distributed_pending_tile_jobs["job-mark"] = job_data
|
||||
|
||||
await self.node.mark_image_completed("job-mark", 7, "img-7")
|
||||
|
||||
self.assertEqual(job_data.completed_images[7], "img-7")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,207 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
try:
|
||||
import torch
|
||||
except ModuleNotFoundError: # pragma: no cover - optional dependency in CI/runtime
|
||||
torch = None
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
|
||||
modes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
|
||||
def tensor_to_pil(image_batch, index):
|
||||
tensor = image_batch[index].detach().cpu().clamp(0, 1)
|
||||
array = (tensor.numpy() * 255.0).astype("uint8")
|
||||
return Image.fromarray(array)
|
||||
|
||||
def pil_to_tensor(image):
|
||||
if torch is None:
|
||||
raise RuntimeError("torch is required for this test helper")
|
||||
array = np.array(image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array).unsqueeze(0)
|
||||
|
||||
def blend_processed_batch_item(
|
||||
result_images,
|
||||
processed_batch,
|
||||
batch_index,
|
||||
blend_fn,
|
||||
x1,
|
||||
y1,
|
||||
ew,
|
||||
eh,
|
||||
tile_mask,
|
||||
padding,
|
||||
):
|
||||
tile_pil = tensor_to_pil(processed_batch, batch_index)
|
||||
if tile_pil.size != (ew, eh):
|
||||
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
|
||||
result_images[batch_index] = blend_fn(
|
||||
result_images[batch_index],
|
||||
tile_pil,
|
||||
x1,
|
||||
y1,
|
||||
(ew, eh),
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
|
||||
image_module.tensor_to_pil = tensor_to_pil
|
||||
image_module.pil_to_tensor = pil_to_tensor
|
||||
image_module.blend_processed_batch_item = blend_processed_batch_item
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
if torch is None:
|
||||
torch_module = types.ModuleType("torch")
|
||||
torch_module.cat = lambda *_args, **_kwargs: None
|
||||
sys.modules["torch"] = torch_module
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_single_gpu_mode_module():
|
||||
package_name = "dist_single_gpu_mode_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
|
||||
_load_module(package_name, "upscale/tile_processing.py", "upscale.tile_processing")
|
||||
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
|
||||
return _load_module(package_name, "upscale/modes/single_gpu.py", "upscale.modes.single_gpu")
|
||||
|
||||
|
||||
single_gpu_module = _load_single_gpu_mode_module()
|
||||
|
||||
|
||||
class _DummySingleGpuNode(single_gpu_module.SingleGpuModeMixin):
|
||||
def __init__(self):
|
||||
self.round_calls = []
|
||||
self.process_batch_calls = 0
|
||||
|
||||
def round_to_multiple(self, value):
|
||||
self.round_calls.append(int(value))
|
||||
return int(value)
|
||||
|
||||
def calculate_tiles(self, _width, _height, _tile_width, _tile_height, _force_uniform_tiles):
|
||||
return [(0, 0)]
|
||||
|
||||
def create_tile_mask(self, _width, _height, _tx, _ty, _tile_width, _tile_height, _mask_blur):
|
||||
return None
|
||||
|
||||
def extract_batch_tile_with_padding(self, source_batch, _tx, _ty, _tile_width, _tile_height, _padding, _force_uniform_tiles):
|
||||
_, h, w, _ = source_batch.shape
|
||||
return source_batch, 0, 0, w, h
|
||||
|
||||
def process_tiles_batch(
|
||||
self,
|
||||
tile_batch,
|
||||
_model,
|
||||
_positive,
|
||||
_negative,
|
||||
_vae,
|
||||
_seed,
|
||||
_steps,
|
||||
_cfg,
|
||||
_sampler_name,
|
||||
_scheduler,
|
||||
_denoise,
|
||||
_tiled_decode,
|
||||
_region,
|
||||
_canvas_shape,
|
||||
):
|
||||
self.process_batch_calls += 1
|
||||
return tile_batch
|
||||
|
||||
def blend_tile(self, _base_image, tile_pil, _x1, _y1, _size, _tile_mask, _padding):
|
||||
return tile_pil
|
||||
|
||||
|
||||
@unittest.skipIf(torch is None, "torch is not installed")
|
||||
class SingleGpuModeTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.node = _DummySingleGpuNode()
|
||||
self.images = torch.rand((2, 8, 8, 3), dtype=torch.float32)
|
||||
self.core_args = single_gpu_module.UpscaleCoreArgs(
|
||||
model=None,
|
||||
positive=None,
|
||||
negative=None,
|
||||
vae=None,
|
||||
seed=1,
|
||||
steps=10,
|
||||
cfg=7.5,
|
||||
sampler_name="euler",
|
||||
scheduler="normal",
|
||||
denoise=0.5,
|
||||
tiled_decode=False,
|
||||
)
|
||||
|
||||
def test_process_single_gpu_preserves_batch_shape(self):
|
||||
result = self.node.process_single_gpu(
|
||||
upscaled_image=self.images,
|
||||
core_args=self.core_args,
|
||||
tile_width=8,
|
||||
tile_height=8,
|
||||
padding=4,
|
||||
mask_blur=2,
|
||||
force_uniform_tiles=True,
|
||||
)[0]
|
||||
|
||||
self.assertEqual(result.shape, self.images.shape)
|
||||
self.assertTrue(torch.allclose(result, self.images, atol=(1.5 / 255.0)))
|
||||
|
||||
def test_process_single_gpu_invokes_rounding_and_batch_processing(self):
|
||||
self.node.process_single_gpu(
|
||||
upscaled_image=self.images,
|
||||
core_args=self.core_args,
|
||||
tile_width=16,
|
||||
tile_height=24,
|
||||
padding=4,
|
||||
mask_blur=2,
|
||||
force_uniform_tiles=True,
|
||||
)
|
||||
|
||||
self.assertEqual(self.node.round_calls[:2], [16, 24])
|
||||
self.assertEqual(self.node.process_batch_calls, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -4,12 +4,13 @@ import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _load_static_mode_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "upscale" / "modes" / "static.py"
|
||||
module_path = Path(__file__).resolve().parents[3] / "upscale" / "modes" / "static.py"
|
||||
package_name = "dist_static_mode_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -65,8 +66,34 @@ def _load_static_mode_module():
|
||||
arr = np.array(image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(arr).unsqueeze(0)
|
||||
|
||||
def _blend_processed_batch_item(
|
||||
result_images,
|
||||
processed_batch,
|
||||
batch_index,
|
||||
blend_fn,
|
||||
x1,
|
||||
y1,
|
||||
ew,
|
||||
eh,
|
||||
tile_mask,
|
||||
padding,
|
||||
):
|
||||
tile_pil = _tensor_to_pil(processed_batch, batch_index)
|
||||
if tile_pil.size != (ew, eh):
|
||||
tile_pil = tile_pil.resize((ew, eh), PILImage.LANCZOS)
|
||||
result_images[batch_index] = blend_fn(
|
||||
result_images[batch_index],
|
||||
tile_pil,
|
||||
x1,
|
||||
y1,
|
||||
(ew, eh),
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
|
||||
image_module.tensor_to_pil = _tensor_to_pil
|
||||
image_module.pil_to_tensor = _pil_to_tensor
|
||||
image_module.blend_processed_batch_item = _blend_processed_batch_item
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
@@ -102,10 +129,10 @@ def _load_static_mode_module():
|
||||
distributed_pending_tile_jobs={},
|
||||
)
|
||||
job_store_module.init_static_job_batched = _noop
|
||||
job_store_module._mark_task_completed = _noop
|
||||
job_store_module._cleanup_job = _noop
|
||||
job_store_module._drain_results_queue = _noop
|
||||
job_store_module._get_completed_count = _noop
|
||||
job_store_module.mark_task_completed = _noop
|
||||
job_store_module.cleanup_job = _noop
|
||||
job_store_module.drain_results_queue = _noop
|
||||
job_store_module.get_completed_count = _noop
|
||||
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
|
||||
|
||||
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
|
||||
@@ -116,6 +143,36 @@ def _load_static_mode_module():
|
||||
job_models_module.TileJobState = _TileJobState
|
||||
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
|
||||
|
||||
processing_args_path = Path(__file__).resolve().parents[3] / "upscale" / "processing_args.py"
|
||||
processing_args_spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.upscale.processing_args",
|
||||
processing_args_path,
|
||||
)
|
||||
processing_args_module = importlib.util.module_from_spec(processing_args_spec)
|
||||
assert processing_args_spec is not None and processing_args_spec.loader is not None
|
||||
sys.modules[processing_args_spec.name] = processing_args_module
|
||||
processing_args_spec.loader.exec_module(processing_args_module)
|
||||
|
||||
tile_processing_path = Path(__file__).resolve().parents[3] / "upscale" / "tile_processing.py"
|
||||
tile_processing_spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.upscale.tile_processing",
|
||||
tile_processing_path,
|
||||
)
|
||||
tile_processing_module = importlib.util.module_from_spec(tile_processing_spec)
|
||||
assert tile_processing_spec is not None and tile_processing_spec.loader is not None
|
||||
sys.modules[tile_processing_spec.name] = tile_processing_module
|
||||
tile_processing_spec.loader.exec_module(tile_processing_module)
|
||||
|
||||
mode_contexts_path = Path(__file__).resolve().parents[3] / "upscale" / "mode_contexts.py"
|
||||
mode_contexts_spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.upscale.mode_contexts",
|
||||
mode_contexts_path,
|
||||
)
|
||||
mode_contexts_module = importlib.util.module_from_spec(mode_contexts_spec)
|
||||
assert mode_contexts_spec is not None and mode_contexts_spec.loader is not None
|
||||
sys.modules[mode_contexts_spec.name] = mode_contexts_module
|
||||
mode_contexts_spec.loader.exec_module(mode_contexts_module)
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.upscale.modes.static", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
@@ -148,14 +205,20 @@ class _FakeStaticWorker(static_mode.StaticModeMixin):
|
||||
def _poll_job_ready(self, *_args, **_kwargs):
|
||||
return self.job_ready
|
||||
|
||||
async def _request_tile_from_master(self, *_args, **_kwargs):
|
||||
async def request_assignment(self, *_args, **_kwargs):
|
||||
self.request_calls += 1
|
||||
return self.tile_sequence.pop(0)
|
||||
tile_idx, estimated_remaining, batched_static = self.tile_sequence.pop(0)
|
||||
return SimpleNamespace(
|
||||
kind="tile" if tile_idx is not None else "none",
|
||||
task_idx=tile_idx,
|
||||
estimated_remaining=estimated_remaining,
|
||||
batched_static=batched_static,
|
||||
)
|
||||
|
||||
async def _send_heartbeat_to_master(self, *_args, **_kwargs):
|
||||
async def send_heartbeat(self, *_args, **_kwargs):
|
||||
self.heartbeat_calls += 1
|
||||
|
||||
async def send_tiles_batch_to_master(
|
||||
async def send_tiles_batch(
|
||||
self,
|
||||
processed_tiles,
|
||||
_multi_job_id,
|
||||
@@ -0,0 +1,179 @@
|
||||
import copy
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from contextlib import nullcontext
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _bootstrap_package(package_name):
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
|
||||
def tensor_to_pil(image_batch, index):
|
||||
tensor = image_batch[index].detach().cpu().clamp(0, 1)
|
||||
array = (tensor.numpy() * 255.0).astype("uint8")
|
||||
return Image.fromarray(array)
|
||||
|
||||
def pil_to_tensor(image):
|
||||
array = np.array(image).astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array).unsqueeze(0)
|
||||
|
||||
image_module.tensor_to_pil = tensor_to_pil
|
||||
image_module.pil_to_tensor = pil_to_tensor
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
usdu_utils_module = types.ModuleType(f"{package_name}.utils.usdu_utils")
|
||||
|
||||
def crop_cond(cond, *_args, **_kwargs):
|
||||
return cond
|
||||
|
||||
def get_crop_region(mask, padding):
|
||||
bbox = mask.getbbox()
|
||||
if bbox is None:
|
||||
return (0, 0, mask.size[0], mask.size[1])
|
||||
x1, y1, x2, y2 = bbox
|
||||
return (
|
||||
max(0, x1 - int(padding)),
|
||||
max(0, y1 - int(padding)),
|
||||
min(mask.size[0], x2 + int(padding)),
|
||||
min(mask.size[1], y2 + int(padding)),
|
||||
)
|
||||
|
||||
def expand_crop(region, width, height, target_w, target_h):
|
||||
x1, y1, x2, y2 = region
|
||||
crop_w = x2 - x1
|
||||
crop_h = y2 - y1
|
||||
if crop_w >= target_w and crop_h >= target_h:
|
||||
return (region, (crop_w, crop_h))
|
||||
new_x2 = min(width, x1 + max(crop_w, target_w))
|
||||
new_y2 = min(height, y1 + max(crop_h, target_h))
|
||||
return ((x1, y1, new_x2, new_y2), (new_x2 - x1, new_y2 - y1))
|
||||
|
||||
usdu_utils_module.crop_cond = crop_cond
|
||||
usdu_utils_module.get_crop_region = get_crop_region
|
||||
usdu_utils_module.expand_crop = expand_crop
|
||||
sys.modules[f"{package_name}.utils.usdu_utils"] = usdu_utils_module
|
||||
|
||||
crop_model_patch_module = types.ModuleType(f"{package_name}.utils.crop_model_patch")
|
||||
crop_model_patch_module.crop_model_cond = lambda model, *_args, **_kwargs: nullcontext(model)
|
||||
sys.modules[f"{package_name}.utils.crop_model_patch"] = crop_model_patch_module
|
||||
|
||||
conditioning_module = types.ModuleType(f"{package_name}.upscale.conditioning")
|
||||
conditioning_module.clone_conditioning = lambda cond, clone_hints=False: copy.deepcopy(cond)
|
||||
sys.modules[f"{package_name}.upscale.conditioning"] = conditioning_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
comfy_samplers = types.ModuleType("comfy.samplers")
|
||||
comfy_model_management = types.ModuleType("comfy.model_management")
|
||||
comfy_module.samplers = comfy_samplers
|
||||
comfy_module.model_management = comfy_model_management
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.samplers"] = comfy_samplers
|
||||
sys.modules["comfy.model_management"] = comfy_model_management
|
||||
|
||||
|
||||
def _load_module(package_name, module_rel_path, module_name):
|
||||
module_path = Path(__file__).resolve().parents[3] / module_rel_path
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.{module_name}",
|
||||
module_path,
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _load_tile_ops_module():
|
||||
package_name = "dist_tile_ops_testpkg"
|
||||
_bootstrap_package(package_name)
|
||||
return _load_module(package_name, "upscale/tile_ops.py", "upscale.tile_ops")
|
||||
|
||||
|
||||
tile_ops_module = _load_tile_ops_module()
|
||||
|
||||
|
||||
class _DummyTileOps(tile_ops_module.TileOpsMixin):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeControl:
|
||||
def __init__(self, cond_hint_original, previous_controlnet=None):
|
||||
self.cond_hint_original = cond_hint_original
|
||||
self.previous_controlnet = previous_controlnet
|
||||
|
||||
|
||||
class TileOpsTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.node = _DummyTileOps()
|
||||
|
||||
def test_round_and_calculate_tiles(self):
|
||||
self.assertEqual(self.node.round_to_multiple(13, 8), 16)
|
||||
tiles = self.node.calculate_tiles(10, 9, 4, 5, force_uniform_tiles=True)
|
||||
self.assertEqual(len(tiles), 6)
|
||||
self.assertEqual(tiles[0], (0, 0))
|
||||
self.assertEqual(tiles[-1], (8, 5))
|
||||
|
||||
def test_create_tile_mask_and_blend_tile(self):
|
||||
base = Image.new("RGB", (8, 8), (0, 0, 0))
|
||||
tile = Image.new("RGB", (4, 4), (255, 255, 255))
|
||||
mask = self.node.create_tile_mask(8, 8, 2, 2, 4, 4, mask_blur=0)
|
||||
|
||||
self.assertEqual(mask.getpixel((0, 0)), 0)
|
||||
self.assertEqual(mask.getpixel((3, 3)), 255)
|
||||
|
||||
blended = self.node.blend_tile(
|
||||
base_image=base,
|
||||
tile_image=tile,
|
||||
x=2,
|
||||
y=2,
|
||||
extracted_size=(4, 4),
|
||||
mask=mask,
|
||||
padding=0,
|
||||
)
|
||||
self.assertEqual(blended.size, (8, 8))
|
||||
self.assertEqual(blended.getpixel((3, 3)), (255, 255, 255))
|
||||
|
||||
def test_slice_conditioning_selects_batch_index(self):
|
||||
control = _FakeControl(cond_hint_original=torch.ones((2, 1, 1)))
|
||||
positive = [[torch.arange(6, dtype=torch.float32).reshape(2, 3), {"control": control, "mask": torch.ones((2, 1, 1))}]]
|
||||
negative = [[torch.arange(6, dtype=torch.float32).reshape(2, 3), {"mask": torch.zeros((2, 1, 1))}]]
|
||||
|
||||
pos_sliced, neg_sliced = self.node._slice_conditioning(positive, negative, batch_idx=1)
|
||||
|
||||
self.assertEqual(tuple(pos_sliced[0][0].shape), (1, 3))
|
||||
self.assertEqual(tuple(neg_sliced[0][0].shape), (1, 3))
|
||||
self.assertEqual(tuple(pos_sliced[0][1]["mask"].shape), (1, 1, 1))
|
||||
self.assertEqual(tuple(neg_sliced[0][1]["mask"].shape), (1, 1, 1))
|
||||
self.assertEqual(tuple(pos_sliced[0][1]["control"].cond_hint_original.shape), (1, 1, 1))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,115 @@
|
||||
import unittest
|
||||
|
||||
from upscale.processing_args import UpscaleCoreArgs
|
||||
from upscale.tile_processing import TileBatchArgs, extract_and_process_tile_batch
|
||||
|
||||
|
||||
class _DummyTileProcessor:
|
||||
def __init__(self):
|
||||
self.extract_calls = []
|
||||
self.process_calls = []
|
||||
|
||||
def extract_batch_tile_with_padding(
|
||||
self,
|
||||
upscaled_image,
|
||||
tx,
|
||||
ty,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
):
|
||||
self.extract_calls.append(
|
||||
{
|
||||
"upscaled_image": upscaled_image,
|
||||
"tx": tx,
|
||||
"ty": ty,
|
||||
"tile_width": tile_width,
|
||||
"tile_height": tile_height,
|
||||
"padding": padding,
|
||||
"force_uniform_tiles": force_uniform_tiles,
|
||||
}
|
||||
)
|
||||
return "tile-batch", 10, 20, 30, 40
|
||||
|
||||
def process_tiles_batch(
|
||||
self,
|
||||
tile_batch,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
region,
|
||||
canvas_shape,
|
||||
):
|
||||
self.process_calls.append(
|
||||
{
|
||||
"tile_batch": tile_batch,
|
||||
"model": model,
|
||||
"positive": positive,
|
||||
"negative": negative,
|
||||
"vae": vae,
|
||||
"seed": seed,
|
||||
"steps": steps,
|
||||
"cfg": cfg,
|
||||
"sampler_name": sampler_name,
|
||||
"scheduler": scheduler,
|
||||
"denoise": denoise,
|
||||
"tiled_decode": tiled_decode,
|
||||
"region": region,
|
||||
"canvas_shape": canvas_shape,
|
||||
}
|
||||
)
|
||||
return "processed-batch"
|
||||
|
||||
|
||||
class TileProcessingTests(unittest.TestCase):
|
||||
def test_extract_and_process_tile_batch_forwards_args_and_computes_region(self):
|
||||
node = _DummyTileProcessor()
|
||||
tile_args = TileBatchArgs(
|
||||
core=UpscaleCoreArgs(
|
||||
model="model",
|
||||
positive="positive",
|
||||
negative="negative",
|
||||
vae="vae",
|
||||
seed=42,
|
||||
steps=20,
|
||||
cfg=7.5,
|
||||
sampler_name="euler",
|
||||
scheduler="normal",
|
||||
denoise=0.4,
|
||||
tiled_decode=False,
|
||||
),
|
||||
tile_width=256,
|
||||
tile_height=192,
|
||||
padding=32,
|
||||
force_uniform_tiles=True,
|
||||
width=1024,
|
||||
height=768,
|
||||
)
|
||||
result = extract_and_process_tile_batch(
|
||||
node=node,
|
||||
upscaled_image="upscaled-image",
|
||||
tx=3,
|
||||
ty=4,
|
||||
args=tile_args,
|
||||
)
|
||||
|
||||
self.assertEqual(result, ("processed-batch", 10, 20, 30, 40))
|
||||
self.assertEqual(len(node.extract_calls), 1)
|
||||
self.assertEqual(len(node.process_calls), 1)
|
||||
self.assertEqual(node.extract_calls[0]["tx"], 3)
|
||||
self.assertEqual(node.extract_calls[0]["ty"], 4)
|
||||
self.assertEqual(node.process_calls[0]["region"], (10, 20, 40, 60))
|
||||
self.assertEqual(node.process_calls[0]["canvas_shape"], (1024, 768))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,128 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _ensure_torch_stubs():
|
||||
if "torch" not in sys.modules:
|
||||
torch_module = types.ModuleType("torch")
|
||||
|
||||
class _Tensor:
|
||||
pass
|
||||
|
||||
torch_module.Tensor = _Tensor
|
||||
torch_module.float32 = object()
|
||||
torch_module.zeros = lambda *_args, **_kwargs: None
|
||||
torch_module.from_numpy = lambda array: array
|
||||
torch_module.cat = lambda *_args, **_kwargs: None
|
||||
|
||||
nn_module = types.ModuleType("torch.nn")
|
||||
functional_module = types.ModuleType("torch.nn.functional")
|
||||
functional_module.interpolate = lambda tensor, **_kwargs: tensor
|
||||
nn_module.functional = functional_module
|
||||
|
||||
torch_module.nn = nn_module
|
||||
|
||||
sys.modules["torch"] = torch_module
|
||||
sys.modules["torch.nn"] = nn_module
|
||||
sys.modules["torch.nn.functional"] = functional_module
|
||||
|
||||
if "torchvision" not in sys.modules:
|
||||
torchvision_module = types.ModuleType("torchvision")
|
||||
transforms_module = types.ModuleType("torchvision.transforms")
|
||||
|
||||
class _GaussianBlur:
|
||||
def __init__(self, *_args, **_kwargs):
|
||||
pass
|
||||
|
||||
def __call__(self, value):
|
||||
return value
|
||||
|
||||
transforms_module.GaussianBlur = _GaussianBlur
|
||||
torchvision_module.transforms = transforms_module
|
||||
|
||||
sys.modules["torchvision"] = torchvision_module
|
||||
sys.modules["torchvision.transforms"] = transforms_module
|
||||
|
||||
|
||||
def _load_usdu_utils_module():
|
||||
_ensure_torch_stubs()
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "usdu_utils.py"
|
||||
spec = importlib.util.spec_from_file_location("dist_test_usdu_utils", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
usdu_utils = _load_usdu_utils_module()
|
||||
|
||||
|
||||
class UsduUtilsGeometryTests(unittest.TestCase):
|
||||
def test_fix_crop_region_decrements_upper_bounds_when_interior(self):
|
||||
self.assertEqual(
|
||||
usdu_utils.fix_crop_region((2, 3, 9, 9), (10, 10)),
|
||||
(2, 3, 8, 8),
|
||||
)
|
||||
self.assertEqual(
|
||||
usdu_utils.fix_crop_region((1, 1, 10, 10), (10, 10)),
|
||||
(1, 1, 10, 10),
|
||||
)
|
||||
|
||||
def test_get_crop_region_with_padding_and_empty_mask(self):
|
||||
mask = Image.new("L", (10, 10), 0)
|
||||
for x in range(2, 5):
|
||||
for y in range(3, 7):
|
||||
mask.putpixel((x, y), 255)
|
||||
|
||||
region = usdu_utils.get_crop_region(mask, pad=1)
|
||||
self.assertEqual(region, (1, 2, 5, 7))
|
||||
|
||||
empty_region = usdu_utils.get_crop_region(Image.new("L", (8, 6), 0), pad=2)
|
||||
self.assertEqual(empty_region, (6, 4, 1, 1))
|
||||
|
||||
def test_expand_crop_and_resize_region(self):
|
||||
expanded_region, expanded_size = usdu_utils.expand_crop(
|
||||
region=(2, 2, 6, 6),
|
||||
width=12,
|
||||
height=12,
|
||||
target_width=8,
|
||||
target_height=10,
|
||||
)
|
||||
self.assertEqual(expanded_region, (0, 0, 8, 10))
|
||||
self.assertEqual(expanded_size, (8, 10))
|
||||
|
||||
resized = usdu_utils.resize_region(
|
||||
region=(2, 2, 6, 10),
|
||||
init_size=(8, 16),
|
||||
resize_size=(16, 32),
|
||||
)
|
||||
self.assertEqual(resized, (4, 4, 12, 20))
|
||||
|
||||
def test_region_intersection(self):
|
||||
self.assertEqual(
|
||||
usdu_utils.region_intersection((0, 0, 5, 5), (3, 2, 8, 7)),
|
||||
(3, 2, 5, 5),
|
||||
)
|
||||
self.assertIsNone(usdu_utils.region_intersection((0, 0, 1, 1), (2, 2, 4, 4)))
|
||||
|
||||
def test_pad_image2_preserves_expected_output_size(self):
|
||||
image = Image.new("RGB", (4, 3), "black")
|
||||
padded = usdu_utils.pad_image2(
|
||||
image,
|
||||
left_pad=2,
|
||||
right_pad=1,
|
||||
top_pad=3,
|
||||
bottom_pad=2,
|
||||
fill=False,
|
||||
blur=False,
|
||||
)
|
||||
self.assertEqual(padded.size, (7, 8))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,7 +1,8 @@
|
||||
import copy
|
||||
from typing import Any
|
||||
|
||||
|
||||
def clone_control_chain(control, clone_hint=True):
|
||||
def clone_control_chain(control: Any, clone_hint: bool = True) -> Any:
|
||||
"""Shallow copy the ControlNet chain, optionally cloning hints but sharing models."""
|
||||
if control is None:
|
||||
return None
|
||||
@@ -14,7 +15,10 @@ def clone_control_chain(control, clone_hint=True):
|
||||
return new_control
|
||||
|
||||
|
||||
def clone_conditioning(cond_list, clone_hints=True):
|
||||
def clone_conditioning(
|
||||
cond_list: list[tuple[Any, dict[str, Any]]] | list[list[Any]],
|
||||
clone_hints: bool = True,
|
||||
) -> list[list[Any]]:
|
||||
"""Clone conditioning without duplicating ControlNet models."""
|
||||
new_cond = []
|
||||
for emb, cond_dict in cond_list:
|
||||
|
||||
+34
-2
@@ -2,7 +2,7 @@ import asyncio
|
||||
|
||||
from ..utils.logging import debug_log
|
||||
from .job_store import ensure_tile_jobs_initialized
|
||||
from .job_timeout import _check_and_requeue_timed_out_workers as _requeue_usdu
|
||||
from .job_timeout import check_and_requeue_timed_out_workers
|
||||
from .job_models import ImageJobState, TileJobState
|
||||
|
||||
|
||||
@@ -22,6 +22,10 @@ class JobStateMixin:
|
||||
return dict(job_data.completed_images)
|
||||
return {}
|
||||
|
||||
async def all_completed_tasks(self, multi_job_id):
|
||||
"""Public completed-task accessor."""
|
||||
return await self._get_all_completed_tasks(multi_job_id)
|
||||
|
||||
async def _get_next_image_index(self, multi_job_id):
|
||||
"""Get next image index from pending queue for master."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -39,6 +43,10 @@ class JobStateMixin:
|
||||
except asyncio.TimeoutError:
|
||||
return None
|
||||
|
||||
async def next_image_index(self, multi_job_id):
|
||||
"""Public next-image accessor."""
|
||||
return await self._get_next_image_index(multi_job_id)
|
||||
|
||||
async def _get_next_tile_index(self, multi_job_id):
|
||||
"""Get next tile index from pending queue for master in static mode."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -56,6 +64,10 @@ class JobStateMixin:
|
||||
except asyncio.TimeoutError:
|
||||
return None
|
||||
|
||||
async def next_tile_index(self, multi_job_id):
|
||||
"""Public next-tile accessor."""
|
||||
return await self._get_next_tile_index(multi_job_id)
|
||||
|
||||
async def _get_total_completed_count(self, multi_job_id):
|
||||
"""Get total count of all completed images (master + workers)."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -67,6 +79,10 @@ class JobStateMixin:
|
||||
return len(job_data.completed_tasks)
|
||||
return 0
|
||||
|
||||
async def total_completed_count(self, multi_job_id):
|
||||
"""Public completed-count accessor."""
|
||||
return await self._get_total_completed_count(multi_job_id)
|
||||
|
||||
async def _get_all_completed_images(self, multi_job_id):
|
||||
"""Get all completed images."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -76,6 +92,10 @@ class JobStateMixin:
|
||||
return job_data.completed_images.copy()
|
||||
return {}
|
||||
|
||||
async def all_completed_images(self, multi_job_id):
|
||||
"""Public completed-image accessor."""
|
||||
return await self._get_all_completed_images(multi_job_id)
|
||||
|
||||
async def _get_pending_count(self, multi_job_id):
|
||||
"""Get count of pending images in the queue."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -87,6 +107,10 @@ class JobStateMixin:
|
||||
return job_data.pending_tasks.qsize()
|
||||
return 0
|
||||
|
||||
async def pending_count(self, multi_job_id):
|
||||
"""Public pending-count accessor."""
|
||||
return await self._get_pending_count(multi_job_id)
|
||||
|
||||
async def _drain_worker_results_queue(self, multi_job_id):
|
||||
"""Drain pending worker results from queue and update completed images."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
@@ -130,6 +154,14 @@ class JobStateMixin:
|
||||
|
||||
return collected
|
||||
|
||||
async def drain_worker_results_queue(self, multi_job_id):
|
||||
"""Public worker-result drain API."""
|
||||
return await self._drain_worker_results_queue(multi_job_id)
|
||||
|
||||
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
|
||||
"""Check for timed out workers and requeue their assigned images."""
|
||||
return await _requeue_usdu(multi_job_id, batch_size)
|
||||
return await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
|
||||
|
||||
async def check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
|
||||
"""Public timeout-requeue API."""
|
||||
return await self._check_and_requeue_timed_out_workers(multi_job_id, batch_size)
|
||||
|
||||
+79
-60
@@ -1,45 +1,51 @@
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from typing import List, Optional
|
||||
|
||||
import server
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from ..utils.logging import debug_log
|
||||
from ..utils.runtime_state import ensure_distributed_runtime_state, get_prompt_server_instance
|
||||
from .job_models import BaseJobState, ImageJobState, TileJobState
|
||||
|
||||
# Configure maximum payload size (50MB default, configurable via environment variable)
|
||||
MAX_PAYLOAD_SIZE = int(os.environ.get('COMFYUI_MAX_PAYLOAD_SIZE', str(50 * 1024 * 1024)))
|
||||
|
||||
|
||||
def ensure_tile_jobs_initialized():
|
||||
@dataclass(frozen=True)
|
||||
class JobQueueInitConfig:
|
||||
mode: str
|
||||
batch_size: int = 0
|
||||
num_tiles_per_image: int = 0
|
||||
all_indices: list[int] | None = None
|
||||
enabled_workers: list[str] | None = None
|
||||
batched_static: bool = False
|
||||
|
||||
|
||||
def ensure_tile_jobs_initialized() -> Any:
|
||||
"""Ensure tile job storage is initialized on the server instance."""
|
||||
prompt_server = server.PromptServer.instance
|
||||
if not hasattr(prompt_server, 'distributed_pending_tile_jobs'):
|
||||
debug_log("Initializing persistent tile job queue on server instance.")
|
||||
prompt_server.distributed_pending_tile_jobs = {}
|
||||
prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
|
||||
else:
|
||||
invalid_job_ids = [
|
||||
job_id
|
||||
for job_id, job_data in prompt_server.distributed_pending_tile_jobs.items()
|
||||
if not isinstance(job_data, BaseJobState)
|
||||
]
|
||||
for job_id in invalid_job_ids:
|
||||
debug_log(f"Removing invalid job state for {job_id}")
|
||||
del prompt_server.distributed_pending_tile_jobs[job_id]
|
||||
prompt_server = get_prompt_server_instance()
|
||||
state = ensure_distributed_runtime_state(prompt_server)
|
||||
if not isinstance(state.distributed_pending_tile_jobs, dict):
|
||||
debug_log("Resetting invalid distributed_pending_tile_jobs state to empty dict.")
|
||||
state.distributed_pending_tile_jobs = {}
|
||||
ensure_distributed_runtime_state(prompt_server)
|
||||
|
||||
invalid_job_ids = [
|
||||
job_id
|
||||
for job_id, job_data in state.distributed_pending_tile_jobs.items()
|
||||
if not isinstance(job_data, BaseJobState)
|
||||
]
|
||||
for job_id in invalid_job_ids:
|
||||
debug_log(f"Removing invalid job state for {job_id}")
|
||||
del state.distributed_pending_tile_jobs[job_id]
|
||||
return prompt_server
|
||||
|
||||
|
||||
async def _init_job_queue(
|
||||
multi_job_id,
|
||||
mode,
|
||||
batch_size=None,
|
||||
num_tiles_per_image=None,
|
||||
all_indices=None,
|
||||
enabled_workers=None,
|
||||
batched_static: bool = False,
|
||||
):
|
||||
multi_job_id: str,
|
||||
config: JobQueueInitConfig,
|
||||
) -> None:
|
||||
"""Unified initialization for job queues in static and dynamic modes."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -47,33 +53,34 @@ async def _init_job_queue(
|
||||
debug_log(f"Queue already exists for {multi_job_id}")
|
||||
return
|
||||
|
||||
if mode == 'dynamic':
|
||||
if config.mode == 'dynamic':
|
||||
job_data = ImageJobState(multi_job_id=multi_job_id)
|
||||
elif mode == 'static':
|
||||
elif config.mode == 'static':
|
||||
job_data = TileJobState(multi_job_id=multi_job_id)
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {mode}")
|
||||
raise ValueError(f"Unknown mode: {config.mode}")
|
||||
|
||||
job_data.worker_status = {w: time.time() for w in enabled_workers or []}
|
||||
job_data.assigned_to_workers = {w: [] for w in enabled_workers or []}
|
||||
job_data.worker_status = {w: time.time() for w in config.enabled_workers or []}
|
||||
job_data.assigned_to_workers = {w: [] for w in config.enabled_workers or []}
|
||||
|
||||
if mode == 'dynamic':
|
||||
job_data.batch_size = int(batch_size or 0)
|
||||
if config.mode == 'dynamic':
|
||||
job_data.batch_size = int(config.batch_size or 0)
|
||||
pending_queue = job_data.pending_images
|
||||
for i in (all_indices or range(int(batch_size or 0))):
|
||||
indices = config.all_indices or list(range(int(config.batch_size or 0)))
|
||||
for i in indices:
|
||||
await pending_queue.put(i)
|
||||
debug_log(f"Initialized image queue with {batch_size} pending items")
|
||||
elif mode == 'static':
|
||||
job_data.num_tiles_per_image = int(num_tiles_per_image or 0)
|
||||
job_data.batch_size = int(batch_size or 0)
|
||||
job_data.batched_static = bool(batched_static)
|
||||
debug_log(f"Initialized image queue with {config.batch_size} pending items")
|
||||
elif config.mode == 'static':
|
||||
job_data.num_tiles_per_image = int(config.num_tiles_per_image or 0)
|
||||
job_data.batch_size = int(config.batch_size or 0)
|
||||
job_data.batched_static = bool(config.batched_static)
|
||||
# For batched static distribution, populate only tile ids [0..num_tiles_per_image-1]
|
||||
pending_queue = job_data.pending_tasks
|
||||
if batched_static and num_tiles_per_image is not None:
|
||||
for i in range(num_tiles_per_image):
|
||||
if config.batched_static and config.num_tiles_per_image > 0:
|
||||
for i in range(config.num_tiles_per_image):
|
||||
await pending_queue.put(i)
|
||||
else:
|
||||
total_tiles = int(batch_size or 0) * int(num_tiles_per_image or 0)
|
||||
total_tiles = int(config.batch_size or 0) * int(config.num_tiles_per_image or 0)
|
||||
for i in range(total_tiles):
|
||||
await pending_queue.put(i)
|
||||
|
||||
@@ -83,16 +90,18 @@ async def _init_job_queue(
|
||||
async def init_dynamic_job(
|
||||
multi_job_id: str,
|
||||
batch_size: int,
|
||||
enabled_workers: List[str],
|
||||
all_indices: Optional[List[int]] = None,
|
||||
):
|
||||
enabled_workers: list[str],
|
||||
all_indices: list[int] | None = None,
|
||||
) -> None:
|
||||
"""Initialize queue for dynamic mode (per-image), with collector fields."""
|
||||
await _init_job_queue(
|
||||
multi_job_id,
|
||||
'dynamic',
|
||||
batch_size=batch_size,
|
||||
all_indices=all_indices or list(range(batch_size)),
|
||||
enabled_workers=enabled_workers,
|
||||
JobQueueInitConfig(
|
||||
mode='dynamic',
|
||||
batch_size=batch_size,
|
||||
all_indices=all_indices or list(range(batch_size)),
|
||||
enabled_workers=enabled_workers,
|
||||
),
|
||||
)
|
||||
debug_log(f"Job {multi_job_id} initialized with {batch_size} images")
|
||||
|
||||
@@ -101,20 +110,22 @@ async def init_static_job_batched(
|
||||
multi_job_id: str,
|
||||
batch_size: int,
|
||||
num_tiles_per_image: int,
|
||||
enabled_workers: List[str],
|
||||
):
|
||||
enabled_workers: list[str],
|
||||
) -> None:
|
||||
"""Initialize queue for static mode (batched-per-tile)."""
|
||||
await _init_job_queue(
|
||||
multi_job_id,
|
||||
'static',
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
enabled_workers=enabled_workers,
|
||||
batched_static=True,
|
||||
JobQueueInitConfig(
|
||||
mode='static',
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
enabled_workers=enabled_workers,
|
||||
batched_static=True,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def _drain_results_queue(multi_job_id):
|
||||
async def drain_results_queue(multi_job_id: str) -> int:
|
||||
"""Drain pending results from queue and update completed_tasks. Returns count drained."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -152,7 +163,7 @@ async def _drain_results_queue(multi_job_id):
|
||||
return collected
|
||||
|
||||
|
||||
async def _get_completed_count(multi_job_id):
|
||||
async def get_completed_count(multi_job_id: str) -> int:
|
||||
"""Get count of completed tasks."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -162,7 +173,7 @@ async def _get_completed_count(multi_job_id):
|
||||
return 0
|
||||
|
||||
|
||||
async def _mark_task_completed(multi_job_id, task_id, result):
|
||||
async def mark_task_completed(multi_job_id: str, task_id: int, result: Any) -> None:
|
||||
"""Mark a task as completed."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -171,10 +182,18 @@ async def _mark_task_completed(multi_job_id, task_id, result):
|
||||
job_data.completed_tasks[task_id] = result
|
||||
|
||||
|
||||
async def _cleanup_job(multi_job_id):
|
||||
async def cleanup_job(multi_job_id: str) -> None:
|
||||
"""Cleanup the job data."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
|
||||
del prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
debug_log(f"Cleaned up job {multi_job_id}")
|
||||
|
||||
|
||||
async def init_job_queue(
|
||||
multi_job_id: str,
|
||||
config: JobQueueInitConfig,
|
||||
) -> None:
|
||||
"""Public entry point for unified job queue initialization."""
|
||||
await _init_job_queue(multi_job_id, config)
|
||||
|
||||
@@ -99,6 +99,10 @@ async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
|
||||
f"Probe diagnostics: online={probe_queue is not None} queue_remaining={probe_queue}"
|
||||
)
|
||||
|
||||
if incomplete_assigned == 0:
|
||||
debug_log(f"Worker {worker} heartbeat stale but all tasks complete; skipping")
|
||||
continue
|
||||
|
||||
if busy:
|
||||
workers_graced.append(worker)
|
||||
debug_log(f"Heartbeat grace: worker {worker} busy via probe; skipping requeue")
|
||||
@@ -148,3 +152,8 @@ async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
|
||||
job_data.assigned_to_workers[worker] = []
|
||||
|
||||
return requeued_count
|
||||
|
||||
|
||||
async def check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
|
||||
"""Public wrapper for worker-timeout requeue checks."""
|
||||
return await _check_and_requeue_timed_out_workers(multi_job_id, total_tasks)
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
class TileOpsCollaborator:
|
||||
"""Typed facade over tile-related operations used by USDU modes."""
|
||||
|
||||
def __init__(self, delegate: Any) -> None:
|
||||
self._delegate = delegate
|
||||
|
||||
def round_to_multiple(self, value: int) -> int:
|
||||
return self._delegate.round_to_multiple(value)
|
||||
|
||||
def calculate_tiles(
|
||||
self,
|
||||
width: int,
|
||||
height: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
force_uniform_tiles: bool,
|
||||
) -> list[tuple[int, int]]:
|
||||
return self._delegate.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
|
||||
def create_tile_mask(
|
||||
self,
|
||||
width: int,
|
||||
height: int,
|
||||
tx: int,
|
||||
ty: int,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
mask_blur: int,
|
||||
) -> Any:
|
||||
return self._delegate.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur)
|
||||
|
||||
def blend_tile(
|
||||
self,
|
||||
base_image: Any,
|
||||
tile_pil: Any,
|
||||
x1: int,
|
||||
y1: int,
|
||||
size: tuple[int, int],
|
||||
tile_mask: Any,
|
||||
padding: int,
|
||||
) -> Any:
|
||||
return self._delegate.blend_tile(base_image, tile_pil, x1, y1, size, tile_mask, padding)
|
||||
|
||||
def slice_conditioning(self, positive: Any, negative: Any, image_idx: int) -> tuple[Any, Any]:
|
||||
return self._delegate.slice_conditioning(positive, negative, image_idx)
|
||||
|
||||
def process_and_blend_tile(
|
||||
self,
|
||||
tile_idx: int,
|
||||
tile_pos: tuple[int, int],
|
||||
upscaled_image: Any,
|
||||
result_image: Any,
|
||||
model: Any,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
vae: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
denoise: float,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
image_width: int,
|
||||
image_height: int,
|
||||
force_uniform_tiles: bool,
|
||||
tiled_decode: bool,
|
||||
batch_idx: int,
|
||||
) -> Any:
|
||||
return self._delegate.process_and_blend_tile(
|
||||
tile_idx,
|
||||
tile_pos,
|
||||
upscaled_image,
|
||||
result_image,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
mask_blur,
|
||||
image_width,
|
||||
image_height,
|
||||
force_uniform_tiles,
|
||||
tiled_decode,
|
||||
batch_idx=batch_idx,
|
||||
)
|
||||
|
||||
|
||||
class JobStateCollaborator:
|
||||
"""Typed facade over job-state operations used by USDU modes."""
|
||||
|
||||
def __init__(self, delegate: Any) -> None:
|
||||
self._delegate = delegate
|
||||
|
||||
async def next_tile_index(self, multi_job_id: str) -> int | None:
|
||||
return await self._delegate.next_tile_index(multi_job_id)
|
||||
|
||||
async def all_completed_tasks(self, multi_job_id: str) -> dict[int, dict[str, Any]]:
|
||||
return await self._delegate.all_completed_tasks(multi_job_id)
|
||||
|
||||
async def check_and_requeue_timed_out_workers(self, multi_job_id: str, expected_total: int) -> int:
|
||||
return await self._delegate.check_and_requeue_timed_out_workers(multi_job_id, expected_total)
|
||||
|
||||
async def next_image_index(self, multi_job_id: str) -> int | None:
|
||||
return await self._delegate.next_image_index(multi_job_id)
|
||||
|
||||
async def drain_worker_results_queue(self, multi_job_id: str) -> int:
|
||||
return await self._delegate.drain_worker_results_queue(multi_job_id)
|
||||
|
||||
async def total_completed_count(self, multi_job_id: str) -> int:
|
||||
return await self._delegate.total_completed_count(multi_job_id)
|
||||
|
||||
async def pending_count(self, multi_job_id: str) -> int:
|
||||
return await self._delegate.pending_count(multi_job_id)
|
||||
|
||||
async def all_completed_images(self, multi_job_id: str) -> dict[int, Any]:
|
||||
return await self._delegate.all_completed_images(multi_job_id)
|
||||
|
||||
|
||||
class WorkerCommsCollaborator:
|
||||
"""Typed facade over worker/master communication operations used by USDU modes."""
|
||||
|
||||
def __init__(self, delegate: Any) -> None:
|
||||
self._delegate = delegate
|
||||
|
||||
async def request_assignment(self, multi_job_id: str, master_url: str, worker_id: str) -> Any:
|
||||
return await self._delegate.request_assignment(multi_job_id, master_url, worker_id)
|
||||
|
||||
async def send_tiles_batch(
|
||||
self,
|
||||
processed_tiles: list[dict[str, Any]],
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
padding: int,
|
||||
worker_id: str,
|
||||
is_final_flush: bool = False,
|
||||
) -> None:
|
||||
await self._delegate.send_tiles_batch(
|
||||
processed_tiles,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
padding,
|
||||
worker_id,
|
||||
is_final_flush=is_final_flush,
|
||||
)
|
||||
|
||||
async def send_heartbeat(self, multi_job_id: str, master_url: str, worker_id: str) -> None:
|
||||
await self._delegate.send_heartbeat(multi_job_id, master_url, worker_id)
|
||||
|
||||
async def check_job_status(self, multi_job_id: str, master_url: str) -> bool:
|
||||
return await self._delegate.check_job_status(multi_job_id, master_url)
|
||||
|
||||
async def async_yield(self) -> None:
|
||||
await self._delegate.async_yield()
|
||||
|
||||
async def send_full_image(
|
||||
self,
|
||||
image_pil: Any,
|
||||
image_idx: int,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str,
|
||||
is_last: bool,
|
||||
) -> None:
|
||||
await self._delegate.send_full_image(
|
||||
image_pil,
|
||||
image_idx,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id,
|
||||
is_last,
|
||||
)
|
||||
|
||||
async def send_worker_complete_signal(self, multi_job_id: str, master_url: str, worker_id: str) -> None:
|
||||
await self._delegate.send_worker_complete_signal(multi_job_id, master_url, worker_id)
|
||||
|
||||
|
||||
class ResultCollectorCollaborator:
|
||||
"""Typed facade over result-collection operations used by dynamic mode."""
|
||||
|
||||
def __init__(self, delegate: Any) -> None:
|
||||
self._delegate = delegate
|
||||
|
||||
async def mark_image_completed(self, multi_job_id: str, image_idx: int, image_pil: Any) -> None:
|
||||
await self._delegate.mark_image_completed(multi_job_id, image_idx, image_pil)
|
||||
|
||||
async def collect_dynamic_images(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
remaining_to_collect: int,
|
||||
num_workers: int,
|
||||
batch_size: int,
|
||||
master_processed_count: int,
|
||||
) -> dict[int, Any]:
|
||||
return await self._delegate.collect_dynamic_images(
|
||||
multi_job_id,
|
||||
remaining_to_collect,
|
||||
num_workers,
|
||||
batch_size,
|
||||
master_processed_count,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SingleGpuModeContext:
|
||||
tile_ops: TileOpsCollaborator
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StaticModeContext:
|
||||
tile_ops: TileOpsCollaborator
|
||||
job_state: JobStateCollaborator
|
||||
worker_comms: WorkerCommsCollaborator
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DynamicModeContext:
|
||||
tile_ops: TileOpsCollaborator
|
||||
job_state: JobStateCollaborator
|
||||
worker_comms: WorkerCommsCollaborator
|
||||
result_collector: ResultCollectorCollaborator
|
||||
+309
-150
@@ -1,12 +1,27 @@
|
||||
import asyncio, torch
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio, time, torch
|
||||
from typing import Any
|
||||
from PIL import Image
|
||||
import comfy.model_management
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.image import tensor_to_pil, pil_to_tensor
|
||||
from ...utils.async_helpers import run_async_in_server_loop
|
||||
from ...utils.config import get_worker_timeout_seconds
|
||||
from ...utils.constants import TILE_WAIT_TIMEOUT, TILE_SEND_TIMEOUT
|
||||
from ...utils.constants import (
|
||||
JOB_POLL_INTERVAL,
|
||||
JOB_POLL_MAX_ATTEMPTS,
|
||||
TILE_WAIT_TIMEOUT,
|
||||
TILE_SEND_TIMEOUT,
|
||||
)
|
||||
from ..job_store import ensure_tile_jobs_initialized, init_dynamic_job
|
||||
from ..mode_contexts import (
|
||||
DynamicModeContext,
|
||||
JobStateCollaborator,
|
||||
ResultCollectorCollaborator,
|
||||
TileOpsCollaborator,
|
||||
WorkerCommsCollaborator,
|
||||
)
|
||||
from ..processing_args import UpscaleCoreArgs
|
||||
|
||||
|
||||
class DynamicModeMixin:
|
||||
@@ -14,16 +29,205 @@ class DynamicModeMixin:
|
||||
Dynamic (per-image queue) USDU mode behaviors for master and worker roles.
|
||||
|
||||
Expected co-mixins on `self`:
|
||||
- TileOpsMixin (`calculate_tiles`, `_slice_conditioning`, `_process_and_blend_tile`).
|
||||
- TileOpsMixin (`calculate_tiles`, `slice_conditioning`, `process_and_blend_tile`).
|
||||
- JobStateMixin (image queue/task completion helpers).
|
||||
- WorkerCommsMixin (`_request_image_from_master`, `_send_full_image_to_master`, `_send_heartbeat_to_master`).
|
||||
- WorkerCommsMixin (`request_assignment`, `send_full_image`, `send_heartbeat`).
|
||||
"""
|
||||
|
||||
def process_master_dynamic(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers):
|
||||
def _build_dynamic_mode_context(self) -> DynamicModeContext:
|
||||
"""Build explicit collaborators for dynamic mode execution."""
|
||||
return DynamicModeContext(
|
||||
tile_ops=TileOpsCollaborator(self),
|
||||
job_state=JobStateCollaborator(self),
|
||||
worker_comms=WorkerCommsCollaborator(self),
|
||||
result_collector=ResultCollectorCollaborator(self),
|
||||
)
|
||||
|
||||
def _poll_job_ready_dynamic(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str | None = None,
|
||||
max_attempts: int = JOB_POLL_MAX_ATTEMPTS,
|
||||
mode_context: DynamicModeContext | None = None,
|
||||
) -> bool:
|
||||
"""Poll master for job readiness using explicit worker-comms collaborator."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
for attempt in range(max_attempts):
|
||||
ready = run_async_in_server_loop(
|
||||
context.worker_comms.check_job_status(multi_job_id, master_url),
|
||||
timeout=5.0,
|
||||
)
|
||||
if ready:
|
||||
if worker_id:
|
||||
debug_log(f"Worker[{worker_id[:8]}] job {multi_job_id} ready after {attempt} attempts")
|
||||
else:
|
||||
debug_log(f"Job {multi_job_id} ready after {attempt} attempts")
|
||||
return True
|
||||
time.sleep(JOB_POLL_INTERVAL)
|
||||
return False
|
||||
|
||||
def _process_image_tiles(
|
||||
self,
|
||||
image_idx: int,
|
||||
local_image: Image.Image,
|
||||
single_tensor: torch.Tensor,
|
||||
core_args: UpscaleCoreArgs,
|
||||
all_tiles: list,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
width: int,
|
||||
height: int,
|
||||
force_uniform_tiles: bool,
|
||||
context: DynamicModeContext,
|
||||
per_tile_callback: Any | None = None,
|
||||
) -> Image.Image:
|
||||
"""Process all tiles for a single image. Returns the processed PIL image."""
|
||||
positive_sliced, negative_sliced = context.tile_ops.slice_conditioning(
|
||||
core_args.positive, core_args.negative, image_idx,
|
||||
)
|
||||
for tile_idx, pos in enumerate(all_tiles):
|
||||
source_tensor = pil_to_tensor(local_image)
|
||||
if single_tensor.is_cuda:
|
||||
source_tensor = source_tensor.cuda()
|
||||
local_image = context.tile_ops.process_and_blend_tile(
|
||||
tile_idx, pos, source_tensor, local_image,
|
||||
core_args.model,
|
||||
positive_sliced,
|
||||
negative_sliced,
|
||||
core_args.vae,
|
||||
core_args.seed,
|
||||
core_args.steps,
|
||||
core_args.cfg,
|
||||
core_args.sampler_name,
|
||||
core_args.scheduler,
|
||||
core_args.denoise,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding, mask_blur, width, height, force_uniform_tiles,
|
||||
core_args.tiled_decode,
|
||||
batch_idx=image_idx,
|
||||
)
|
||||
if per_tile_callback:
|
||||
per_tile_callback()
|
||||
return local_image
|
||||
|
||||
def _handle_master_dynamic_idle_state(
|
||||
self,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
consecutive_retries,
|
||||
max_consecutive_retries,
|
||||
mode_context: DynamicModeContext | None = None,
|
||||
):
|
||||
"""Handle queue-empty state while master waits for worker progress."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
drained_count = run_async_in_server_loop(
|
||||
context.job_state.drain_worker_results_queue(multi_job_id),
|
||||
timeout=5.0,
|
||||
)
|
||||
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
|
||||
|
||||
requeued_count = run_async_in_server_loop(
|
||||
context.job_state.check_and_requeue_timed_out_workers(multi_job_id, batch_size),
|
||||
timeout=5.0,
|
||||
)
|
||||
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
|
||||
|
||||
if requeued_count > 0:
|
||||
log(f"Requeued {requeued_count} images from timed out workers")
|
||||
return False, True, 0
|
||||
|
||||
completed_now = run_async_in_server_loop(
|
||||
context.job_state.total_completed_count(multi_job_id),
|
||||
timeout=1.0,
|
||||
)
|
||||
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
|
||||
if completed_now >= batch_size:
|
||||
return True, False, consecutive_retries
|
||||
|
||||
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
|
||||
|
||||
pending_count = run_async_in_server_loop(
|
||||
context.job_state.pending_count(multi_job_id),
|
||||
timeout=1.0,
|
||||
)
|
||||
if pending_count > 0:
|
||||
return False, True, 0
|
||||
|
||||
consecutive_retries += 1
|
||||
if consecutive_retries >= max_consecutive_retries:
|
||||
log(f"Max retries ({max_consecutive_retries}) reached. Forcing collection of remaining results.")
|
||||
return True, False, consecutive_retries
|
||||
|
||||
if drained_count > 0:
|
||||
debug_log(f"Master idle drain picked up {drained_count} images while waiting for workers")
|
||||
debug_log("Waiting for workers")
|
||||
run_async_in_server_loop(asyncio.sleep(2), timeout=3.0)
|
||||
return False, True, consecutive_retries
|
||||
|
||||
def _finalize_master_dynamic_results(
|
||||
self,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_workers,
|
||||
processed_count,
|
||||
result_images,
|
||||
upscaled_image,
|
||||
mode_context: DynamicModeContext | None = None,
|
||||
):
|
||||
"""Collect remaining worker outputs and convert final images back to tensor."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
all_completed = run_async_in_server_loop(
|
||||
context.job_state.all_completed_images(multi_job_id),
|
||||
timeout=5.0,
|
||||
)
|
||||
remaining_to_collect = batch_size - len(all_completed)
|
||||
if remaining_to_collect > 0:
|
||||
debug_log(f"Waiting for {remaining_to_collect} more images from workers")
|
||||
collection_timeout = float(get_worker_timeout_seconds())
|
||||
collected_images = run_async_in_server_loop(
|
||||
context.result_collector.collect_dynamic_images(
|
||||
multi_job_id,
|
||||
remaining_to_collect,
|
||||
num_workers,
|
||||
batch_size,
|
||||
processed_count,
|
||||
),
|
||||
timeout=collection_timeout,
|
||||
)
|
||||
all_completed.update(collected_images)
|
||||
|
||||
for idx, processed_img in all_completed.items():
|
||||
if idx < batch_size:
|
||||
result_images[idx] = processed_img
|
||||
|
||||
result_tensor = (
|
||||
torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
|
||||
if batch_size > 1
|
||||
else pil_to_tensor(result_images[0])
|
||||
)
|
||||
if upscaled_image.is_cuda:
|
||||
result_tensor = result_tensor.cuda()
|
||||
return result_tensor
|
||||
|
||||
def process_master_dynamic(
|
||||
self,
|
||||
upscaled_image: torch.Tensor,
|
||||
core_args: UpscaleCoreArgs,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
multi_job_id: str,
|
||||
enabled_workers: list[str],
|
||||
mode_context: DynamicModeContext | None = None,
|
||||
) -> tuple[torch.Tensor]:
|
||||
"""Dynamic mode for large batches - assigns whole images to workers dynamically, including master."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
# Get batch size and dimensions
|
||||
batch_size, height, width, _ = upscaled_image.shape
|
||||
num_workers = len(enabled_workers)
|
||||
@@ -36,7 +240,7 @@ class DynamicModeMixin:
|
||||
debug_log(f"Processing {batch_size} images dynamically across master + {num_workers} workers.")
|
||||
|
||||
# Calculate tiles for processing
|
||||
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
|
||||
# Initialize job queue for communication
|
||||
try:
|
||||
@@ -61,7 +265,7 @@ class DynamicModeMixin:
|
||||
while processed_count < batch_size:
|
||||
# Try to get an image to process
|
||||
image_idx = run_async_in_server_loop(
|
||||
self._get_next_image_index(multi_job_id),
|
||||
context.job_state.next_image_index(multi_job_id),
|
||||
timeout=5.0 # Short timeout to allow frequent checks
|
||||
)
|
||||
|
||||
@@ -74,38 +278,25 @@ class DynamicModeMixin:
|
||||
# Process locally
|
||||
single_tensor = upscaled_image[image_idx:image_idx+1]
|
||||
local_image = result_images[image_idx]
|
||||
image_seed = seed
|
||||
|
||||
# Pre-slice conditioning once per image (not per tile)
|
||||
positive_sliced, negative_sliced = self._slice_conditioning(positive, negative, image_idx)
|
||||
|
||||
for tile_idx, pos in enumerate(all_tiles):
|
||||
source_tensor = pil_to_tensor(local_image)
|
||||
if single_tensor.is_cuda:
|
||||
source_tensor = source_tensor.cuda()
|
||||
local_image = self._process_and_blend_tile(
|
||||
tile_idx, pos, source_tensor, local_image,
|
||||
model, positive_sliced, negative_sliced, vae, image_seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height,
|
||||
padding, mask_blur, width, height, force_uniform_tiles,
|
||||
tiled_decode, batch_idx=image_idx
|
||||
)
|
||||
|
||||
# Yield after each tile to minimize worker downtime
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1)
|
||||
# Note: No per-tile drain here – that's what makes this "per-image"
|
||||
|
||||
|
||||
local_image = self._process_image_tiles(
|
||||
image_idx, local_image, single_tensor, core_args, all_tiles,
|
||||
tile_width, tile_height, padding, mask_blur, width, height,
|
||||
force_uniform_tiles, context,
|
||||
per_tile_callback=lambda: run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1),
|
||||
)
|
||||
|
||||
result_images[image_idx] = local_image
|
||||
|
||||
# Mark as completed
|
||||
run_async_in_server_loop(
|
||||
self._mark_image_completed(multi_job_id, image_idx, local_image),
|
||||
context.result_collector.mark_image_completed(multi_job_id, image_idx, local_image),
|
||||
timeout=5.0
|
||||
)
|
||||
|
||||
# NEW: Drain after the full image is marked complete (catches workers who finished during master's processing)
|
||||
drained_count = run_async_in_server_loop(
|
||||
self._drain_worker_results_queue(multi_job_id),
|
||||
context.job_state.drain_worker_results_queue(multi_job_id),
|
||||
timeout=5.0
|
||||
)
|
||||
|
||||
@@ -114,133 +305,106 @@ class DynamicModeMixin:
|
||||
|
||||
# NEW: Log overall progress (includes master's image + any drained workers)
|
||||
completed_now = run_async_in_server_loop(
|
||||
self._get_total_completed_count(multi_job_id),
|
||||
context.job_state.total_completed_count(multi_job_id),
|
||||
timeout=1.0
|
||||
)
|
||||
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
|
||||
|
||||
# Yield to allow workers to get new images after completing one
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1)
|
||||
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
|
||||
else:
|
||||
# Queue empty: collect any queued worker results to update progress
|
||||
drained_count = run_async_in_server_loop(
|
||||
self._drain_worker_results_queue(multi_job_id),
|
||||
timeout=5.0
|
||||
should_break, should_continue, consecutive_retries = self._handle_master_dynamic_idle_state(
|
||||
multi_job_id=multi_job_id,
|
||||
batch_size=batch_size,
|
||||
consecutive_retries=consecutive_retries,
|
||||
max_consecutive_retries=max_consecutive_retries,
|
||||
mode_context=context,
|
||||
)
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield after drain
|
||||
|
||||
# Check for timed out workers and requeue their images
|
||||
requeued_count = run_async_in_server_loop(
|
||||
self._check_and_requeue_timed_out_workers(multi_job_id, batch_size),
|
||||
timeout=5.0
|
||||
)
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield after requeue
|
||||
|
||||
if requeued_count > 0:
|
||||
log(f"Requeued {requeued_count} images from timed out workers")
|
||||
consecutive_retries = 0 # Reset since we have work to do
|
||||
continue
|
||||
|
||||
# Now check total completed (includes newly collected)
|
||||
completed_now = run_async_in_server_loop(
|
||||
self._get_total_completed_count(multi_job_id),
|
||||
timeout=1.0
|
||||
)
|
||||
|
||||
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
|
||||
|
||||
if completed_now >= batch_size:
|
||||
if should_break:
|
||||
break
|
||||
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield before pending check
|
||||
|
||||
# Check if there are pending images in the queue (could be requeued)
|
||||
pending_count = run_async_in_server_loop(
|
||||
self._get_pending_count(multi_job_id),
|
||||
timeout=1.0
|
||||
)
|
||||
|
||||
if pending_count > 0:
|
||||
consecutive_retries = 0 # Reset retries since there's work to do
|
||||
if should_continue:
|
||||
continue
|
||||
|
||||
consecutive_retries += 1
|
||||
if consecutive_retries >= max_consecutive_retries:
|
||||
log(f"Max retries ({max_consecutive_retries}) reached. Forcing collection of remaining results.")
|
||||
break # Force exit to collection phase
|
||||
|
||||
debug_log("Waiting for workers")
|
||||
# Use async sleep to allow event loop to process worker requests
|
||||
run_async_in_server_loop(asyncio.sleep(2), timeout=3.0)
|
||||
|
||||
debug_log(f"Master processed {processed_count} images locally")
|
||||
|
||||
# Get all completed images to check what needs to be collected
|
||||
all_completed = run_async_in_server_loop(
|
||||
self._get_all_completed_images(multi_job_id),
|
||||
timeout=5.0
|
||||
result_tensor = self._finalize_master_dynamic_results(
|
||||
multi_job_id=multi_job_id,
|
||||
batch_size=batch_size,
|
||||
num_workers=num_workers,
|
||||
processed_count=processed_count,
|
||||
result_images=result_images,
|
||||
upscaled_image=upscaled_image,
|
||||
mode_context=context,
|
||||
)
|
||||
|
||||
# Calculate how many we still need to collect
|
||||
remaining_to_collect = batch_size - len(all_completed)
|
||||
|
||||
if remaining_to_collect > 0:
|
||||
debug_log(f"Waiting for {remaining_to_collect} more images from workers")
|
||||
# Use the unified worker timeout for the collection phase
|
||||
collection_timeout = float(get_worker_timeout_seconds())
|
||||
collected_images = run_async_in_server_loop(
|
||||
self._async_collect_dynamic_images(multi_job_id, remaining_to_collect, num_workers, batch_size, processed_count),
|
||||
timeout=collection_timeout
|
||||
)
|
||||
|
||||
# Merge collected with already completed
|
||||
all_completed.update(collected_images)
|
||||
|
||||
# Update result images with all completed images
|
||||
for idx, processed_img in all_completed.items():
|
||||
if idx < batch_size:
|
||||
result_images[idx] = processed_img
|
||||
|
||||
# Convert back to tensor
|
||||
result_tensor = torch.cat([pil_to_tensor(img) for img in result_images], dim=0) if batch_size > 1 else pil_to_tensor(result_images[0])
|
||||
if upscaled_image.is_cuda:
|
||||
result_tensor = result_tensor.cuda()
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Job {multi_job_id} complete")
|
||||
log(f"Completed processing all {batch_size} images")
|
||||
return (result_tensor,)
|
||||
|
||||
def process_worker_dynamic(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
worker_id, enabled_worker_ids, dynamic_threshold):
|
||||
def process_worker_dynamic(
|
||||
self,
|
||||
upscaled_image: torch.Tensor,
|
||||
core_args: UpscaleCoreArgs,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str,
|
||||
enabled_worker_ids: list[str] | str,
|
||||
dynamic_threshold: int,
|
||||
mode_context: DynamicModeContext | None = None,
|
||||
) -> tuple[torch.Tensor]:
|
||||
"""Worker processing in dynamic mode - processes whole images."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
# Round tile dimensions
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
tile_height = self.round_to_multiple(tile_height)
|
||||
tile_width = context.tile_ops.round_to_multiple(tile_width)
|
||||
tile_height = context.tile_ops.round_to_multiple(tile_height)
|
||||
|
||||
# Get dimensions and tile grid
|
||||
batch_size, height, width, _ = upscaled_image.shape
|
||||
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
log(f"USDU Dist Worker[{worker_id[:8]}]: Processing image queue | Batch {batch_size}")
|
||||
|
||||
# Keep track of processed images for is_last detection
|
||||
processed_count = 0
|
||||
|
||||
# Ensure the dynamic job queue exists on the master. In delegate-only
|
||||
# mode the master's USDU node is replaced with a lightweight collector
|
||||
# that cannot initialise the queue itself, so workers do it instead.
|
||||
# The call is idempotent — if the master already created the queue this
|
||||
# is a no-op.
|
||||
from ...utils.worker_ids import coerce_enabled_worker_ids as _coerce
|
||||
_enabled = _coerce(enabled_worker_ids)
|
||||
run_async_in_server_loop(
|
||||
context.worker_comms.init_dynamic_job_on_master(
|
||||
multi_job_id, master_url, batch_size, _enabled,
|
||||
),
|
||||
timeout=10.0,
|
||||
)
|
||||
|
||||
# Poll for job readiness to avoid races during master init
|
||||
max_poll_attempts = 20 # ~20s at 1s sleep
|
||||
if not self._poll_job_ready(multi_job_id, master_url, worker_id=worker_id, max_attempts=max_poll_attempts):
|
||||
max_poll_attempts = JOB_POLL_MAX_ATTEMPTS
|
||||
if not self._poll_job_ready_dynamic(
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id=worker_id,
|
||||
max_attempts=max_poll_attempts,
|
||||
mode_context=context,
|
||||
):
|
||||
log(f"Job {multi_job_id} not ready after {max_poll_attempts} attempts, aborting")
|
||||
return (upscaled_image,)
|
||||
|
||||
# Loop to request and process images
|
||||
while True:
|
||||
# Request an image to process
|
||||
image_idx, estimated_remaining = run_async_in_server_loop(
|
||||
self._request_image_from_master(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_WAIT_TIMEOUT
|
||||
assignment = run_async_in_server_loop(
|
||||
context.worker_comms.request_assignment(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_WAIT_TIMEOUT,
|
||||
)
|
||||
image_idx = assignment.task_idx if assignment.kind == "image" else None
|
||||
estimated_remaining = assignment.estimated_remaining
|
||||
|
||||
if image_idx is None:
|
||||
debug_log(f"USDU Dist Worker - No more images to process")
|
||||
@@ -259,39 +423,34 @@ class DynamicModeMixin:
|
||||
local_image = tensor_to_pil(single_tensor, 0).copy()
|
||||
|
||||
# Process all tiles for this image
|
||||
image_seed = seed
|
||||
|
||||
# Pre-slice conditioning once per image (not per tile)
|
||||
positive_sliced, negative_sliced = self._slice_conditioning(positive, negative, image_idx)
|
||||
|
||||
for tile_idx, pos in enumerate(all_tiles):
|
||||
source_tensor = pil_to_tensor(local_image)
|
||||
if single_tensor.is_cuda:
|
||||
source_tensor = source_tensor.cuda()
|
||||
local_image = self._process_and_blend_tile(
|
||||
tile_idx, pos, source_tensor, local_image,
|
||||
model, positive_sliced, negative_sliced, vae, image_seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height,
|
||||
padding, mask_blur, width, height, force_uniform_tiles,
|
||||
tiled_decode, batch_idx=image_idx
|
||||
)
|
||||
run_async_in_server_loop(
|
||||
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0
|
||||
)
|
||||
local_image = self._process_image_tiles(
|
||||
image_idx, local_image, single_tensor, core_args, all_tiles,
|
||||
tile_width, tile_height, padding, mask_blur, width, height,
|
||||
force_uniform_tiles, context,
|
||||
per_tile_callback=lambda: run_async_in_server_loop(
|
||||
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0,
|
||||
),
|
||||
)
|
||||
|
||||
# Send processed image back to master
|
||||
try:
|
||||
# Use the estimated remaining to determine if this is the last image
|
||||
is_last = is_last_for_worker
|
||||
run_async_in_server_loop(
|
||||
self._send_full_image_to_master(local_image, image_idx, multi_job_id,
|
||||
master_url, worker_id, is_last),
|
||||
context.worker_comms.send_full_image(
|
||||
local_image,
|
||||
image_idx,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id,
|
||||
is_last,
|
||||
),
|
||||
timeout=TILE_SEND_TIMEOUT
|
||||
)
|
||||
# Send heartbeat after processing
|
||||
run_async_in_server_loop(
|
||||
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
|
||||
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0
|
||||
)
|
||||
if is_last:
|
||||
@@ -304,7 +463,7 @@ class DynamicModeMixin:
|
||||
debug_log(f"Worker[{worker_id[:8]}] processed {processed_count} images, sending completion signal")
|
||||
try:
|
||||
run_async_in_server_loop(
|
||||
self._send_worker_complete_signal(multi_job_id, master_url, worker_id),
|
||||
context.worker_comms.send_worker_complete_signal(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_SEND_TIMEOUT
|
||||
)
|
||||
except Exception as e:
|
||||
|
||||
+56
-24
@@ -1,29 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math, torch
|
||||
from PIL import Image
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.image import tensor_to_pil, pil_to_tensor
|
||||
from typing import Any
|
||||
from ...utils.logging import log
|
||||
from ...utils.image import blend_processed_batch_item, pil_to_tensor, tensor_to_pil
|
||||
from ..mode_contexts import SingleGpuModeContext, TileOpsCollaborator
|
||||
from ..processing_args import UpscaleCoreArgs
|
||||
from ..tile_processing import TileBatchArgs, extract_and_process_tile_batch
|
||||
|
||||
|
||||
class SingleGpuModeMixin:
|
||||
def process_single_gpu(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode):
|
||||
def _build_single_gpu_mode_context(self) -> SingleGpuModeContext:
|
||||
"""Build explicit collaborators for single-GPU mode execution."""
|
||||
return SingleGpuModeContext(tile_ops=TileOpsCollaborator(self))
|
||||
|
||||
def process_single_gpu(
|
||||
self,
|
||||
upscaled_image: torch.Tensor,
|
||||
core_args: UpscaleCoreArgs,
|
||||
tile_width: int,
|
||||
tile_height: int,
|
||||
padding: int,
|
||||
mask_blur: int,
|
||||
force_uniform_tiles: bool,
|
||||
mode_context: SingleGpuModeContext | None = None,
|
||||
) -> tuple[torch.Tensor]:
|
||||
"""Process all tiles on a single GPU (no distribution), batching per tile like USDU."""
|
||||
context = mode_context or self._build_single_gpu_mode_context()
|
||||
ops = context.tile_ops
|
||||
# Round tile dimensions
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
tile_height = self.round_to_multiple(tile_height)
|
||||
tile_width = ops.round_to_multiple(tile_width)
|
||||
tile_height = ops.round_to_multiple(tile_height)
|
||||
|
||||
# Get image dimensions and batch size
|
||||
batch_size, height, width, _ = upscaled_image.shape
|
||||
|
||||
# Calculate all tiles
|
||||
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
|
||||
rows = math.ceil(height / tile_height)
|
||||
cols = math.ceil(width / tile_width)
|
||||
log(
|
||||
f"USDU Dist: Single GPU | Canvas {width}x{height} | Tile {tile_width}x{tile_height} | Grid {rows}x{cols} ({len(all_tiles)} tiles/image) | Batch {batch_size}"
|
||||
)
|
||||
tile_batch_args = TileBatchArgs(
|
||||
core=core_args,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
width=width,
|
||||
height=height,
|
||||
)
|
||||
|
||||
# Prepare result images list
|
||||
result_images = []
|
||||
@@ -34,7 +62,7 @@ class SingleGpuModeMixin:
|
||||
# Precompute tile masks once
|
||||
tile_masks = []
|
||||
for tx, ty in all_tiles:
|
||||
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
tile_masks.append(ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
|
||||
# Process tiles batched across images
|
||||
for tile_idx, (tx, ty) in enumerate(all_tiles):
|
||||
@@ -43,25 +71,29 @@ class SingleGpuModeMixin:
|
||||
if upscaled_image.is_cuda:
|
||||
source_batch = source_batch.cuda()
|
||||
|
||||
# Extract batched tile
|
||||
tile_batch, x1, y1, ew, eh = self.extract_batch_tile_with_padding(
|
||||
source_batch, tx, ty, tile_width, tile_height, padding, force_uniform_tiles
|
||||
processed_batch, x1, y1, ew, eh = extract_and_process_tile_batch(
|
||||
node=ops,
|
||||
upscaled_image=source_batch,
|
||||
tx=tx,
|
||||
ty=ty,
|
||||
args=tile_batch_args,
|
||||
)
|
||||
|
||||
# Process batch
|
||||
region = (x1, y1, x1 + ew, y1 + eh)
|
||||
processed_batch = self.process_tiles_batch(tile_batch, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tiled_decode, region, (width, height))
|
||||
|
||||
# Blend results back into each image using cached mask
|
||||
tile_mask = tile_masks[tile_idx]
|
||||
for b in range(batch_size):
|
||||
tile_pil = tensor_to_pil(processed_batch, b)
|
||||
# Resize back to extracted size
|
||||
if tile_pil.size != (ew, eh):
|
||||
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
|
||||
result_images[b] = self.blend_tile(result_images[b], tile_pil, x1, y1, (ew, eh), tile_mask, padding)
|
||||
blend_processed_batch_item(
|
||||
result_images,
|
||||
processed_batch,
|
||||
b,
|
||||
ops.blend_tile,
|
||||
x1,
|
||||
y1,
|
||||
ew,
|
||||
eh,
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
|
||||
# Convert back to tensor
|
||||
result_tensors = [pil_to_tensor(img) for img in result_images]
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user