Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e2125c7529 | ||
|
|
a91f9fb081 | ||
|
|
6874d2735f | ||
|
|
0e9ee7fbac | ||
|
|
5386de10e3 | ||
|
|
41f9e44945 | ||
|
|
7792d613bb | ||
|
|
6e698512b7 | ||
|
|
79cc9f5ad5 | ||
|
|
4eba87cba3 | ||
|
|
27e08de94c | ||
|
|
c8453a139b | ||
|
|
aae831e1e5 | ||
|
|
e7ab67733b | ||
|
|
dd55ff740e | ||
|
|
a6d0b82d35 |
@@ -1,4 +1,3 @@
|
||||
# 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@v1
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
with:
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -7,8 +7,3 @@ __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,22 +27,12 @@
|
||||
- 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
|
||||
|
||||
---
|
||||
|
||||
## 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.
|
||||
|
||||
---
|
||||
#### 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
|
||||
|
||||
---
|
||||
|
||||
## Worker Types
|
||||
|
||||
@@ -61,7 +51,6 @@ 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
|
||||
@@ -92,19 +81,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.
|
||||
@@ -122,12 +111,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
|
||||
@@ -153,40 +142,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 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)
|
||||
* **Endpoint:** `POST /distributed/queue`
|
||||
* **Functionality:** Accepts a standard ComfyUI workflow JSON, automatically distributes it to available workers, and returns the execution ID.
|
||||
* **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 Batch Divider** | Splits audio batches for multi-GPU output |
|
||||
> **⚠️ 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 |
|
||||
|
||||
@@ -206,7 +195,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>
|
||||
Currently, it is not compatible with the ComfyUI desktop app.
|
||||
Yes, it does now.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
@@ -249,3 +238,4 @@ Buy me a coffee at: https://buymeacoffee.com/robertvoy
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+24
-6
@@ -1,11 +1,29 @@
|
||||
"""ComfyUI-Distributed package entrypoint."""
|
||||
from __future__ import annotations
|
||||
# Import everything needed from the main module
|
||||
from .distributed import (
|
||||
NODE_CLASS_MAPPINGS as DISTRIBUTED_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as DISTRIBUTED_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
|
||||
from .bootstrap.entrypoint import build_node_mappings, initialize_runtime
|
||||
# Import utilities
|
||||
from .utils.config import ensure_config_exists, CONFIG_FILE
|
||||
from .utils.logging import debug_log
|
||||
|
||||
# Import distributed upscale nodes
|
||||
from .nodes.distributed_upscale import (
|
||||
NODE_CLASS_MAPPINGS as UPSCALE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as UPSCALE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
|
||||
initialize_runtime()
|
||||
ensure_config_exists()
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
|
||||
# Merge node mappings
|
||||
NODE_CLASS_MAPPINGS = {**DISTRIBUTED_CLASS_MAPPINGS, **UPSCALE_CLASS_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**DISTRIBUTED_DISPLAY_NAME_MAPPINGS, **UPSCALE_DISPLAY_NAME_MAPPINGS}
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
debug_log("Loaded Distributed nodes.")
|
||||
debug_log(f"Config file: {CONFIG_FILE}")
|
||||
debug_log(f"Available nodes: {list(NODE_CLASS_MAPPINGS.keys())}")
|
||||
|
||||
+5
-21
@@ -1,21 +1,5 @@
|
||||
"""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"]
|
||||
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
|
||||
|
||||
+79
-97
@@ -1,15 +1,32 @@
|
||||
from __future__ import annotations
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from aiohttp import web
|
||||
import server
|
||||
|
||||
from ..utils.config import config_transaction, load_config
|
||||
from ..utils.logging import debug_log
|
||||
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.network import handle_api_error, normalize_host
|
||||
from .endpoint_policy import run_authorized_endpoint
|
||||
|
||||
|
||||
def _positive_int(value: int) -> bool:
|
||||
def _positive_int(value):
|
||||
return value > 0
|
||||
|
||||
|
||||
@@ -69,103 +86,68 @@ 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: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
return web.json_response(load_config())
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
async def get_config_endpoint(request):
|
||||
config = load_config()
|
||||
return web.json_response(config)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config")
|
||||
async def update_config_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def update_config_endpoint(request):
|
||||
"""Bulk config update with schema validation."""
|
||||
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)
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, f"Invalid JSON payload: {e}", 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 == "workers":
|
||||
errors.extend(_validate_workers_payload(value))
|
||||
if errors:
|
||||
continue
|
||||
if key in _SETTINGS_FIELDS:
|
||||
validated_settings[key] = value
|
||||
else:
|
||||
validated_root[key] = value
|
||||
|
||||
if key in _SETTINGS_FIELDS:
|
||||
validated_settings[key] = value
|
||||
else:
|
||||
validated_root[key] = value
|
||||
|
||||
if errors:
|
||||
return await handle_api_error(request, errors, 400)
|
||||
if errors:
|
||||
return web.json_response({
|
||||
"status": "error",
|
||||
"error": errors,
|
||||
"message": "; ".join(errors),
|
||||
}, status=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})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e)
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/queue_status/{job_id}")
|
||||
async def queue_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def queue_status_endpoint(request):
|
||||
"""Check if a job queue is initialized."""
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
job_id = request.match_info['job_id']
|
||||
|
||||
# Import to ensure initialization
|
||||
@@ -177,12 +159,12 @@ async def queue_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
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})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_worker")
|
||||
async def update_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
async def update_worker_endpoint(request):
|
||||
try:
|
||||
data = await request.json()
|
||||
worker_id = data.get("worker_id")
|
||||
|
||||
@@ -221,12 +203,12 @@ async def update_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
)
|
||||
|
||||
return web.json_response({"status": "success"})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/delete_worker")
|
||||
async def delete_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def _operation() -> web.StreamResponse:
|
||||
async def delete_worker_endpoint(request):
|
||||
try:
|
||||
data = await request.json()
|
||||
worker_id = data.get("worker_id")
|
||||
|
||||
@@ -253,13 +235,13 @@ async def delete_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
"status": "success",
|
||||
"message": f"Worker {removed_worker.get('name', worker_id)} deleted"
|
||||
})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_setting")
|
||||
async def update_setting_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def update_setting_endpoint(request):
|
||||
"""Updates a specific key in the settings object."""
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
data = await request.json()
|
||||
key = data.get("key")
|
||||
value = data.get("value")
|
||||
@@ -276,13 +258,13 @@ async def update_setting_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
config['settings'][key] = value
|
||||
|
||||
return web.json_response({"status": "success", "message": f"Setting '{key}' updated."})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/config/update_master")
|
||||
async def update_master_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def update_master_endpoint(request):
|
||||
"""Updates master configuration."""
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
data = await request.json()
|
||||
|
||||
async with config_transaction() as config:
|
||||
@@ -291,5 +273,5 @@ async def update_master_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
_apply_field_patch(config['master'], data, _MASTER_FIELDS)
|
||||
|
||||
return web.json_response({"status": "success", "message": "Master configuration updated."})
|
||||
|
||||
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 400)
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
"""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)
|
||||
+43
-90
@@ -1,5 +1,3 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import io
|
||||
@@ -7,7 +5,6 @@ import os
|
||||
import base64
|
||||
import binascii
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from aiohttp import web
|
||||
import server
|
||||
@@ -18,21 +15,22 @@ 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
|
||||
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
|
||||
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 .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 _runtime_state():
|
||||
return ensure_distributed_runtime_state()
|
||||
|
||||
|
||||
def _decode_image_sync(image_path: str) -> dict[str, Any]:
|
||||
def _decode_image_sync(image_path):
|
||||
"""Decode image/video file and compute hash in a threadpool worker."""
|
||||
import base64
|
||||
import hashlib
|
||||
@@ -78,7 +76,7 @@ def _decode_image_sync(image_path: str) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _check_file_sync(filename: str, expected_hash: str) -> dict[str, Any]:
|
||||
def _check_file_sync(filename, expected_hash):
|
||||
"""Check file presence and hash in a threadpool worker."""
|
||||
import hashlib
|
||||
import folder_paths
|
||||
@@ -103,7 +101,7 @@ def _check_file_sync(filename: str, expected_hash: str) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _decode_canonical_png_tensor(image_payload: str) -> torch.Tensor:
|
||||
def _decode_canonical_png_tensor(image_payload):
|
||||
"""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.")
|
||||
@@ -134,7 +132,7 @@ def _decode_canonical_png_tensor(image_payload: str) -> torch.Tensor:
|
||||
raise ValueError(f"Failed to decode PNG image payload: {exc}") from exc
|
||||
|
||||
|
||||
def _decode_audio_payload(audio_payload: dict[str, Any]) -> dict[str, Any]:
|
||||
def _decode_audio_payload(audio_payload):
|
||||
"""Decode canonical audio payload into an AUDIO dict."""
|
||||
from ..utils.audio_payload import decode_audio_payload
|
||||
|
||||
@@ -142,20 +140,17 @@ def _decode_audio_payload(audio_payload: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/prepare_job")
|
||||
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
|
||||
async def prepare_job_endpoint(request):
|
||||
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)
|
||||
|
||||
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()
|
||||
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()
|
||||
|
||||
debug_log(f"Prepared queue for job {multi_job_id}")
|
||||
return web.json_response({"status": "success"})
|
||||
@@ -163,13 +158,9 @@ async def prepare_job_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
return await handle_api_error(request, e)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/clear_memory")
|
||||
async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def clear_memory_endpoint(request):
|
||||
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)
|
||||
@@ -186,16 +177,12 @@ async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
try:
|
||||
mm.unload_all_models()
|
||||
except AttributeError as e:
|
||||
warning = f"Model unload warning: {e}"
|
||||
warnings.append(warning)
|
||||
debug_log(warning)
|
||||
debug_log(f"Warning during model unload: {e}")
|
||||
|
||||
try:
|
||||
mm.soft_empty_cache()
|
||||
except Exception as e:
|
||||
warning = f"Cache clear warning: {e}"
|
||||
warnings.append(warning)
|
||||
debug_log(warning)
|
||||
debug_log(f"Warning during cache clear: {e}")
|
||||
|
||||
for _ in range(3):
|
||||
gc.collect()
|
||||
@@ -204,16 +191,6 @@ async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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:
|
||||
@@ -222,16 +199,13 @@ async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
debug_log(f"VRAM clear failed: {e}")
|
||||
return await handle_api_error(request, f"GPU memory clear failed: {e}", 500)
|
||||
debug_log(f"Partial VRAM clear completed with warning: {e}")
|
||||
return web.json_response({"status": "success", "message": "GPU memory cleared (with warnings)"})
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/queue")
|
||||
async def distributed_queue_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def distributed_queue_endpoint(request):
|
||||
"""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:
|
||||
@@ -243,7 +217,7 @@ async def distributed_queue_endpoint(request: web.Request) -> web.StreamResponse
|
||||
return await handle_api_error(request, exc, 400)
|
||||
|
||||
try:
|
||||
prompt_id, worker_count = await orchestrate_distributed_execution(
|
||||
prompt_id, prompt_number, worker_count, node_errors = await orchestrate_distributed_execution(
|
||||
payload.prompt,
|
||||
payload.workflow_meta,
|
||||
payload.client_id,
|
||||
@@ -253,17 +227,17 @@ async def distributed_queue_endpoint(request: web.Request) -> web.StreamResponse
|
||||
)
|
||||
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: web.Request) -> web.StreamResponse:
|
||||
async def load_image_endpoint(request):
|
||||
"""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")
|
||||
@@ -279,11 +253,8 @@ async def load_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/check_file")
|
||||
async def check_file_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def check_file_endpoint(request):
|
||||
"""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")
|
||||
@@ -300,10 +271,7 @@ async def check_file_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/job_complete")
|
||||
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
|
||||
async def job_complete_endpoint(request):
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception as exc:
|
||||
@@ -318,7 +286,7 @@ async def job_complete_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
batch_idx = data.get("batch_idx")
|
||||
image_payload = data.get("image")
|
||||
audio_payload = data.get("audio")
|
||||
is_last_raw = data.get("is_last")
|
||||
is_last = data.get("is_last")
|
||||
|
||||
errors = []
|
||||
if not isinstance(job_id, str) or not job_id.strip():
|
||||
@@ -331,12 +299,8 @@ async def job_complete_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
errors.append("image: expected non-empty base64 PNG string")
|
||||
if audio_payload is not None and not isinstance(audio_payload, dict):
|
||||
errors.append("audio: expected object when provided")
|
||||
|
||||
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 not isinstance(is_last, bool):
|
||||
errors.append("is_last: expected boolean")
|
||||
if errors:
|
||||
return await handle_api_error(request, errors, 400)
|
||||
|
||||
@@ -345,32 +309,21 @@ async def job_complete_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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 runtime_state.distributed_jobs_lock:
|
||||
pending = runtime_state.distributed_pending_jobs.get(multi_job_id)
|
||||
async with prompt_server.distributed_jobs_lock:
|
||||
pending = prompt_server.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(
|
||||
queue_item
|
||||
{
|
||||
"tensor": tensor,
|
||||
"worker_id": worker_id,
|
||||
"image_index": int(batch_idx),
|
||||
"is_last": is_last,
|
||||
"audio": decoded_audio,
|
||||
}
|
||||
)
|
||||
queue_size = pending.qsize()
|
||||
break
|
||||
|
||||
@@ -1,7 +1,4 @@
|
||||
"""Orchestration helpers used by distributed queue execution."""
|
||||
|
||||
__all__ = [
|
||||
"dispatch",
|
||||
"media_sync",
|
||||
"prompt_transform",
|
||||
]
|
||||
# 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
|
||||
|
||||
@@ -1,33 +1,45 @@
|
||||
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
|
||||
from ...utils.trace_logger import trace_debug, trace_info
|
||||
from ..schemas import coerce_positive_int
|
||||
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
|
||||
|
||||
_least_busy_rr_index = 0
|
||||
|
||||
|
||||
async def worker_is_active(worker: dict[str, Any]) -> bool:
|
||||
async def worker_is_active(worker):
|
||||
"""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: dict[str, Any]) -> bool:
|
||||
async def worker_ws_is_active(worker):
|
||||
"""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, headers=headers)
|
||||
ws = await session.ws_connect(url, heartbeat=20, timeout=3)
|
||||
await ws.close()
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
@@ -47,13 +59,7 @@ async def _probe_worker_active(worker, use_websocket, semaphore):
|
||||
return worker, is_active
|
||||
|
||||
|
||||
async def _dispatch_via_websocket(
|
||||
worker_url,
|
||||
payload,
|
||||
client_id,
|
||||
headers: dict[str, str],
|
||||
timeout=60.0,
|
||||
):
|
||||
async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
|
||||
"""Open a fresh worker websocket, dispatch one prompt, wait for ack, then close."""
|
||||
request_id = uuid.uuid4().hex
|
||||
ws_payload = {
|
||||
@@ -67,12 +73,7 @@ async def _dispatch_via_websocket(
|
||||
ws_url = f"{ws_url}/distributed/worker_ws"
|
||||
session = await get_client_session()
|
||||
|
||||
async with session.ws_connect(
|
||||
ws_url,
|
||||
heartbeat=20,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
) as ws:
|
||||
async with session.ws_connect(ws_url, heartbeat=20, timeout=timeout) as ws:
|
||||
await ws.send_json(ws_payload)
|
||||
async for msg in ws:
|
||||
if msg.type == aiohttp.WSMsgType.TEXT:
|
||||
@@ -95,13 +96,13 @@ async def _dispatch_via_websocket(
|
||||
|
||||
|
||||
async def dispatch_worker_prompt(
|
||||
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:
|
||||
worker,
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id=None,
|
||||
use_websocket=False,
|
||||
trace_execution_id=None,
|
||||
):
|
||||
"""Send the prepared prompt to a worker ComfyUI instance."""
|
||||
worker_url = build_worker_url(worker)
|
||||
url = build_worker_url(worker, "/prompt")
|
||||
@@ -113,7 +114,6 @@ 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,7 +122,6 @@ async def dispatch_worker_prompt(
|
||||
"workflow": workflow_meta,
|
||||
},
|
||||
client_id,
|
||||
ws_headers,
|
||||
)
|
||||
return
|
||||
except Exception as exc:
|
||||
@@ -143,14 +142,14 @@ async def dispatch_worker_prompt(
|
||||
|
||||
|
||||
async def select_active_workers(
|
||||
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]:
|
||||
workers,
|
||||
use_websocket,
|
||||
delegate_master,
|
||||
trace_execution_id=None,
|
||||
probe_concurrency=8,
|
||||
):
|
||||
"""Probe workers and return (active_workers, updated_delegate_master)."""
|
||||
probe_limit = coerce_positive_int(probe_concurrency, 8)
|
||||
probe_limit = parse_positive_int(probe_concurrency, 8)
|
||||
probe_semaphore = asyncio.Semaphore(probe_limit)
|
||||
|
||||
if trace_execution_id and workers:
|
||||
@@ -224,16 +223,16 @@ def _select_idle_round_robin(statuses):
|
||||
|
||||
|
||||
async def select_least_busy_worker(
|
||||
workers: list[dict[str, Any]],
|
||||
trace_execution_id: str | None = None,
|
||||
probe_concurrency: int = 8,
|
||||
probe_timeout: float = 3.0,
|
||||
) -> dict[str, Any] | None:
|
||||
workers,
|
||||
trace_execution_id=None,
|
||||
probe_concurrency=8,
|
||||
probe_timeout=3.0,
|
||||
):
|
||||
"""Select one worker by queue depth, round-robin among idle workers."""
|
||||
if not workers:
|
||||
return None
|
||||
|
||||
probe_limit = coerce_positive_int(probe_concurrency, 8)
|
||||
probe_limit = parse_positive_int(probe_concurrency, 8)
|
||||
probe_semaphore = asyncio.Semaphore(probe_limit)
|
||||
statuses = await asyncio.gather(
|
||||
*[
|
||||
@@ -267,50 +266,3 @@ 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,12 +3,9 @@ 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
|
||||
@@ -36,7 +33,7 @@ def _normalize_media_reference(value):
|
||||
return None
|
||||
|
||||
|
||||
def convert_paths_for_platform(obj: Any, target_separator: str) -> Any:
|
||||
def convert_paths_for_platform(obj, target_separator):
|
||||
"""Recursively normalize likely file paths for the worker platform separator."""
|
||||
if target_separator not in ("/", "\\"):
|
||||
return obj
|
||||
@@ -127,16 +124,12 @@ def _load_media_file_sync(filename):
|
||||
return file_bytes, file_hash, mime_type
|
||||
|
||||
|
||||
async def fetch_worker_path_separator(
|
||||
worker: dict[str, Any],
|
||||
trace_execution_id: str | None = None,
|
||||
) -> str | None:
|
||||
async def fetch_worker_path_separator(worker, trace_execution_id=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, headers=headers, timeout=aiohttp.ClientTimeout(total=5)) as resp:
|
||||
async with session.get(url, timeout=aiohttp.ClientTimeout(total=5)) as resp:
|
||||
if resp.status != 200:
|
||||
return None
|
||||
payload = await resp.json()
|
||||
@@ -153,7 +146,6 @@ async def fetch_worker_path_separator(
|
||||
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")
|
||||
@@ -161,7 +153,6 @@ 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:
|
||||
@@ -202,11 +193,7 @@ async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_
|
||||
return True, worker_path
|
||||
|
||||
|
||||
async def sync_worker_media(
|
||||
worker: dict[str, Any],
|
||||
prompt_obj: dict[str, Any],
|
||||
trace_execution_id: str | None = None,
|
||||
) -> None:
|
||||
async def sync_worker_media(worker, prompt_obj, trace_execution_id=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,6 +1,5 @@
|
||||
import json
|
||||
from collections import deque
|
||||
from typing import Any
|
||||
|
||||
from ...utils.logging import debug_log
|
||||
|
||||
@@ -28,7 +27,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: str, target_class: str) -> bool:
|
||||
def has_upstream(self, start_node_id, target_class):
|
||||
cache_key = (str(start_node_id), target_class)
|
||||
if cache_key in self._upstream_cache:
|
||||
return self._upstream_cache[cache_key]
|
||||
@@ -89,38 +88,6 @@ 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
|
||||
@@ -128,7 +95,6 @@ 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)
|
||||
|
||||
@@ -159,193 +125,15 @@ def _find_upstream_nodes(prompt_obj, start_ids):
|
||||
return connected
|
||||
|
||||
|
||||
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 _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 _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
|
||||
return False
|
||||
|
||||
|
||||
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 _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
|
||||
|
||||
next_id = _create_numeric_id_generator(prompt_obj)
|
||||
|
||||
# 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
|
||||
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
|
||||
|
||||
# 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_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]:
|
||||
def prune_prompt_for_worker(prompt_obj):
|
||||
"""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")
|
||||
branch_ids = find_nodes_by_class(prompt_obj, "DistributedBranch")
|
||||
distributed_ids = collector_ids + branch_collector_ids + upscale_ids + branch_ids
|
||||
distributed_ids = collector_ids + upscale_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)
|
||||
@@ -358,7 +146,7 @@ def prune_prompt_for_worker(prompt_obj: dict[str, Any]) -> dict[str, Any]:
|
||||
if dist_id not in pruned_prompt:
|
||||
continue
|
||||
downstream = _find_downstream_nodes(prompt_obj, [dist_id])
|
||||
has_removed_downstream = any(node_id != dist_id and node_id not in connected for node_id in downstream)
|
||||
has_removed_downstream = any(node_id != dist_id for node_id in downstream)
|
||||
if has_removed_downstream:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
@@ -374,10 +162,7 @@ def prune_prompt_for_worker(prompt_obj: dict[str, Any]) -> dict[str, Any]:
|
||||
return pruned_prompt
|
||||
|
||||
|
||||
def prepare_delegate_master_prompt(
|
||||
prompt_obj: dict[str, Any],
|
||||
collector_ids: list[str],
|
||||
) -> dict[str, Any]:
|
||||
def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
"""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)
|
||||
@@ -409,17 +194,6 @@ def prepare_delegate_master_prompt(
|
||||
collector_entry = pruned_prompt.get(collector_id)
|
||||
if not collector_entry:
|
||||
continue
|
||||
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] = {
|
||||
"class_type": "DistributedEmptyImage",
|
||||
@@ -440,145 +214,32 @@ def prepare_delegate_master_prompt(
|
||||
return pruned_prompt
|
||||
|
||||
|
||||
def generate_job_id_map(prompt_index: PromptIndex, prefix: str) -> dict[str, str]:
|
||||
def generate_job_id_map(prompt_index, prefix):
|
||||
"""Create stable per-node job IDs for distributed nodes."""
|
||||
job_map = {}
|
||||
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")
|
||||
distributed_nodes = prompt_index.nodes_for_class("DistributedCollector") + prompt_index.nodes_for_class(
|
||||
"UltimateSDUpscaleDistributed"
|
||||
)
|
||||
for node_id in distributed_nodes:
|
||||
job_map[node_id] = f"{prefix}_{node_id}"
|
||||
return job_map
|
||||
|
||||
|
||||
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"):
|
||||
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"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
|
||||
inputs = node.setdefault("inputs", {})
|
||||
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:
|
||||
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
|
||||
if is_master:
|
||||
inputs["worker_id"] = ""
|
||||
else:
|
||||
inputs["worker_id"] = f"worker_{worker_index_map.get(participant_id, 0)}"
|
||||
|
||||
|
||||
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(
|
||||
def _override_collector_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
@@ -588,77 +249,87 @@ def _override_branch_collector_nodes(
|
||||
enabled_json,
|
||||
delegate_master,
|
||||
):
|
||||
"""Configure DistributedBranchCollector nodes for branch convergence."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedBranchCollector"):
|
||||
"""Configure DistributedCollector nodes for master or worker role."""
|
||||
for node_id in prompt_index.nodes_for_class("DistributedCollector"):
|
||||
node = prompt_copy.get(node_id)
|
||||
if not isinstance(node, dict):
|
||||
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
|
||||
if prompt_index.has_upstream(node_id, "UltimateSDUpscaleDistributed"):
|
||||
node.setdefault("inputs", {})["pass_through"] = True
|
||||
continue
|
||||
|
||||
inputs = node.setdefault("inputs", {})
|
||||
# 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["multi_job_id"] = job_id_map.get(node_id, node_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"] = str(participant_id)
|
||||
inputs["worker_id"] = 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: 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]:
|
||||
prompt_copy,
|
||||
participant_id,
|
||||
enabled_worker_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
delegate_master,
|
||||
prompt_index,
|
||||
):
|
||||
"""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_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(
|
||||
_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(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
@@ -668,9 +339,14 @@ def apply_participant_overrides(
|
||||
enabled_json,
|
||||
delegate_master,
|
||||
)
|
||||
|
||||
if not is_master:
|
||||
prompt_copy = _prune_worker_downstream_of_branch_collectors(prompt_copy)
|
||||
prompt_copy = _ensure_worker_output_node(prompt_copy)
|
||||
_override_upscale_nodes(
|
||||
prompt_copy,
|
||||
prompt_index,
|
||||
is_master,
|
||||
participant_id,
|
||||
job_id_map,
|
||||
master_url,
|
||||
enabled_json,
|
||||
)
|
||||
|
||||
return prompt_copy
|
||||
|
||||
+165
-245
@@ -1,7 +1,8 @@
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import server
|
||||
|
||||
from ..utils.async_helpers import queue_prompt_payload
|
||||
from ..utils.config import load_config
|
||||
@@ -12,13 +13,11 @@ from ..utils.constants import (
|
||||
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
|
||||
)
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.network import build_master_url
|
||||
from ..utils.runtime_state import ensure_distributed_runtime_state, get_prompt_server_instance
|
||||
from ..utils.network import build_master_url, build_master_callback_url
|
||||
from ..utils.trace_logger import trace_debug
|
||||
from .schemas import coerce_positive_float, coerce_positive_int
|
||||
from .schemas import parse_positive_float, parse_positive_int
|
||||
from .orchestration.dispatch import (
|
||||
dispatch_worker_prompt,
|
||||
rank_workers_by_load,
|
||||
select_active_workers,
|
||||
select_least_busy_worker,
|
||||
)
|
||||
@@ -32,21 +31,33 @@ 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."""
|
||||
ensure_distributed_runtime_state(server_instance)
|
||||
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()
|
||||
|
||||
|
||||
async def _ensure_distributed_queue(job_id):
|
||||
"""Ensure a queue exists for the given distributed job ID."""
|
||||
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()
|
||||
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()
|
||||
|
||||
|
||||
def _resolve_enabled_workers(config, requested_ids=None):
|
||||
@@ -85,19 +96,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 = coerce_positive_int(
|
||||
worker_probe_concurrency = parse_positive_int(
|
||||
settings.get("worker_probe_concurrency"),
|
||||
ORCHESTRATION_WORKER_PROBE_CONCURRENCY,
|
||||
)
|
||||
worker_prep_concurrency = coerce_positive_int(
|
||||
worker_prep_concurrency = parse_positive_int(
|
||||
settings.get("worker_prep_concurrency"),
|
||||
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
|
||||
)
|
||||
media_sync_concurrency = coerce_positive_int(
|
||||
media_sync_concurrency = parse_positive_int(
|
||||
settings.get("media_sync_concurrency"),
|
||||
ORCHESTRATION_MEDIA_SYNC_CONCURRENCY,
|
||||
)
|
||||
media_sync_timeout_seconds = coerce_positive_float(
|
||||
media_sync_timeout_seconds = parse_positive_float(
|
||||
settings.get("media_sync_timeout_seconds"),
|
||||
ORCHESTRATION_MEDIA_SYNC_TIMEOUT,
|
||||
)
|
||||
@@ -133,6 +144,7 @@ async def _prepare_worker_payload(
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
config,
|
||||
delegate_master,
|
||||
trace_execution_id,
|
||||
worker_prep_semaphore,
|
||||
@@ -142,6 +154,11 @@ 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"
|
||||
@@ -156,7 +173,7 @@ async def _prepare_worker_payload(
|
||||
worker["id"],
|
||||
enabled_ids,
|
||||
job_id_map,
|
||||
master_url,
|
||||
worker_master_url,
|
||||
delegate_master,
|
||||
prompt_index,
|
||||
)
|
||||
@@ -180,16 +197,60 @@ async def _prepare_worker_payload(
|
||||
return worker, worker_prompt
|
||||
|
||||
|
||||
async def _select_execution_workers(
|
||||
workers,
|
||||
use_websocket,
|
||||
delegate_master,
|
||||
load_balance_requested,
|
||||
has_branch_nodes,
|
||||
master_url,
|
||||
execution_trace_id,
|
||||
worker_probe_concurrency,
|
||||
async def orchestrate_distributed_execution(
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
enabled_worker_ids=None,
|
||||
delegate_master=None,
|
||||
trace_execution_id=None,
|
||||
):
|
||||
"""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,
|
||||
@@ -201,6 +262,7 @@ async def _select_execution_workers(
|
||||
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",
|
||||
@@ -225,6 +287,7 @@ async def _select_execution_workers(
|
||||
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(
|
||||
@@ -233,6 +296,7 @@ async def _select_execution_workers(
|
||||
)
|
||||
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,
|
||||
@@ -246,47 +310,29 @@ async def _select_execution_workers(
|
||||
active_workers = []
|
||||
delegate_master = False
|
||||
|
||||
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,
|
||||
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", {}),
|
||||
)
|
||||
|
||||
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():
|
||||
if job_id:
|
||||
runtime_state.distributed_job_allowed_workers[str(job_id)] = set(allowed)
|
||||
await _ensure_distributed_queue(job_id)
|
||||
|
||||
|
||||
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,
|
||||
@@ -298,171 +344,19 @@ def _build_master_prompt(
|
||||
prompt_index,
|
||||
)
|
||||
|
||||
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,
|
||||
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."
|
||||
)
|
||||
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,
|
||||
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, 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,
|
||||
)
|
||||
else:
|
||||
master_prompt = prepare_delegate_master_prompt(master_prompt, collector_ids)
|
||||
|
||||
if active_workers:
|
||||
trace_debug(
|
||||
@@ -470,29 +364,55 @@ async def orchestrate_distributed_execution(
|
||||
"Active distributed workers: "
|
||||
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers),
|
||||
)
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
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
|
||||
]
|
||||
)
|
||||
|
||||
prompt_id = await queue_prompt_payload(master_prompt, workflow_meta, client_id)
|
||||
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,
|
||||
)
|
||||
prompt_id = queue_result["prompt_id"]
|
||||
prompt_number = queue_result["number"]
|
||||
node_errors = queue_result.get("node_errors", {})
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
f"Orchestration complete: prompt_id={prompt_id}, dispatched_workers={len(worker_payloads)}, delegate_master={delegate_master}",
|
||||
)
|
||||
return prompt_id, len(worker_payloads)
|
||||
return prompt_id, prompt_number, len(worker_payloads), node_errors
|
||||
|
||||
+15
-20
@@ -1,20 +1,16 @@
|
||||
from dataclasses import dataclass
|
||||
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
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QueueRequestPayload:
|
||||
prompt: dict[str, Any]
|
||||
workflow_meta: dict[str, Any] | None
|
||||
prompt: Dict[str, Any]
|
||||
workflow_meta: Any
|
||||
client_id: str
|
||||
delegate_master: bool | None
|
||||
enabled_worker_ids: list[str]
|
||||
trace_execution_id: str | None
|
||||
delegate_master: Optional[bool]
|
||||
enabled_worker_ids: List[str]
|
||||
auto_prepare: bool
|
||||
trace_execution_id: Optional[str]
|
||||
|
||||
|
||||
def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
@@ -22,6 +18,11 @@ 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:
|
||||
@@ -34,10 +35,6 @@ 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:
|
||||
@@ -54,10 +51,7 @@ 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")
|
||||
try:
|
||||
enabled_ids = require_enabled_worker_ids(enabled_ids_raw)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"enabled_worker_ids invalid: {exc}") from exc
|
||||
enabled_ids = [str(worker_id).strip() for worker_id in enabled_ids_raw if str(worker_id).strip()]
|
||||
|
||||
delegate_master = data.get("delegate_master")
|
||||
if delegate_master is not None and not isinstance(delegate_master, bool):
|
||||
@@ -76,9 +70,10 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
|
||||
|
||||
return QueueRequestPayload(
|
||||
prompt=prompt,
|
||||
workflow_meta=workflow_meta,
|
||||
workflow_meta=data.get("workflow"),
|
||||
client_id=client_id,
|
||||
delegate_master=delegate_master,
|
||||
enabled_worker_ids=enabled_ids,
|
||||
auto_prepare=auto_prepare,
|
||||
trace_execution_id=trace_execution_id,
|
||||
)
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
"""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)
|
||||
+15
-61
@@ -1,29 +1,3 @@
|
||||
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):
|
||||
@@ -44,43 +18,26 @@ def require_fields(data: dict, *fields) -> list[str]:
|
||||
return missing
|
||||
|
||||
|
||||
def validate_worker_id(worker_id: str, config: dict[str, Any]) -> bool:
|
||||
def validate_worker_id(worker_id: str, config: dict) -> bool:
|
||||
"""Return True when worker_id exists in config['workers']."""
|
||||
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.")
|
||||
|
||||
worker_id_str = str(worker_id)
|
||||
workers = (config or {}).get("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.")
|
||||
return any(str(worker.get("id")) == worker_id_str for worker in workers)
|
||||
|
||||
|
||||
def require_positive_int(value, field_name: str = "value") -> int:
|
||||
"""Strict positive-int parser for request boundary validation."""
|
||||
def validate_positive_int(value, field_name: str) -> str | None:
|
||||
"""Validate positive integers and return an error string when invalid."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"Field '{field_name}' must be a positive integer.") from exc
|
||||
except (TypeError, ValueError):
|
||||
return f"Field '{field_name}' must be a positive integer."
|
||||
if parsed <= 0:
|
||||
raise ValueError(f"Field '{field_name}' must be a positive integer.")
|
||||
return parsed
|
||||
return f"Field '{field_name}' must be a positive integer."
|
||||
return None
|
||||
|
||||
|
||||
def coerce_positive_int(value, default: int) -> int:
|
||||
"""Permissive positive-int coercion used for non-boundary defaults."""
|
||||
def parse_positive_int(value, default: int) -> int:
|
||||
"""Parse value as positive int, returning default on failure."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError):
|
||||
@@ -88,13 +45,10 @@ def coerce_positive_int(value, default: int) -> int:
|
||||
return max(1, parsed)
|
||||
|
||||
|
||||
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
|
||||
def parse_positive_float(value, default: float) -> float:
|
||||
"""Parse value as positive float, returning default on failure."""
|
||||
try:
|
||||
parsed = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return fallback
|
||||
return parsed if parsed > 0.0 else fallback
|
||||
return max(0.0, float(default))
|
||||
return max(0.0, parsed)
|
||||
|
||||
+29
-41
@@ -1,63 +1,51 @@
|
||||
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: web.Request) -> web.StreamResponse:
|
||||
async def tunnel_status_endpoint(request):
|
||||
"""Return Cloudflare tunnel status and last known details."""
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
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,
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
return web.json_response({
|
||||
"status": "success",
|
||||
"tunnel": status,
|
||||
"master_host": master_host
|
||||
})
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/tunnel/start")
|
||||
async def tunnel_start_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def tunnel_start_endpoint(request):
|
||||
"""Start a Cloudflare tunnel pointing at the current ComfyUI server."""
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
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"),
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
|
||||
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)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/tunnel/stop")
|
||||
async def tunnel_stop_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def tunnel_stop_endpoint(request):
|
||||
"""Stop the managed Cloudflare tunnel if running."""
|
||||
|
||||
async def _operation() -> web.StreamResponse:
|
||||
try:
|
||||
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"),
|
||||
}
|
||||
)
|
||||
|
||||
return await run_authorized_endpoint(request, _operation)
|
||||
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)
|
||||
|
||||
+21
-137
@@ -1,5 +1,3 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import io
|
||||
import time
|
||||
@@ -8,78 +6,15 @@ 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 ensure_tile_jobs_initialized, init_dynamic_job
|
||||
from ..upscale.payload_parsers import parse_tiles_from_form
|
||||
from ..upscale.job_store import MAX_PAYLOAD_SIZE, ensure_tile_jobs_initialized
|
||||
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: web.Request) -> web.StreamResponse:
|
||||
auth_error = await authorization_error_or_none(request)
|
||||
if auth_error is not None:
|
||||
return auth_error
|
||||
async def heartbeat_endpoint(request):
|
||||
try:
|
||||
data = await request.json()
|
||||
worker_id = data.get('worker_id')
|
||||
@@ -93,8 +28,6 @@ async def heartbeat_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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"})
|
||||
@@ -105,38 +38,24 @@ async def heartbeat_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/submit_tiles")
|
||||
async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def submit_tiles_endpoint(request):
|
||||
"""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:
|
||||
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)
|
||||
if content_length and int(content_length) > 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')
|
||||
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)
|
||||
is_last = data.get('is_last', 'False').lower() == 'true'
|
||||
|
||||
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()
|
||||
|
||||
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)
|
||||
batch_size = int(data.get('batch_size', 0))
|
||||
|
||||
# Handle completion signal
|
||||
if batch_size == 0 and is_last:
|
||||
@@ -145,8 +64,6 @@ async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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,
|
||||
@@ -156,7 +73,7 @@ async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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)
|
||||
|
||||
@@ -166,8 +83,6 @@ async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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:
|
||||
@@ -191,28 +106,17 @@ async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/submit_image")
|
||||
async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def submit_image_endpoint(request):
|
||||
"""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:
|
||||
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)
|
||||
if content_length and int(content_length) > 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')
|
||||
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)
|
||||
is_last = data.get('is_last', 'False').lower() == 'true'
|
||||
|
||||
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)
|
||||
@@ -221,10 +125,7 @@ async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
# Handle image submission
|
||||
if 'full_image' in data and 'image_idx' in data:
|
||||
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)
|
||||
image_idx = int(data.get('image_idx'))
|
||||
img_data = data['full_image'].file.read()
|
||||
img = Image.open(io.BytesIO(img_data)).convert("RGB")
|
||||
|
||||
@@ -235,8 +136,6 @@ async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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,
|
||||
@@ -252,8 +151,6 @@ async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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,
|
||||
@@ -269,11 +166,8 @@ async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/request_image")
|
||||
async def request_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def request_image_endpoint(request):
|
||||
"""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')
|
||||
@@ -296,8 +190,6 @@ async def request_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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)
|
||||
@@ -307,33 +199,25 @@ async def request_image_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
if mode == 'dynamic':
|
||||
debug_log(f"UltimateSDUpscale API - Assigned image {task_idx} to worker {worker_id}")
|
||||
return web.json_response(
|
||||
{
|
||||
"kind": "image",
|
||||
"task_idx": task_idx,
|
||||
"estimated_remaining": remaining,
|
||||
}
|
||||
)
|
||||
return web.json_response({"image_idx": task_idx, "estimated_remaining": remaining})
|
||||
debug_log(f"UltimateSDUpscale API - Assigned tile {task_idx} to worker {worker_id}")
|
||||
return web.json_response({
|
||||
"kind": "tile",
|
||||
"task_idx": task_idx,
|
||||
"tile_idx": task_idx,
|
||||
"estimated_remaining": remaining,
|
||||
"batched_static": job_data.batched_static,
|
||||
})
|
||||
except asyncio.TimeoutError:
|
||||
return web.json_response({"kind": "none", "task_idx": None, "estimated_remaining": 0})
|
||||
if mode == 'dynamic':
|
||||
return web.json_response({"image_idx": None})
|
||||
return web.json_response({"tile_idx": None})
|
||||
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: web.Request) -> web.StreamResponse:
|
||||
async def job_status_endpoint(request):
|
||||
"""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})
|
||||
|
||||
+47
-108
@@ -1,13 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
import platform
|
||||
import subprocess # nosec B404 - commands are fixed and never shell-expanded
|
||||
import subprocess
|
||||
import socket
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import aiohttp
|
||||
@@ -25,27 +22,27 @@ from ..utils.network import (
|
||||
)
|
||||
from ..utils.constants import CHUNK_SIZE
|
||||
from ..workers import get_worker_manager
|
||||
from .request_guards import authorization_error_or_none
|
||||
from .schemas import (
|
||||
distributed_auth_headers,
|
||||
require_fields,
|
||||
require_worker_id,
|
||||
)
|
||||
from .schemas import require_fields, validate_worker_id
|
||||
from ..workers.detection import (
|
||||
get_machine_id,
|
||||
is_docker_environment,
|
||||
is_runpod_environment,
|
||||
)
|
||||
from ..utils.async_helpers import PromptValidationError, queue_prompt_payload
|
||||
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 {}
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/worker_ws")
|
||||
async def worker_ws_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def worker_ws_endpoint(request):
|
||||
"""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)
|
||||
|
||||
@@ -116,11 +113,8 @@ async def worker_ws_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/worker/clear_launching")
|
||||
async def clear_launching_state(request: web.Request) -> web.StreamResponse:
|
||||
async def clear_launching_state(request):
|
||||
"""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()
|
||||
@@ -130,10 +124,8 @@ async def clear_launching_state(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
worker_id = str(data.get("worker_id")).strip()
|
||||
config = load_config()
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
|
||||
# Clear launching flag in managed processes
|
||||
if worker_id in wm.processes:
|
||||
@@ -147,9 +139,8 @@ async def clear_launching_state(request: web.Request) -> web.StreamResponse:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
|
||||
def get_network_ips() -> list[str]:
|
||||
def get_network_ips():
|
||||
"""Get all network IPs, trying multiple methods."""
|
||||
command_timeout = 5.0
|
||||
ips = []
|
||||
hostname = socket.gethostname()
|
||||
|
||||
@@ -160,8 +151,8 @@ def get_network_ips() -> list[str]:
|
||||
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) as exc:
|
||||
debug_log(f"get_network_ips: getaddrinfo failed for hostname {hostname}: {exc}")
|
||||
except (socket.gaierror, OSError):
|
||||
pass
|
||||
|
||||
# Method 2: Try to connect to external server and get local IP
|
||||
try:
|
||||
@@ -171,19 +162,14 @@ def get_network_ips() -> list[str]:
|
||||
s.close()
|
||||
if local_ip not in ips:
|
||||
ips.append(local_ip)
|
||||
except (OSError, socket.error) as exc:
|
||||
debug_log(f"get_network_ips: UDP local IP probe failed: {exc}")
|
||||
except (OSError, socket.error):
|
||||
pass
|
||||
|
||||
# Method 3: Platform-specific commands
|
||||
try:
|
||||
if platform.system() == "Windows":
|
||||
# Windows ipconfig
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ipconfig"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
result = subprocess.run(["ipconfig"], capture_output=True, text=True)
|
||||
lines = result.stdout.split('\n')
|
||||
for i, line in enumerate(lines):
|
||||
if 'IPv4' in line and i + 1 < len(lines):
|
||||
@@ -193,20 +179,10 @@ def get_network_ips() -> list[str]:
|
||||
else:
|
||||
# Unix/Linux/Mac ifconfig or ip addr
|
||||
try:
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ip", "addr"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
result = subprocess.run(["ip", "addr"], capture_output=True, text=True)
|
||||
except (FileNotFoundError, OSError):
|
||||
try:
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
["ifconfig"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=command_timeout,
|
||||
)
|
||||
result = subprocess.run(["ifconfig"], capture_output=True, text=True)
|
||||
except (FileNotFoundError, OSError):
|
||||
result = None
|
||||
|
||||
@@ -217,13 +193,13 @@ def get_network_ips() -> list[str]:
|
||||
ip = match.group(1)
|
||||
if ip and ip not in ips:
|
||||
ips.append(ip)
|
||||
except (OSError, subprocess.SubprocessError) as exc:
|
||||
debug_log(f"get_network_ips: platform command probe failed: {exc}")
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
pass
|
||||
|
||||
return ips
|
||||
|
||||
|
||||
def get_recommended_ip(ips: list[str]) -> str | None:
|
||||
def get_recommended_ip(ips):
|
||||
"""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)
|
||||
@@ -258,7 +234,7 @@ def get_recommended_ip(ips: list[str]) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _get_cuda_info() -> tuple[int | None, int, int]:
|
||||
def _get_cuda_info():
|
||||
"""Detect CUDA device index and total physical GPU count.
|
||||
|
||||
Returns (cuda_device, cuda_device_count, physical_device_count).
|
||||
@@ -274,7 +250,7 @@ def _get_cuda_info() -> tuple[int | None, int, int]:
|
||||
if visible_devices:
|
||||
cuda_device = visible_devices[0]
|
||||
try:
|
||||
result = subprocess.run( # nosec B603 - static command, no user input
|
||||
result = subprocess.run(
|
||||
['nvidia-smi', '--query-gpu=name', '--format=csv,noheader'],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
@@ -298,7 +274,7 @@ def _get_cuda_info() -> tuple[int | None, int, int]:
|
||||
return None, 0, 0
|
||||
|
||||
|
||||
def _collect_network_info_sync() -> dict[str, Any]:
|
||||
def _collect_network_info_sync():
|
||||
"""Collect network/cuda info in a worker thread to avoid blocking route handlers."""
|
||||
cuda_device, cuda_device_count, physical_device_count = _get_cuda_info()
|
||||
hostname = socket.gethostname()
|
||||
@@ -313,7 +289,7 @@ def _collect_network_info_sync() -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _read_worker_log_sync(log_file: str, lines_to_read: int) -> dict[str, Any]:
|
||||
def _read_worker_log_sync(log_file, lines_to_read):
|
||||
"""Read worker log content from disk in a threadpool worker."""
|
||||
file_size = os.path.getsize(log_file)
|
||||
|
||||
@@ -349,12 +325,7 @@ def _read_worker_log_sync(log_file: str, lines_to_read: int) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _parse_positive_int_query(
|
||||
value: Any,
|
||||
default: int,
|
||||
minimum: int = 1,
|
||||
maximum: int | None = 10000,
|
||||
) -> int:
|
||||
def _parse_positive_int_query(value, default, minimum=1, maximum=10000):
|
||||
"""Parse bounded positive integer query params with sane fallback."""
|
||||
try:
|
||||
parsed = int(value)
|
||||
@@ -366,7 +337,7 @@ def _parse_positive_int_query(
|
||||
return parsed
|
||||
|
||||
|
||||
def _find_worker_by_id(config: dict[str, Any], worker_id: str) -> dict[str, Any] | None:
|
||||
def _find_worker_by_id(config, worker_id):
|
||||
worker_id_str = str(worker_id).strip()
|
||||
for worker in config.get("workers", []):
|
||||
if str(worker.get("id")).strip() == worker_id_str:
|
||||
@@ -375,11 +346,8 @@ def _find_worker_by_id(config: dict[str, Any], worker_id: str) -> dict[str, Any]
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/local_log")
|
||||
async def get_local_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_local_log_endpoint(request):
|
||||
"""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:
|
||||
@@ -423,11 +391,8 @@ async def get_local_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/network_info")
|
||||
async def get_network_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_network_info_endpoint(request):
|
||||
"""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)
|
||||
@@ -441,11 +406,8 @@ async def get_network_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/system_info")
|
||||
async def get_system_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_system_info_endpoint(request):
|
||||
"""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
|
||||
|
||||
@@ -468,11 +430,8 @@ async def get_system_info_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
return await handle_api_error(request, e, 500)
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/launch_worker")
|
||||
async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def launch_worker_endpoint(request):
|
||||
"""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()
|
||||
@@ -484,10 +443,8 @@ async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
# Find worker config
|
||||
config = load_config()
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 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)
|
||||
@@ -506,7 +463,7 @@ async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
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)
|
||||
@@ -534,11 +491,8 @@ async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.post("/distributed/stop_worker")
|
||||
async def stop_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def stop_worker_endpoint(request):
|
||||
"""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()
|
||||
@@ -548,10 +502,8 @@ async def stop_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
worker_id = str(data.get("worker_id")).strip()
|
||||
config = load_config()
|
||||
try:
|
||||
worker_id = require_worker_id(worker_id, config)
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 404)
|
||||
if not validate_worker_id(worker_id, config):
|
||||
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
|
||||
|
||||
success, message = wm.stop_worker(worker_id)
|
||||
|
||||
@@ -569,11 +521,8 @@ async def stop_worker_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/managed_workers")
|
||||
async def get_managed_workers_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_managed_workers_endpoint(request):
|
||||
"""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({
|
||||
@@ -585,11 +534,8 @@ async def get_managed_workers_endpoint(request: web.Request) -> web.StreamRespon
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/local-worker-status")
|
||||
async def get_local_worker_status_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_local_worker_status_endpoint(request):
|
||||
"""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 = {}
|
||||
@@ -658,11 +604,8 @@ async def get_local_worker_status_endpoint(request: web.Request) -> web.StreamRe
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/worker_log/{worker_id}")
|
||||
async def get_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_worker_log_endpoint(request):
|
||||
"""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']
|
||||
@@ -704,11 +647,8 @@ async def get_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
|
||||
|
||||
@server.PromptServer.instance.routes.get("/distributed/remote_worker_log/{worker_id}")
|
||||
async def get_remote_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
|
||||
async def get_remote_worker_log_endpoint(request):
|
||||
"""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()
|
||||
@@ -731,7 +671,6 @@ async def get_remote_worker_log_endpoint(request: web.Request) -> web.StreamResp
|
||||
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:
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
"""Bootstrap entrypoints for ComfyUI-Distributed."""
|
||||
|
||||
from .entrypoint import build_node_mappings, initialize_runtime
|
||||
|
||||
__all__ = ["build_node_mappings", "initialize_runtime"]
|
||||
@@ -1,66 +0,0 @@
|
||||
"""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())}")
|
||||
+25
-1
@@ -1,4 +1,28 @@
|
||||
"""Project-level pytest collection rules for plugin-style package layout."""
|
||||
# conftest.py — project-level pytest configuration.
|
||||
#
|
||||
# Problem: custom_nodes/ComfyUI-Distributed/__init__.py uses relative imports
|
||||
# (from .distributed 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
|
||||
|
||||
collect_ignore = [
|
||||
"__init__.py",
|
||||
|
||||
+47
-11
@@ -1,15 +1,51 @@
|
||||
"""Compatibility facade for legacy imports from `distributed`."""
|
||||
from __future__ import annotations
|
||||
"""
|
||||
ComfyUI-Distributed: thin entry point.
|
||||
All implementation lives in workers/, nodes/, api/.
|
||||
"""
|
||||
import atexit
|
||||
import os
|
||||
|
||||
from .bootstrap.entrypoint import (
|
||||
build_node_mappings,
|
||||
initialize_runtime,
|
||||
import server
|
||||
|
||||
from .utils.config import ensure_config_exists
|
||||
from .utils.logging import debug_log
|
||||
from .utils.network import cleanup_client_session
|
||||
from .workers import get_worker_manager
|
||||
from .workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
from .upscale.job_store import ensure_tile_jobs_initialized
|
||||
from .nodes import (
|
||||
NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS,
|
||||
ImageBatchDivider,
|
||||
DistributedCollectorNode,
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from . import api # noqa: F401 - triggers all @routes.* registrations
|
||||
from .api.queue_orchestration import ensure_distributed_state
|
||||
|
||||
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
|
||||
ensure_config_exists()
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"initialize_runtime",
|
||||
]
|
||||
# Aiohttp session cleanup
|
||||
async def _cleanup_session():
|
||||
await cleanup_client_session()
|
||||
|
||||
|
||||
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
|
||||
|
||||
# Initialize distributed job state on prompt_server
|
||||
prompt_server = server.PromptServer.instance
|
||||
ensure_distributed_state(prompt_server)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# Worker startup
|
||||
if not os.environ.get('COMFYUI_IS_WORKER'):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
|
||||
@@ -53,6 +53,7 @@ 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"
|
||||
}
|
||||
```
|
||||
@@ -73,6 +74,10 @@ 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>]`.
|
||||
@@ -100,7 +105,8 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
|
||||
```json
|
||||
{
|
||||
"prompt_id": "<uuid>",
|
||||
"worker_count": 2
|
||||
"worker_count": 2,
|
||||
"auto_prepare_supported": true
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
@@ -10,13 +10,9 @@ from .utilities import (
|
||||
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,
|
||||
@@ -26,8 +22,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DistributedCollector": "Distributed Collector",
|
||||
"DistributedBranch": "Distributed Branch",
|
||||
"DistributedBranchCollector": "Distributed Branch Collector",
|
||||
"DistributedSeed": "Distributed Seed",
|
||||
"DistributedModelName": "Distributed Model Name",
|
||||
"DistributedValue": "Distributed Value",
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
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)
|
||||
@@ -1,298 +0,0 @@
|
||||
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)
|
||||
+304
-309
@@ -1,49 +1,32 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import torch
|
||||
import io
|
||||
import json
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
import base64
|
||||
|
||||
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(cls: type["DistributedCollectorNode"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -56,65 +39,108 @@ class DistributedCollectorNode:
|
||||
),
|
||||
},
|
||||
"optional": { "audio": ("AUDIO",) },
|
||||
"hidden": build_distributed_hidden_inputs(
|
||||
include_worker_batch_size=True,
|
||||
include_pass_through=True,
|
||||
),
|
||||
"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}),
|
||||
},
|
||||
}
|
||||
|
||||
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 _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_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)
|
||||
|
||||
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,
|
||||
)
|
||||
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, 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):
|
||||
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)
|
||||
|
||||
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 context.multi_job_id or context.pass_through:
|
||||
if context.pass_through:
|
||||
if not multi_job_id or pass_through:
|
||||
if 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=images,
|
||||
audio=audio,
|
||||
load_balance=load_balance,
|
||||
context=context,
|
||||
images,
|
||||
audio,
|
||||
load_balance,
|
||||
multi_job_id,
|
||||
is_worker,
|
||||
master_url,
|
||||
enabled_worker_ids,
|
||||
worker_batch_size,
|
||||
worker_id,
|
||||
delegate_only,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
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:
|
||||
async def send_batch_to_master(self, image_batch, audio, multi_job_id, master_url, worker_id):
|
||||
"""Send image batch to master via canonical JSON envelopes."""
|
||||
batch_size = image_batch.shape[0]
|
||||
if batch_size == 0:
|
||||
@@ -143,7 +169,6 @@ class DistributedCollectorNode:
|
||||
async with session.post(
|
||||
url,
|
||||
json=payload,
|
||||
headers=distributed_auth_headers(load_config()),
|
||||
timeout=aiohttp.ClientTimeout(total=60),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
@@ -269,265 +294,235 @@ class DistributedCollectorNode:
|
||||
else:
|
||||
raise ValueError("No image data collected from master or 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}")
|
||||
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
|
||||
debug_log(f"Worker - Job {multi_job_id} complete. Sending {images.shape[0]} 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)
|
||||
|
||||
# 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")
|
||||
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})")
|
||||
images_on_cpu = 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...")
|
||||
|
||||
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)
|
||||
# Ensure master images are contiguous
|
||||
images_on_cpu = ensure_contiguous(images_on_cpu)
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
# 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()
|
||||
|
||||
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))
|
||||
# NEW: Initialize progress bar for workers (total = num_workers)
|
||||
p = ProgressBar(num_workers)
|
||||
|
||||
def mark_worker_done(done_worker_id):
|
||||
done_worker_id = str(done_worker_id)
|
||||
if done_worker_id not in expected_workers:
|
||||
debug_log(
|
||||
"Collector probe: worker "
|
||||
f"{wid} online={payload is not None} queue_remaining={queue_remaining}"
|
||||
f"Master - Ignoring completion from unexpected worker {done_worker_id} for job {multi_job_id}"
|
||||
)
|
||||
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}")
|
||||
return any_busy
|
||||
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
|
||||
|
||||
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}"
|
||||
)
|
||||
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)
|
||||
|
||||
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:
|
||||
while len(workers_done) < num_workers:
|
||||
# Check for user interruption to abort collection promptly
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
continue
|
||||
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}"
|
||||
)
|
||||
|
||||
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)"
|
||||
# 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"
|
||||
)
|
||||
log(
|
||||
f"Master - Heartbeat timeout. Still waiting for workers: {list(missing_workers)} "
|
||||
f"(elapsed={elapsed:.1f}s)"
|
||||
)
|
||||
|
||||
# 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]
|
||||
|
||||
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})")
|
||||
|
||||
if await self._probe_missing_workers_busy(missing_workers):
|
||||
last_activity = time.time()
|
||||
base_timeout = float(get_worker_timeout_seconds())
|
||||
continue
|
||||
# Combine audio from master and workers
|
||||
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
|
||||
|
||||
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,
|
||||
)
|
||||
return (combined, combined_audio)
|
||||
except Exception as e:
|
||||
log(f"Master - Error combining images: {e}")
|
||||
# Return just the master images as fallback
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
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)),
|
||||
}
|
||||
+82
-320
@@ -1,28 +1,25 @@
|
||||
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: Callable[..., Any]) -> Callable[..., Any]:
|
||||
def sync_wrapper(async_func):
|
||||
"""Decorator to wrap async methods for synchronous execution."""
|
||||
@wraps(async_func)
|
||||
def sync_func(self, *args: Any, **kwargs: Any) -> Any:
|
||||
def sync_func(self, *args, **kwargs):
|
||||
# Use run_async_in_server_loop for ComfyUI compatibility
|
||||
return run_async_in_server_loop(
|
||||
async_func(self, *args, **kwargs),
|
||||
@@ -30,10 +27,21 @@ def sync_wrapper(async_func: Callable[..., Any]) -> Callable[..., Any]:
|
||||
)
|
||||
return sync_func
|
||||
|
||||
|
||||
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)
|
||||
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 []
|
||||
|
||||
class UltimateSDUpscaleDistributed(
|
||||
DynamicModeMixin,
|
||||
@@ -73,14 +81,7 @@ class UltimateSDUpscaleDistributed(
|
||||
debug_log("UltimateSDUpscaleDistributed - Node initialized")
|
||||
|
||||
@classmethod
|
||||
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}),
|
||||
}
|
||||
)
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"upscaled_image": ("IMAGE",),
|
||||
@@ -101,7 +102,15 @@ class UltimateSDUpscaleDistributed(
|
||||
"force_uniform_tiles": ("BOOLEAN", {"default": True}),
|
||||
"tiled_decode": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": hidden_inputs,
|
||||
"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}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -113,47 +122,12 @@ class UltimateSDUpscaleDistributed(
|
||||
"""Force re-execution."""
|
||||
return float("nan") # Always re-execute
|
||||
|
||||
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, ...]:
|
||||
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):
|
||||
"""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])
|
||||
@@ -168,15 +142,9 @@ class UltimateSDUpscaleDistributed(
|
||||
)
|
||||
if not multi_job_id:
|
||||
# No distributed processing, run single GPU version
|
||||
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,
|
||||
)
|
||||
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)
|
||||
|
||||
if is_worker:
|
||||
# Worker mode: process tiles synchronously
|
||||
@@ -187,64 +155,24 @@ class UltimateSDUpscaleDistributed(
|
||||
worker_id, enabled_worker_ids, dynamic_threshold)
|
||||
else:
|
||||
# Master mode: distribute and collect synchronously
|
||||
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,
|
||||
)
|
||||
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)
|
||||
|
||||
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, ...]:
|
||||
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):
|
||||
"""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 = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
enabled_workers = json.loads(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
|
||||
@@ -253,47 +181,31 @@ 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 for single-image runs.
|
||||
if num_workers > 0 and batch_size <= 1 and num_tiles_per_image > 1:
|
||||
# and there is more than one tile to process, even if batch == 1.
|
||||
if num_workers > 0 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=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,
|
||||
)
|
||||
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)
|
||||
|
||||
# Static mode - enhanced with health monitoring and retry logic
|
||||
return self._process_worker_static_sync(upscaled_image, core_args,
|
||||
return self._process_worker_static_sync(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, multi_job_id, master_url,
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
worker_id, enabled_workers)
|
||||
|
||||
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, ...]:
|
||||
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):
|
||||
"""Unified master processing with enhanced monitoring and failure handling."""
|
||||
# Round tile dimensions
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
@@ -312,206 +224,56 @@ class UltimateSDUpscaleDistributed(
|
||||
)
|
||||
|
||||
# Parse enabled workers
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
enabled_workers = json.loads(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,
|
||||
# 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:
|
||||
# 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:
|
||||
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=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,
|
||||
)
|
||||
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)
|
||||
|
||||
elif mode == "dynamic":
|
||||
# Dynamic mode for large batches
|
||||
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,
|
||||
)
|
||||
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)
|
||||
|
||||
# Static mode - enhanced with unified job management
|
||||
return self._process_master_static_sync(upscaled_image, core_args,
|
||||
return self._process_master_static_sync(upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, multi_job_id, enabled_workers,
|
||||
force_uniform_tiles, tiled_decode, 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:
|
||||
"""Determine mode from worker availability and configured threshold."""
|
||||
"""Determines processing mode per requested policy:
|
||||
- any workers => prefer static (tile-based) for USDU
|
||||
- no workers => single_gpu
|
||||
"""
|
||||
if num_workers == 0:
|
||||
return "single_gpu"
|
||||
threshold = max(1, int(dynamic_threshold))
|
||||
if int(batch_size) >= threshold:
|
||||
return "dynamic"
|
||||
# Default to static when distributed; master/worker may still override if special cases arise
|
||||
return "static"
|
||||
|
||||
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)
|
||||
|
||||
# Ensure initialization before registering routes
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# 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
|
||||
}
|
||||
|
||||
@@ -1,38 +0,0 @@
|
||||
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
|
||||
@@ -1,46 +0,0 @@
|
||||
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
|
||||
@@ -1,16 +0,0 @@
|
||||
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
|
||||
+48
-73
@@ -1,10 +1,7 @@
|
||||
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]]:
|
||||
@@ -31,7 +28,7 @@ class DistributedSeed:
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedSeed"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"seed": ("INT", {
|
||||
@@ -42,7 +39,8 @@ class DistributedSeed:
|
||||
}),
|
||||
},
|
||||
"hidden": {
|
||||
**build_worker_identity_hidden_inputs(),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -51,29 +49,30 @@ class DistributedSeed:
|
||||
FUNCTION = "distribute"
|
||||
CATEGORY = "utils"
|
||||
|
||||
def distribute(
|
||||
self,
|
||||
seed: int,
|
||||
is_worker: bool = False,
|
||||
worker_id: str = "",
|
||||
enabled_worker_ids: str = "[]",
|
||||
) -> tuple[int]:
|
||||
def distribute(self, seed, is_worker=False, worker_id=""):
|
||||
if not is_worker:
|
||||
# Master node: pass through original values
|
||||
debug_log(f"Distributor - Master: seed={seed}")
|
||||
return (seed,)
|
||||
else:
|
||||
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
|
||||
worker_index = parse_worker_index(worker_id, enabled_workers)
|
||||
if worker_index is not None:
|
||||
# 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)
|
||||
|
||||
offset = worker_index + 1
|
||||
new_seed = seed + offset
|
||||
debug_log(f"Distributor - Worker {worker_index}: seed={seed} → {new_seed}")
|
||||
return (new_seed,)
|
||||
|
||||
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': no worker index resolved")
|
||||
# Fallback: return original seed
|
||||
return (seed,)
|
||||
except (ValueError, IndexError) as e:
|
||||
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': {e}")
|
||||
# Fallback: return original seed
|
||||
return (seed,)
|
||||
|
||||
|
||||
# Define ByPassTypeTuple for flexible return types
|
||||
@@ -93,14 +92,15 @@ class DistributedValue:
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedValue"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"default_value": ("STRING", {"default": ""}),
|
||||
"worker_values": ("STRING", {"default": "{}"}),
|
||||
},
|
||||
"hidden": {
|
||||
**build_worker_identity_hidden_inputs(),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
"worker_id": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -110,7 +110,7 @@ class DistributedValue:
|
||||
CATEGORY = "utils"
|
||||
|
||||
@staticmethod
|
||||
def _coerce(value: Any, value_type: str) -> Any:
|
||||
def _coerce(value, value_type):
|
||||
"""Convert a string value to the requested type."""
|
||||
if value_type == "INT":
|
||||
return int(float(value))
|
||||
@@ -119,71 +119,51 @@ class DistributedValue:
|
||||
return value # STRING and COMBO stay as strings
|
||||
|
||||
@staticmethod
|
||||
def _coerce_safe(value: Any, value_type: str) -> Any:
|
||||
def _coerce_safe(value, value_type):
|
||||
"""Best-effort coercion with graceful fallback to original value."""
|
||||
try:
|
||||
return DistributedValue._coerce(value, value_type)
|
||||
except (TypeError, ValueError):
|
||||
return value
|
||||
|
||||
@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]:
|
||||
def distribute(self, default_value, worker_values="{}", is_worker=False, worker_id=""):
|
||||
values = {}
|
||||
value_type = "STRING"
|
||||
|
||||
try:
|
||||
raw_values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
|
||||
values = dict(raw_values) if isinstance(raw_values, dict) else {}
|
||||
values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
|
||||
if not isinstance(values, dict):
|
||||
values = {}
|
||||
except json.JSONDecodeError as e:
|
||||
debug_log(f"DistributedValue - Error parsing worker_values: {e}")
|
||||
values = {}
|
||||
|
||||
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
|
||||
value_type = values.get("_type", "STRING")
|
||||
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,)
|
||||
|
||||
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:
|
||||
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:
|
||||
coerced = self._coerce(raw, value_type)
|
||||
debug_log(f"DistributedValue - Worker key {lookup_key}: returning '{coerced}'")
|
||||
debug_log(f"DistributedValue - Worker {idx}: returning '{coerced}'")
|
||||
return (coerced,)
|
||||
except (TypeError, ValueError) as e:
|
||||
debug_log(f"DistributedValue - Error coercing worker value for key {lookup_key}: {e}")
|
||||
except (ValueError, IndexError) as e:
|
||||
debug_log(f"DistributedValue - Error: {e}")
|
||||
debug_log(f"DistributedValue - Worker fallback: returning default '{coerced_default}'")
|
||||
return (coerced_default,)
|
||||
|
||||
class DistributedModelName:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedModelName"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"default": ""}),
|
||||
@@ -200,7 +180,7 @@ class DistributedModelName:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "utils"
|
||||
|
||||
def _stringify(self, value: Any) -> str:
|
||||
def _stringify(self, value):
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, (int, float, bool)):
|
||||
@@ -210,7 +190,7 @@ class DistributedModelName:
|
||||
except Exception:
|
||||
return str(value)
|
||||
|
||||
def _update_workflow(self, extra_pnginfo: Any, unique_id: Any, values: list[str]) -> None:
|
||||
def _update_workflow(self, extra_pnginfo, unique_id, values):
|
||||
if not extra_pnginfo:
|
||||
return
|
||||
info = extra_pnginfo[0] if isinstance(extra_pnginfo, list) else extra_pnginfo
|
||||
@@ -228,12 +208,7 @@ class DistributedModelName:
|
||||
if node:
|
||||
node["widgets_values"] = [values]
|
||||
|
||||
def log_input(
|
||||
self,
|
||||
text: Any,
|
||||
unique_id: Any = None,
|
||||
extra_pnginfo: Any = None,
|
||||
) -> dict[str, Any]:
|
||||
def log_input(self, text, unique_id=None, extra_pnginfo=None):
|
||||
values = []
|
||||
if isinstance(text, list):
|
||||
for val in text:
|
||||
@@ -249,7 +224,7 @@ class DistributedModelName:
|
||||
return {"ui": {"text": values}, "result": (values,)}
|
||||
|
||||
class ByPassTypeTuple(tuple):
|
||||
def __getitem__(self, index: int) -> Any:
|
||||
def __getitem__(self, index):
|
||||
if index > 0:
|
||||
index = 0
|
||||
item = super().__getitem__(index)
|
||||
@@ -259,7 +234,7 @@ class ByPassTypeTuple(tuple):
|
||||
|
||||
class ImageBatchDivider:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["ImageBatchDivider"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -280,7 +255,7 @@ class ImageBatchDivider:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "image"
|
||||
|
||||
def divide_batch(self, images: torch.Tensor, divide_by: int) -> tuple[torch.Tensor, ...]:
|
||||
def divide_batch(self, images, divide_by):
|
||||
total_splits = max(1, min(int(divide_by), 10))
|
||||
total_frames = images.shape[0]
|
||||
empty_tensor = images[:0]
|
||||
@@ -297,7 +272,7 @@ class AudioBatchDivider:
|
||||
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["AudioBatchDivider"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"audio": ("AUDIO",),
|
||||
@@ -318,7 +293,7 @@ class AudioBatchDivider:
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "audio"
|
||||
|
||||
def divide_audio(self, audio: dict[str, Any], divide_by: int) -> tuple[dict[str, Any], ...]:
|
||||
def divide_audio(self, audio, divide_by):
|
||||
import torch
|
||||
|
||||
waveform = audio.get("waveform")
|
||||
@@ -358,7 +333,7 @@ class DistributedEmptyImage:
|
||||
"""Produces an empty IMAGE batch used when the master delegates all work."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls: type["DistributedEmptyImage"]) -> dict[str, Any]:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"height": ("INT", {"default": 64, "min": 1, "max": 4096, "step": 1}),
|
||||
@@ -371,7 +346,7 @@ class DistributedEmptyImage:
|
||||
FUNCTION = "create"
|
||||
CATEGORY = "image"
|
||||
|
||||
def create(self, height: int, width: int, channels: int) -> tuple[torch.Tensor]:
|
||||
def create(self, height, width, channels):
|
||||
import torch
|
||||
|
||||
shape = (0, height, width, channels)
|
||||
|
||||
+2
-5
@@ -1,12 +1,9 @@
|
||||
[project]
|
||||
name = "ComfyUI-Distributed"
|
||||
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
|
||||
version = "1.4.1"
|
||||
version = "1.4.4"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = [
|
||||
"aiohttp>=3.9,<4",
|
||||
"Pillow>=10,<12",
|
||||
]
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/robertvoy/ComfyUI-Distributed"
|
||||
|
||||
@@ -1,163 +0,0 @@
|
||||
"""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
|
||||
@@ -1,423 +0,0 @@
|
||||
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,19 +3,9 @@ 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):
|
||||
@@ -35,18 +25,45 @@ def _load_config_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "config_routes.py"
|
||||
package_name = "dist_api_config_testpkg"
|
||||
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
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
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
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()
|
||||
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
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
@@ -56,16 +73,7 @@ def _load_config_routes_module():
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
|
||||
async def _handle_api_error(_request, error, status=500):
|
||||
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,
|
||||
)
|
||||
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
|
||||
|
||||
network_module.handle_api_error = _handle_api_error
|
||||
network_module.normalize_host = lambda value: value
|
||||
@@ -81,12 +89,6 @@ 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)
|
||||
@@ -94,7 +96,8 @@ def _load_config_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
|
||||
return module
|
||||
|
||||
@@ -115,7 +118,9 @@ 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):
|
||||
with patch.object(config_routes, "load_config", return_value=cfg), patch.object(
|
||||
config_routes, "save_config", return_value=True
|
||||
):
|
||||
response = await config_routes.update_config_endpoint(_FakeRequest({"debug": True}))
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
|
||||
@@ -11,14 +11,6 @@ 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):
|
||||
@@ -27,9 +19,8 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
def __init__(self, payload, headers=None):
|
||||
def __init__(self, payload):
|
||||
self._payload = payload
|
||||
self.headers = headers or {}
|
||||
|
||||
async def json(self):
|
||||
return self._payload
|
||||
@@ -39,19 +30,53 @@ def _load_job_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "job_routes.py"
|
||||
package_name = "dist_api_queue_testpkg"
|
||||
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True)
|
||||
# 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
|
||||
|
||||
# aiohttp.web stub
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
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 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={},
|
||||
)
|
||||
install_server_stub(prompt_server_instance)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
# torch stub (only needed to satisfy import)
|
||||
created_torch_stub = False
|
||||
@@ -132,53 +157,12 @@ 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", 1))
|
||||
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 7, 1, {}))
|
||||
sys.modules[f"{package_name}.api.queue_orchestration"] = queue_orchestration_module
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -188,6 +172,7 @@ 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):
|
||||
@@ -208,6 +193,7 @@ 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"),
|
||||
)
|
||||
|
||||
@@ -220,11 +206,13 @@ def _load_job_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
cleanup_optional_module("torch", created_torch_stub)
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
if created_pil_stub:
|
||||
cleanup_optional_module("PIL.Image", True)
|
||||
cleanup_optional_module("PIL", True)
|
||||
sys.modules.pop("PIL.Image", None)
|
||||
sys.modules.pop("PIL", None)
|
||||
|
||||
return module
|
||||
|
||||
@@ -233,7 +221,7 @@ job_routes = _load_job_routes_module()
|
||||
|
||||
|
||||
class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_id(self):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_metadata(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"prompt": {"1": {"class_type": "Node"}},
|
||||
@@ -245,12 +233,15 @@ class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
with patch.object(
|
||||
job_routes,
|
||||
"orchestrate_distributed_execution",
|
||||
new=AsyncMock(return_value=("prompt_123", 2)),
|
||||
new=AsyncMock(return_value=("prompt_123", 42, 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(
|
||||
@@ -287,10 +278,8 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
|
||||
async def test_job_complete_accepts_audio_payload(self):
|
||||
queue = asyncio.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"}}
|
||||
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
job_routes.prompt_server.distributed_pending_jobs = {"job-1": queue}
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
@@ -313,28 +302,6 @@ 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_rejects_worker_not_in_job_allowlist(self):
|
||||
queue = asyncio.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": "job-allow",
|
||||
"worker_id": "worker-unexpected",
|
||||
"batch_idx": 0,
|
||||
"image": "data:image/png;base64,AAAA",
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
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, 403)
|
||||
self.assertIn("unauthorized", response.payload.get("message", "").lower())
|
||||
|
||||
def test_decode_audio_payload_rejects_bad_shape(self):
|
||||
bad = {
|
||||
"sample_rate": 44100,
|
||||
@@ -4,29 +4,30 @@ 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"
|
||||
|
||||
bootstrap_test_package(
|
||||
package_name,
|
||||
with_api=True,
|
||||
with_utils=True,
|
||||
with_orchestration=True,
|
||||
)
|
||||
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
|
||||
@@ -47,17 +48,22 @@ def _load_media_sync_module():
|
||||
trace_module.trace_info = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.trace_logger"] = trace_module
|
||||
|
||||
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
|
||||
created_aiohttp_stub = False
|
||||
if "aiohttp" not in sys.modules:
|
||||
created_aiohttp_stub = True
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
|
||||
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 _ClientTimeout:
|
||||
def __init__(self, total=None):
|
||||
pass
|
||||
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _AiohttpResponse(payload, status=status)
|
||||
)
|
||||
class _FormData:
|
||||
def add_field(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
aiohttp_module.ClientTimeout = _ClientTimeout
|
||||
aiohttp_module.FormData = _FormData
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.api.orchestration.media_sync",
|
||||
@@ -67,7 +73,8 @@ def _load_media_sync_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
|
||||
return module
|
||||
|
||||
@@ -253,7 +260,3 @@ class RewritePromptMediaInputsTests(unittest.TestCase):
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
class _AiohttpResponse:
|
||||
def __init__(self, payload, status=200):
|
||||
self.payload = payload
|
||||
self.status = status
|
||||
|
||||
@@ -1,195 +0,0 @@
|
||||
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()
|
||||
@@ -1,95 +0,0 @@
|
||||
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,14 +9,6 @@ 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):
|
||||
@@ -38,30 +30,43 @@ 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"
|
||||
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_upscale=True)
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_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.")
|
||||
api_pkg = types.ModuleType(f"{package_name}.api")
|
||||
api_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.api"] = api_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
|
||||
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
|
||||
upscale_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.upscale"] = upscale_pkg
|
||||
|
||||
install_request_guards_stub(package_name)
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
prompt_server_holder = {
|
||||
"value": types.SimpleNamespace(
|
||||
@@ -70,23 +75,23 @@ def _load_usdu_routes_module():
|
||||
)
|
||||
}
|
||||
|
||||
created_aiohttp_stub = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
install_server_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
|
||||
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
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):
|
||||
@@ -98,21 +103,6 @@ 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")
|
||||
@@ -160,7 +150,6 @@ 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
|
||||
|
||||
@@ -169,7 +158,8 @@ def _load_usdu_routes_module():
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
|
||||
module.web = types.SimpleNamespace(
|
||||
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
@@ -232,8 +222,7 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
response = await usdu_routes.request_image_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("kind"), "image")
|
||||
self.assertEqual(response.payload.get("task_idx"), 7)
|
||||
self.assertEqual(response.payload.get("image_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)
|
||||
@@ -249,8 +238,7 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
response = await usdu_routes.request_image_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("kind"), "tile")
|
||||
self.assertEqual(response.payload.get("task_idx"), 4)
|
||||
self.assertEqual(response.payload.get("tile_idx"), 4)
|
||||
self.assertTrue(response.payload.get("batched_static"))
|
||||
self.assertEqual(job_data.assigned_to_workers["worker-a"], [4])
|
||||
|
||||
|
||||
@@ -8,14 +8,6 @@ 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):
|
||||
@@ -24,11 +16,10 @@ class _FakeResponse:
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
def __init__(self, payload=None, match_info=None, query=None, headers=None):
|
||||
def __init__(self, payload=None, match_info=None, query=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
|
||||
@@ -58,8 +49,8 @@ class _FakeHTTPClientSession:
|
||||
self._status = status
|
||||
self.calls = []
|
||||
|
||||
def get(self, url, params=None, headers=None, timeout=None):
|
||||
self.calls.append({"url": url, "params": params, "headers": headers, "timeout": timeout})
|
||||
def get(self, url, params=None, timeout=None):
|
||||
self.calls.append({"url": url, "params": params, "timeout": timeout})
|
||||
return _FakeHTTPClientResponse(self._payload, status=self._status)
|
||||
|
||||
|
||||
@@ -71,12 +62,12 @@ class _DummyWorkerManager:
|
||||
worker_id = str(worker["id"])
|
||||
self.processes[worker_id] = {
|
||||
"pid": 12345,
|
||||
"log_file": f"/tmp/distributed_worker_{worker_id}.log", # nosec B108 - deterministic fake test path
|
||||
"log_file": f"/tmp/distributed_worker_{worker_id}.log",
|
||||
"process": None,
|
||||
}
|
||||
return 12345
|
||||
|
||||
def is_process_running(self, _pid):
|
||||
def _is_process_running(self, _pid):
|
||||
return False
|
||||
|
||||
def save_processes(self):
|
||||
@@ -98,7 +89,21 @@ def _load_worker_routes_module():
|
||||
module_path = Path(__file__).resolve().parents[2] / "api" / "worker_routes.py"
|
||||
package_name = "dist_api_worker_testpkg"
|
||||
|
||||
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_workers=True)
|
||||
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
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
@@ -114,10 +119,59 @@ 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 = install_aiohttp_stub(
|
||||
lambda payload, status=200: _FakeResponse(payload, status=status)
|
||||
)
|
||||
install_server_stub()
|
||||
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_torch_stub = False
|
||||
if "torch" not in sys.modules:
|
||||
@@ -194,29 +248,19 @@ 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)
|
||||
|
||||
cleanup_optional_module("aiohttp", created_aiohttp_stub)
|
||||
cleanup_optional_module("torch", created_torch_stub)
|
||||
if created_aiohttp_stub:
|
||||
sys.modules.pop("aiohttp", None)
|
||||
if created_torch_stub:
|
||||
sys.modules.pop("torch", None)
|
||||
|
||||
return module
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# This conftest.py marks the tests/ directory as the pytest collection root,
|
||||
# preventing pytest from traversing into the parent package's __init__.py.
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
@@ -6,7 +5,7 @@ import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class _FakePromptQueue:
|
||||
class _PromptQueue:
|
||||
def __init__(self):
|
||||
self.items = []
|
||||
|
||||
@@ -14,16 +13,10 @@ class _FakePromptQueue:
|
||||
self.items.append(item)
|
||||
|
||||
|
||||
class _FakePromptServer:
|
||||
def __init__(self):
|
||||
self.number = 0
|
||||
self.prompt_queue = _FakePromptQueue()
|
||||
def _load_async_helpers_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "async_helpers.py"
|
||||
package_name = "dist_async_helpers_testpkg"
|
||||
|
||||
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]
|
||||
@@ -36,68 +29,62 @@ def _bootstrap_package(package_name):
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.get_server_loop = asyncio.get_event_loop
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
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"], {})
|
||||
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 = ()
|
||||
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,
|
||||
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
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
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, fake_prompt_server
|
||||
return module, prompt_server
|
||||
|
||||
|
||||
class AsyncHelpersQueuePromptPayloadTests(unittest.TestCase):
|
||||
def test_queue_prompt_payload_includes_create_time_metadata(self):
|
||||
async_helpers_module, fake_prompt_server = _load_async_helpers_module()
|
||||
async_helpers, prompt_server = _load_async_helpers_module()
|
||||
|
||||
prompt_id = asyncio.run(
|
||||
async_helpers_module.queue_prompt_payload(
|
||||
{"1": {"class_type": "KSampler", "inputs": {}}},
|
||||
workflow_meta={"id": "workflow-1"},
|
||||
client_id="client-1",
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
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.assertIsInstance(result["prompt_id"], str)
|
||||
self.assertTrue(result["prompt_id"])
|
||||
self.assertEqual(result["number"], 12)
|
||||
self.assertEqual(result["node_errors"], {})
|
||||
|
||||
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__":
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
|
||||
|
||||
def _load_utilities_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
|
||||
package_name = "dist_divider_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -0,0 +1,205 @@
|
||||
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_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"}
|
||||
@@ -10,7 +10,7 @@ from unittest.mock import patch
|
||||
|
||||
|
||||
def _load_config_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "config.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "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_state().cache)
|
||||
self.assertIsNone(config._config_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[3] / "workers" / "detection.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "workers" / "detection.py"
|
||||
package_name = "dist_det_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -32,7 +32,6 @@ 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")
|
||||
@@ -157,7 +156,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}) # nosec B104 - explicit wildcard host test case
|
||||
result = await detection.is_local_worker({"host": "0.0.0.0", "port": 8188})
|
||||
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[3] / "api" / "orchestration" / "dispatch.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "api" / "orchestration" / "dispatch.py"
|
||||
|
||||
package_name = "dist_dispatch_testpkg"
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
@@ -249,50 +249,6 @@ 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()
|
||||
@@ -13,7 +13,7 @@ class DistributedValueTests(unittest.TestCase):
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
|
||||
pkg_name = "dv_test_pkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -28,18 +28,6 @@ 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
|
||||
@@ -49,16 +37,6 @@ 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
|
||||
)
|
||||
@@ -9,7 +9,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_job_timeout_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "upscale" / "job_timeout.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "upscale" / "job_timeout.py"
|
||||
package_name = "dist_job_timeout_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -6,7 +6,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_network_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "utils" / "network.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "network.py"
|
||||
|
||||
package_name = "dist_utils_testpkg"
|
||||
package_module = types.ModuleType(package_name)
|
||||
@@ -17,10 +17,6 @@ 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)
|
||||
@@ -94,14 +90,48 @@ 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": ""}}
|
||||
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190) # nosec B104 - wildcard bind normalization test
|
||||
cfg = {"master": {"host": "", "port": 8001}}
|
||||
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190)
|
||||
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()
|
||||
@@ -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[3] / "upscale" / "payload_parsers.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "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
|
||||
@@ -7,7 +7,7 @@ from pathlib import Path
|
||||
|
||||
|
||||
def _load_prompt_transform_module():
|
||||
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "prompt_transform.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "api" / "orchestration" / "prompt_transform.py"
|
||||
package_name = "dist_pt_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -81,38 +81,13 @@ def _delegate_prompt():
|
||||
}
|
||||
|
||||
|
||||
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)
|
||||
idx = pt.PromptIndex(prompt)
|
||||
job_id_map = pt.generate_job_id_map(idx, "run")
|
||||
return pt.apply_participant_overrides(
|
||||
prompt_copy,
|
||||
prompt,
|
||||
participant_id=participant_id,
|
||||
enabled_worker_ids=enabled_worker_ids,
|
||||
job_id_map=job_id_map,
|
||||
@@ -275,21 +250,6 @@ class PrunePromptForWorkerTests(unittest.TestCase):
|
||||
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
|
||||
@@ -367,22 +327,6 @@ class GenerateJobIdMapTests(unittest.TestCase):
|
||||
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"), {})
|
||||
@@ -474,9 +418,9 @@ class ApplyOverridesSeedTests(unittest.TestCase):
|
||||
result = _apply(self._seed_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_id_uses_canonical_worker_id(self):
|
||||
def test_worker_id_reflects_index_in_enabled_list(self):
|
||||
result = _apply(self._seed_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker_1")
|
||||
|
||||
def test_master_sets_is_worker_false(self):
|
||||
result = _apply(self._seed_prompt(), "master")
|
||||
@@ -486,11 +430,6 @@ class ApplyOverridesSeedTests(unittest.TestCase):
|
||||
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
|
||||
@@ -537,9 +476,9 @@ class ApplyOverridesValueTests(unittest.TestCase):
|
||||
result = _apply(self._value_prompt(), "worker-a")
|
||||
self.assertTrue(result["1"]["inputs"]["is_worker"])
|
||||
|
||||
def test_worker_id_uses_canonical_worker_id(self):
|
||||
def test_worker_id_reflects_index_in_enabled_list(self):
|
||||
result = _apply(self._value_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
|
||||
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker_1")
|
||||
|
||||
def test_master_sets_is_worker_false(self):
|
||||
result = _apply(self._value_prompt(), "master")
|
||||
@@ -549,157 +488,6 @@ class ApplyOverridesValueTests(unittest.TestCase):
|
||||
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[3] / "api" / "queue_request.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "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,26 +26,27 @@ class QueueRequestPayloadTests(unittest.TestCase):
|
||||
|
||||
def test_normalizes_enabled_worker_ids(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data["enabled_worker_ids"] = [" worker-a ", "worker-b", "worker-a"]
|
||||
payload_data["enabled_worker_ids"] = ["a", 2, 3]
|
||||
payload_data["delegate_master"] = True
|
||||
payload = parse_queue_request_payload(
|
||||
payload_data
|
||||
)
|
||||
self.assertEqual(payload.enabled_worker_ids, ["worker-a", "worker-b"])
|
||||
self.assertEqual(payload.enabled_worker_ids, ["a", "2", "3"])
|
||||
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": "w3"}, {"name": "no-id"}]
|
||||
payload_data["workers"] = [{"id": "w1"}, "w2", {"id": 3}, {"name": "no-id"}]
|
||||
payload = parse_queue_request_payload(
|
||||
payload_data
|
||||
)
|
||||
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "w3"])
|
||||
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "3"])
|
||||
|
||||
def test_supports_workflow_prompt_fallback(self):
|
||||
def test_supports_auto_prepare_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"},
|
||||
@@ -55,6 +56,7 @@ 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()
|
||||
@@ -72,6 +74,10 @@ 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)
|
||||
@@ -85,7 +91,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_is_ignored_for_backward_compat(self):
|
||||
def test_auto_prepare_false_still_falls_back_to_workflow_prompt(self):
|
||||
payload_data = self._base_payload()
|
||||
payload_data.pop("prompt", None)
|
||||
payload_data["auto_prepare"] = False
|
||||
@@ -94,6 +100,13 @@ 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()
|
||||
@@ -107,12 +120,6 @@ 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"])
|
||||
@@ -4,13 +4,12 @@ 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[3] / "upscale" / "modes" / "static.py"
|
||||
module_path = Path(__file__).resolve().parents[1] / "upscale" / "modes" / "static.py"
|
||||
package_name = "dist_static_mode_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
@@ -66,34 +65,8 @@ 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")
|
||||
@@ -129,10 +102,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")
|
||||
@@ -143,36 +116,6 @@ 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
|
||||
@@ -205,20 +148,14 @@ class _FakeStaticWorker(static_mode.StaticModeMixin):
|
||||
def _poll_job_ready(self, *_args, **_kwargs):
|
||||
return self.job_ready
|
||||
|
||||
async def request_assignment(self, *_args, **_kwargs):
|
||||
async def _request_tile_from_master(self, *_args, **_kwargs):
|
||||
self.request_calls += 1
|
||||
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,
|
||||
)
|
||||
return self.tile_sequence.pop(0)
|
||||
|
||||
async def send_heartbeat(self, *_args, **_kwargs):
|
||||
async def _send_heartbeat_to_master(self, *_args, **_kwargs):
|
||||
self.heartbeat_calls += 1
|
||||
|
||||
async def send_tiles_batch(
|
||||
async def send_tiles_batch_to_master(
|
||||
self,
|
||||
processed_tiles,
|
||||
_multi_job_id,
|
||||
@@ -0,0 +1,134 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from argparse import Namespace
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
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__ = []
|
||||
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"
|
||||
sys.modules[f"{package_name}.utils.process"] = process_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 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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1 +0,0 @@
|
||||
"""Unit tests."""
|
||||
@@ -1 +0,0 @@
|
||||
"""Unit tests."""
|
||||
@@ -1,80 +0,0 @@
|
||||
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()
|
||||
@@ -1,282 +0,0 @@
|
||||
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()
|
||||
@@ -1,228 +0,0 @@
|
||||
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()
|
||||
@@ -1,248 +0,0 @@
|
||||
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()
|
||||
@@ -1 +0,0 @@
|
||||
"""Unit tests."""
|
||||
@@ -1,71 +0,0 @@
|
||||
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()
|
||||
@@ -1,40 +0,0 @@
|
||||
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()
|
||||
@@ -1,228 +0,0 @@
|
||||
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()
|
||||
@@ -1,32 +0,0 @@
|
||||
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()
|
||||
@@ -1,154 +0,0 @@
|
||||
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()
|
||||
@@ -1 +0,0 @@
|
||||
"""Unit tests."""
|
||||
@@ -1,53 +0,0 @@
|
||||
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()
|
||||
@@ -1,247 +0,0 @@
|
||||
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()
|
||||
@@ -1,39 +0,0 @@
|
||||
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()
|
||||
@@ -1,131 +0,0 @@
|
||||
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()
|
||||
@@ -1,164 +0,0 @@
|
||||
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()
|
||||
@@ -1,248 +0,0 @@
|
||||
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()
|
||||
@@ -1,207 +0,0 @@
|
||||
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()
|
||||
@@ -1,179 +0,0 @@
|
||||
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()
|
||||
@@ -1,115 +0,0 @@
|
||||
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()
|
||||
@@ -1,128 +0,0 @@
|
||||
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,8 +1,7 @@
|
||||
import copy
|
||||
from typing import Any
|
||||
|
||||
|
||||
def clone_control_chain(control: Any, clone_hint: bool = True) -> Any:
|
||||
def clone_control_chain(control, clone_hint=True):
|
||||
"""Shallow copy the ControlNet chain, optionally cloning hints but sharing models."""
|
||||
if control is None:
|
||||
return None
|
||||
@@ -15,10 +14,7 @@ def clone_control_chain(control: Any, clone_hint: bool = True) -> Any:
|
||||
return new_control
|
||||
|
||||
|
||||
def clone_conditioning(
|
||||
cond_list: list[tuple[Any, dict[str, Any]]] | list[list[Any]],
|
||||
clone_hints: bool = True,
|
||||
) -> list[list[Any]]:
|
||||
def clone_conditioning(cond_list, clone_hints=True):
|
||||
"""Clone conditioning without duplicating ControlNet models."""
|
||||
new_cond = []
|
||||
for emb, cond_dict in cond_list:
|
||||
|
||||
+2
-34
@@ -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
|
||||
from .job_timeout import _check_and_requeue_timed_out_workers as _requeue_usdu
|
||||
from .job_models import ImageJobState, TileJobState
|
||||
|
||||
|
||||
@@ -22,10 +22,6 @@ 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()
|
||||
@@ -43,10 +39,6 @@ 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()
|
||||
@@ -64,10 +56,6 @@ 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()
|
||||
@@ -79,10 +67,6 @@ 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()
|
||||
@@ -92,10 +76,6 @@ 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()
|
||||
@@ -107,10 +87,6 @@ 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()
|
||||
@@ -154,14 +130,6 @@ 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 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)
|
||||
return await _requeue_usdu(multi_job_id, batch_size)
|
||||
|
||||
+60
-79
@@ -1,51 +1,45 @@
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import List, Optional
|
||||
|
||||
import server
|
||||
|
||||
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)))
|
||||
|
||||
|
||||
@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:
|
||||
def ensure_tile_jobs_initialized():
|
||||
"""Ensure tile job storage is initialized on the server instance."""
|
||||
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]
|
||||
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]
|
||||
return prompt_server
|
||||
|
||||
|
||||
async def _init_job_queue(
|
||||
multi_job_id: str,
|
||||
config: JobQueueInitConfig,
|
||||
) -> None:
|
||||
multi_job_id,
|
||||
mode,
|
||||
batch_size=None,
|
||||
num_tiles_per_image=None,
|
||||
all_indices=None,
|
||||
enabled_workers=None,
|
||||
batched_static: bool = False,
|
||||
):
|
||||
"""Unified initialization for job queues in static and dynamic modes."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -53,34 +47,33 @@ async def _init_job_queue(
|
||||
debug_log(f"Queue already exists for {multi_job_id}")
|
||||
return
|
||||
|
||||
if config.mode == 'dynamic':
|
||||
if mode == 'dynamic':
|
||||
job_data = ImageJobState(multi_job_id=multi_job_id)
|
||||
elif config.mode == 'static':
|
||||
elif mode == 'static':
|
||||
job_data = TileJobState(multi_job_id=multi_job_id)
|
||||
else:
|
||||
raise ValueError(f"Unknown mode: {config.mode}")
|
||||
raise ValueError(f"Unknown mode: {mode}")
|
||||
|
||||
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 []}
|
||||
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 []}
|
||||
|
||||
if config.mode == 'dynamic':
|
||||
job_data.batch_size = int(config.batch_size or 0)
|
||||
if mode == 'dynamic':
|
||||
job_data.batch_size = int(batch_size or 0)
|
||||
pending_queue = job_data.pending_images
|
||||
indices = config.all_indices or list(range(int(config.batch_size or 0)))
|
||||
for i in indices:
|
||||
for i in (all_indices or range(int(batch_size or 0))):
|
||||
await pending_queue.put(i)
|
||||
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)
|
||||
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)
|
||||
# For batched static distribution, populate only tile ids [0..num_tiles_per_image-1]
|
||||
pending_queue = job_data.pending_tasks
|
||||
if config.batched_static and config.num_tiles_per_image > 0:
|
||||
for i in range(config.num_tiles_per_image):
|
||||
if batched_static and num_tiles_per_image is not None:
|
||||
for i in range(num_tiles_per_image):
|
||||
await pending_queue.put(i)
|
||||
else:
|
||||
total_tiles = int(config.batch_size or 0) * int(config.num_tiles_per_image or 0)
|
||||
total_tiles = int(batch_size or 0) * int(num_tiles_per_image or 0)
|
||||
for i in range(total_tiles):
|
||||
await pending_queue.put(i)
|
||||
|
||||
@@ -90,18 +83,16 @@ async def _init_job_queue(
|
||||
async def init_dynamic_job(
|
||||
multi_job_id: str,
|
||||
batch_size: int,
|
||||
enabled_workers: list[str],
|
||||
all_indices: list[int] | None = None,
|
||||
) -> None:
|
||||
enabled_workers: List[str],
|
||||
all_indices: Optional[List[int]] = None,
|
||||
):
|
||||
"""Initialize queue for dynamic mode (per-image), with collector fields."""
|
||||
await _init_job_queue(
|
||||
multi_job_id,
|
||||
JobQueueInitConfig(
|
||||
mode='dynamic',
|
||||
batch_size=batch_size,
|
||||
all_indices=all_indices or list(range(batch_size)),
|
||||
enabled_workers=enabled_workers,
|
||||
),
|
||||
'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")
|
||||
|
||||
@@ -110,22 +101,20 @@ async def init_static_job_batched(
|
||||
multi_job_id: str,
|
||||
batch_size: int,
|
||||
num_tiles_per_image: int,
|
||||
enabled_workers: list[str],
|
||||
) -> None:
|
||||
enabled_workers: List[str],
|
||||
):
|
||||
"""Initialize queue for static mode (batched-per-tile)."""
|
||||
await _init_job_queue(
|
||||
multi_job_id,
|
||||
JobQueueInitConfig(
|
||||
mode='static',
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
enabled_workers=enabled_workers,
|
||||
batched_static=True,
|
||||
),
|
||||
'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: str) -> int:
|
||||
async def _drain_results_queue(multi_job_id):
|
||||
"""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:
|
||||
@@ -163,7 +152,7 @@ async def drain_results_queue(multi_job_id: str) -> int:
|
||||
return collected
|
||||
|
||||
|
||||
async def get_completed_count(multi_job_id: str) -> int:
|
||||
async def _get_completed_count(multi_job_id):
|
||||
"""Get count of completed tasks."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -173,7 +162,7 @@ async def get_completed_count(multi_job_id: str) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
async def mark_task_completed(multi_job_id: str, task_id: int, result: Any) -> None:
|
||||
async def _mark_task_completed(multi_job_id, task_id, result):
|
||||
"""Mark a task as completed."""
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
@@ -182,18 +171,10 @@ async def mark_task_completed(multi_job_id: str, task_id: int, result: Any) -> N
|
||||
job_data.completed_tasks[task_id] = result
|
||||
|
||||
|
||||
async def cleanup_job(multi_job_id: str) -> None:
|
||||
async def _cleanup_job(multi_job_id):
|
||||
"""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,10 +99,6 @@ 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")
|
||||
@@ -152,8 +148,3 @@ 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)
|
||||
|
||||
@@ -1,238 +0,0 @@
|
||||
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
|
||||
+151
-310
@@ -1,27 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio, time, torch
|
||||
from typing import Any
|
||||
import asyncio, torch
|
||||
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 (
|
||||
JOB_POLL_INTERVAL,
|
||||
JOB_POLL_MAX_ATTEMPTS,
|
||||
TILE_WAIT_TIMEOUT,
|
||||
TILE_SEND_TIMEOUT,
|
||||
)
|
||||
from ...utils.constants import 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:
|
||||
@@ -29,205 +14,16 @@ 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_assignment`, `send_full_image`, `send_heartbeat`).
|
||||
- WorkerCommsMixin (`_request_image_from_master`, `_send_full_image_to_master`, `_send_heartbeat_to_master`).
|
||||
"""
|
||||
|
||||
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]:
|
||||
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):
|
||||
"""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)
|
||||
@@ -240,7 +36,7 @@ class DynamicModeMixin:
|
||||
debug_log(f"Processing {batch_size} images dynamically across master + {num_workers} workers.")
|
||||
|
||||
# Calculate tiles for processing
|
||||
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
|
||||
# Initialize job queue for communication
|
||||
try:
|
||||
@@ -265,7 +61,7 @@ class DynamicModeMixin:
|
||||
while processed_count < batch_size:
|
||||
# Try to get an image to process
|
||||
image_idx = run_async_in_server_loop(
|
||||
context.job_state.next_image_index(multi_job_id),
|
||||
self._get_next_image_index(multi_job_id),
|
||||
timeout=5.0 # Short timeout to allow frequent checks
|
||||
)
|
||||
|
||||
@@ -278,25 +74,38 @@ class DynamicModeMixin:
|
||||
# Process locally
|
||||
single_tensor = upscaled_image[image_idx:image_idx+1]
|
||||
local_image = result_images[image_idx]
|
||||
|
||||
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),
|
||||
)
|
||||
|
||||
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"
|
||||
|
||||
result_images[image_idx] = local_image
|
||||
|
||||
# Mark as completed
|
||||
run_async_in_server_loop(
|
||||
context.result_collector.mark_image_completed(multi_job_id, image_idx, local_image),
|
||||
self._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(
|
||||
context.job_state.drain_worker_results_queue(multi_job_id),
|
||||
self._drain_worker_results_queue(multi_job_id),
|
||||
timeout=5.0
|
||||
)
|
||||
|
||||
@@ -305,106 +114,133 @@ class DynamicModeMixin:
|
||||
|
||||
# NEW: Log overall progress (includes master's image + any drained workers)
|
||||
completed_now = run_async_in_server_loop(
|
||||
context.job_state.total_completed_count(multi_job_id),
|
||||
self._get_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(context.worker_comms.async_yield(), timeout=0.1)
|
||||
run_async_in_server_loop(self._async_yield(), timeout=0.1)
|
||||
else:
|
||||
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,
|
||||
# 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
|
||||
)
|
||||
if should_break:
|
||||
break
|
||||
if should_continue:
|
||||
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:
|
||||
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
|
||||
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")
|
||||
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,
|
||||
|
||||
# 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
|
||||
)
|
||||
|
||||
# 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: 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]:
|
||||
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):
|
||||
"""Worker processing in dynamic mode - processes whole images."""
|
||||
context = mode_context or self._build_dynamic_mode_context()
|
||||
# Round tile dimensions
|
||||
tile_width = context.tile_ops.round_to_multiple(tile_width)
|
||||
tile_height = context.tile_ops.round_to_multiple(tile_height)
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
tile_height = self.round_to_multiple(tile_height)
|
||||
|
||||
# Get dimensions and tile grid
|
||||
batch_size, height, width, _ = upscaled_image.shape
|
||||
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = self.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 = 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,
|
||||
):
|
||||
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):
|
||||
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
|
||||
assignment = run_async_in_server_loop(
|
||||
context.worker_comms.request_assignment(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_WAIT_TIMEOUT,
|
||||
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
|
||||
)
|
||||
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")
|
||||
@@ -423,34 +259,39 @@ class DynamicModeMixin:
|
||||
local_image = tensor_to_pil(single_tensor, 0).copy()
|
||||
|
||||
# Process all tiles for this 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.send_heartbeat(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0,
|
||||
),
|
||||
)
|
||||
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
|
||||
)
|
||||
|
||||
# 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(
|
||||
context.worker_comms.send_full_image(
|
||||
local_image,
|
||||
image_idx,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id,
|
||||
is_last,
|
||||
),
|
||||
self._send_full_image_to_master(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(
|
||||
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
|
||||
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0
|
||||
)
|
||||
if is_last:
|
||||
@@ -463,7 +304,7 @@ class DynamicModeMixin:
|
||||
debug_log(f"Worker[{worker_id[:8]}] processed {processed_count} images, sending completion signal")
|
||||
try:
|
||||
run_async_in_server_loop(
|
||||
context.worker_comms.send_worker_complete_signal(multi_job_id, master_url, worker_id),
|
||||
self._send_worker_complete_signal(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_SEND_TIMEOUT
|
||||
)
|
||||
except Exception as e:
|
||||
|
||||
+24
-56
@@ -1,57 +1,29 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math, torch
|
||||
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
|
||||
from PIL import Image
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.image import tensor_to_pil, pil_to_tensor
|
||||
|
||||
|
||||
class SingleGpuModeMixin:
|
||||
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]:
|
||||
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):
|
||||
"""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 = ops.round_to_multiple(tile_width)
|
||||
tile_height = ops.round_to_multiple(tile_height)
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
tile_height = self.round_to_multiple(tile_height)
|
||||
|
||||
# Get image dimensions and batch size
|
||||
batch_size, height, width, _ = upscaled_image.shape
|
||||
|
||||
# Calculate all tiles
|
||||
all_tiles = ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = self.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 = []
|
||||
@@ -62,7 +34,7 @@ class SingleGpuModeMixin:
|
||||
# Precompute tile masks once
|
||||
tile_masks = []
|
||||
for tx, ty in all_tiles:
|
||||
tile_masks.append(ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
tile_masks.append(self.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):
|
||||
@@ -71,29 +43,25 @@ class SingleGpuModeMixin:
|
||||
if upscaled_image.is_cuda:
|
||||
source_batch = source_batch.cuda()
|
||||
|
||||
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,
|
||||
# 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
|
||||
)
|
||||
|
||||
# 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):
|
||||
blend_processed_batch_item(
|
||||
result_images,
|
||||
processed_batch,
|
||||
b,
|
||||
ops.blend_tile,
|
||||
x1,
|
||||
y1,
|
||||
ew,
|
||||
eh,
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
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)
|
||||
|
||||
# Convert back to tensor
|
||||
result_tensors = [pil_to_tensor(img) for img in result_images]
|
||||
|
||||
+239
-354
@@ -2,7 +2,7 @@ import asyncio, time, torch
|
||||
from PIL import Image
|
||||
import comfy.model_management
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.image import blend_processed_batch_item, pil_to_tensor, tensor_to_pil
|
||||
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 (
|
||||
@@ -15,16 +15,9 @@ from ...utils.constants import (
|
||||
)
|
||||
from ..job_store import (
|
||||
ensure_tile_jobs_initialized, init_static_job_batched,
|
||||
cleanup_job, drain_results_queue, get_completed_count, mark_task_completed,
|
||||
_mark_task_completed, _cleanup_job, _drain_results_queue, _get_completed_count,
|
||||
)
|
||||
from ..job_models import TileJobState
|
||||
from ..mode_contexts import (
|
||||
JobStateCollaborator,
|
||||
StaticModeContext,
|
||||
TileOpsCollaborator,
|
||||
WorkerCommsCollaborator,
|
||||
)
|
||||
from ..tile_processing import TileBatchArgs, extract_and_process_tile_batch
|
||||
|
||||
|
||||
class StaticModeMixin:
|
||||
@@ -34,30 +27,14 @@ class StaticModeMixin:
|
||||
Expected co-mixins on `self`:
|
||||
- TileOpsMixin (`calculate_tiles`, tile extract/blend helpers).
|
||||
- JobStateMixin (`_get_next_tile_index`, `_get_all_completed_tasks`, requeue checks).
|
||||
- WorkerCommsMixin (`send_tiles_batch`, `request_assignment`, `send_heartbeat`).
|
||||
- WorkerCommsMixin (`send_tiles_batch_to_master`, `_request_tile_from_master`, `_send_heartbeat_to_master`).
|
||||
"""
|
||||
|
||||
def _build_static_mode_context(self) -> StaticModeContext:
|
||||
"""Build explicit collaborators for static mode execution."""
|
||||
return StaticModeContext(
|
||||
tile_ops=TileOpsCollaborator(self),
|
||||
job_state=JobStateCollaborator(self),
|
||||
worker_comms=WorkerCommsCollaborator(self),
|
||||
)
|
||||
|
||||
def _poll_job_ready(
|
||||
self,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id=None,
|
||||
max_attempts=JOB_POLL_MAX_ATTEMPTS,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
def _poll_job_ready(self, multi_job_id, master_url, worker_id=None, max_attempts=JOB_POLL_MAX_ATTEMPTS):
|
||||
"""Poll master for job readiness to avoid worker/master initialization race."""
|
||||
context = mode_context or self._build_static_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),
|
||||
self._check_job_status(multi_job_id, master_url),
|
||||
timeout=5.0
|
||||
)
|
||||
if ready:
|
||||
@@ -78,27 +55,30 @@ class StaticModeMixin:
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
):
|
||||
"""Extract one tile position for the whole batch and process it."""
|
||||
tx, ty = all_tiles[tile_id]
|
||||
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,
|
||||
tile_batch, x1, y1, ew, eh = self.extract_batch_tile_with_padding(
|
||||
upscaled_image, tx, ty, tile_width, tile_height, padding, force_uniform_tiles
|
||||
)
|
||||
processed_batch, x1, y1, ew, eh = extract_and_process_tile_batch(
|
||||
node=self,
|
||||
upscaled_image=upscaled_image,
|
||||
tx=tx,
|
||||
ty=ty,
|
||||
args=tile_batch_args,
|
||||
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)
|
||||
)
|
||||
return processed_batch, x1, y1, ew, eh
|
||||
|
||||
@@ -110,14 +90,12 @@ class StaticModeMixin:
|
||||
padding,
|
||||
worker_id,
|
||||
is_final_flush=False,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
"""Send accumulated tile payloads to master and return a fresh accumulator."""
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
if not processed_tiles:
|
||||
if is_final_flush:
|
||||
run_async_in_server_loop(
|
||||
context.worker_comms.send_tiles_batch(
|
||||
self.send_tiles_batch_to_master(
|
||||
[],
|
||||
multi_job_id,
|
||||
master_url,
|
||||
@@ -129,7 +107,7 @@ class StaticModeMixin:
|
||||
)
|
||||
return processed_tiles
|
||||
run_async_in_server_loop(
|
||||
context.worker_comms.send_tiles_batch(
|
||||
self.send_tiles_batch_to_master(
|
||||
processed_tiles,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
@@ -155,13 +133,21 @@ class StaticModeMixin:
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
"""Process one tile_id across the batch and blend into result_images."""
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
source_batch = torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
|
||||
if upscaled_image.is_cuda:
|
||||
source_batch = source_batch.cuda()
|
||||
@@ -173,7 +159,17 @@ class StaticModeMixin:
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
@@ -181,39 +177,30 @@ class StaticModeMixin:
|
||||
out_bs = processed_batch.shape[0] if hasattr(processed_batch, "shape") else batch_size
|
||||
processed_items = min(batch_size, out_bs)
|
||||
for b in range(processed_items):
|
||||
blend_processed_batch_item(
|
||||
result_images,
|
||||
processed_batch,
|
||||
b,
|
||||
context.tile_ops.blend_tile,
|
||||
x1,
|
||||
y1,
|
||||
ew,
|
||||
eh,
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
tile_pil = tensor_to_pil(processed_batch, b)
|
||||
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)
|
||||
global_idx = b * num_tiles_per_image + tile_id
|
||||
run_async_in_server_loop(
|
||||
mark_task_completed(multi_job_id, global_idx, {'batch_idx': b, 'tile_idx': tile_id}),
|
||||
_mark_task_completed(multi_job_id, global_idx, {'batch_idx': b, 'tile_idx': tile_id}),
|
||||
timeout=5.0
|
||||
)
|
||||
return processed_items
|
||||
|
||||
def _process_worker_static_sync(self, upscaled_image, core_args,
|
||||
def _process_worker_static_sync(self, upscaled_image, model, positive, negative, vae,
|
||||
seed, steps, cfg, sampler_name, scheduler, denoise,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, multi_job_id, master_url,
|
||||
worker_id, enabled_workers,
|
||||
mode_context: StaticModeContext | None = None):
|
||||
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
|
||||
worker_id, enabled_workers):
|
||||
"""Worker static mode processing with optional dynamic queue pulling."""
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
# Round tile dimensions
|
||||
tile_width = context.tile_ops.round_to_multiple(tile_width)
|
||||
tile_height = context.tile_ops.round_to_multiple(tile_height)
|
||||
tile_width = self.round_to_multiple(tile_width)
|
||||
tile_height = self.round_to_multiple(tile_height)
|
||||
|
||||
# Get dimensions and calculate tiles
|
||||
_, height, width, _ = upscaled_image.shape
|
||||
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
|
||||
num_tiles_per_image = len(all_tiles)
|
||||
batch_size = upscaled_image.shape[0]
|
||||
total_tiles = batch_size * num_tiles_per_image
|
||||
@@ -225,31 +212,24 @@ class StaticModeMixin:
|
||||
working_images.append(image_pil.copy())
|
||||
tile_masks = []
|
||||
for tx, ty in all_tiles:
|
||||
tile_masks.append(context.tile_ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
|
||||
# Dynamic queue mode (static processing): process batched-per-tile
|
||||
log(f"USDU Dist Worker[{worker_id[:8]}]: Canvas {width}x{height} | Tile {tile_width}x{tile_height} | Tiles/image {num_tiles_per_image} | Batch {batch_size}")
|
||||
processed_count = 0
|
||||
|
||||
max_poll_attempts = JOB_POLL_MAX_ATTEMPTS
|
||||
if not self._poll_job_ready(
|
||||
multi_job_id,
|
||||
master_url,
|
||||
worker_id=worker_id,
|
||||
max_attempts=max_poll_attempts,
|
||||
mode_context=context,
|
||||
):
|
||||
if not self._poll_job_ready(multi_job_id, master_url, worker_id=worker_id, max_attempts=max_poll_attempts):
|
||||
log(f"Job {multi_job_id} not ready after {max_poll_attempts} attempts, aborting")
|
||||
return (upscaled_image,)
|
||||
|
||||
# Main processing loop - pull tile ids from queue
|
||||
while True:
|
||||
# Request a tile to process
|
||||
assignment = run_async_in_server_loop(
|
||||
context.worker_comms.request_assignment(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_WAIT_TIMEOUT,
|
||||
tile_idx, estimated_remaining, batched_static = run_async_in_server_loop(
|
||||
self._request_tile_from_master(multi_job_id, master_url, worker_id),
|
||||
timeout=TILE_WAIT_TIMEOUT
|
||||
)
|
||||
tile_idx = assignment.task_idx if assignment.kind == "tile" else None
|
||||
|
||||
if tile_idx is None:
|
||||
debug_log(f"Worker[{worker_id[:8]}] - No more tiles to process")
|
||||
@@ -270,7 +250,17 @@ class StaticModeMixin:
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
@@ -279,7 +269,7 @@ class StaticModeMixin:
|
||||
tile_pil = tensor_to_pil(processed_batch, b)
|
||||
if tile_pil.size != (ew, eh):
|
||||
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
|
||||
working_images[b] = context.tile_ops.blend_tile(
|
||||
working_images[b] = self.blend_tile(
|
||||
working_images[b],
|
||||
tile_pil,
|
||||
x1,
|
||||
@@ -303,7 +293,7 @@ class StaticModeMixin:
|
||||
# Send heartbeat
|
||||
try:
|
||||
run_async_in_server_loop(
|
||||
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
|
||||
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
|
||||
timeout=5.0
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -312,39 +302,20 @@ class StaticModeMixin:
|
||||
# Send tiles in batches within loop
|
||||
if len(processed_tiles) >= MAX_BATCH:
|
||||
processed_tiles = self._flush_tiles_to_master(
|
||||
processed_tiles,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
padding,
|
||||
worker_id,
|
||||
is_final_flush=False,
|
||||
mode_context=context,
|
||||
processed_tiles, multi_job_id, master_url, padding, worker_id, is_final_flush=False
|
||||
)
|
||||
|
||||
# Send any remaining tiles
|
||||
processed_tiles = self._flush_tiles_to_master(
|
||||
processed_tiles,
|
||||
multi_job_id,
|
||||
master_url,
|
||||
padding,
|
||||
worker_id,
|
||||
is_final_flush=True,
|
||||
mode_context=context,
|
||||
processed_tiles, multi_job_id, master_url, padding, worker_id, is_final_flush=True
|
||||
)
|
||||
|
||||
debug_log(f"Worker {worker_id} completed all assigned and requeued tiles")
|
||||
return (upscaled_image,)
|
||||
|
||||
async def _async_collect_and_monitor_static(
|
||||
self,
|
||||
multi_job_id,
|
||||
total_tiles,
|
||||
expected_total,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
async def _async_collect_and_monitor_static(self, multi_job_id, total_tiles, expected_total):
|
||||
"""Async helper for collection and monitoring in static mode.
|
||||
Returns collected tasks dict. Caller should check if all tasks are complete."""
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
last_progress_log = time.time()
|
||||
progress_interval = 5.0
|
||||
last_heartbeat_check = time.time()
|
||||
@@ -357,18 +328,18 @@ class StaticModeMixin:
|
||||
raise comfy.model_management.InterruptProcessingException()
|
||||
|
||||
# Drain any pending results
|
||||
collected_count = await drain_results_queue(multi_job_id)
|
||||
collected_count = await _drain_results_queue(multi_job_id)
|
||||
|
||||
# Check and requeue timed-out workers periodically
|
||||
current_time = time.time()
|
||||
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
|
||||
requeued_count = await context.job_state.check_and_requeue_timed_out_workers(multi_job_id, expected_total)
|
||||
requeued_count = await self._check_and_requeue_timed_out_workers(multi_job_id, expected_total)
|
||||
if requeued_count > 0:
|
||||
log(f"Requeued {requeued_count} tasks from timed-out workers")
|
||||
last_heartbeat_check = current_time
|
||||
|
||||
# Get current completion count
|
||||
completed_count = await get_completed_count(multi_job_id)
|
||||
completed_count = await _get_completed_count(multi_job_id)
|
||||
|
||||
# Progress logging
|
||||
if current_time - last_progress_log >= progress_interval:
|
||||
@@ -395,175 +366,159 @@ class StaticModeMixin:
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Get all completed tasks for return
|
||||
return await context.job_state.all_completed_tasks(multi_job_id)
|
||||
return await self._get_all_completed_tasks(multi_job_id)
|
||||
|
||||
def _init_master_static_state(
|
||||
self,
|
||||
upscaled_image,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
enabled_workers,
|
||||
all_tiles,
|
||||
tile_width,
|
||||
tile_height,
|
||||
mask_blur,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
def _process_master_static_sync(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,
|
||||
all_tiles, num_tiles_per_image):
|
||||
"""Static mode master processing with optional dynamic queue pulling."""
|
||||
batch_size = upscaled_image.shape[0]
|
||||
_, height, width, _ = upscaled_image.shape
|
||||
result_images = [tensor_to_pil(upscaled_image[b:b + 1], 0).copy() for b in range(batch_size)]
|
||||
|
||||
total_tiles = batch_size * num_tiles_per_image
|
||||
|
||||
# Convert batch to PIL list for processing
|
||||
result_images = []
|
||||
for b in range(batch_size):
|
||||
image_pil = tensor_to_pil(upscaled_image[b:b+1], 0)
|
||||
result_images.append(image_pil.copy())
|
||||
|
||||
# Initialize queue: pending queue holds tile ids (batched per tile)
|
||||
log("USDU Dist: Using tile queue distribution")
|
||||
run_async_in_server_loop(
|
||||
init_static_job_batched(multi_job_id, batch_size, num_tiles_per_image, enabled_workers),
|
||||
timeout=10.0,
|
||||
timeout=10.0
|
||||
)
|
||||
debug_log(
|
||||
f"Initialized tile-id queue with {num_tiles_per_image} ids for batch {batch_size}"
|
||||
)
|
||||
debug_log(f"Initialized tile-id queue with {num_tiles_per_image} ids for batch {batch_size}")
|
||||
|
||||
tile_masks = [
|
||||
context.tile_ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur)
|
||||
for tx, ty in all_tiles
|
||||
]
|
||||
return result_images, tile_masks, width, height
|
||||
# Precompute masks for all tile positions to avoid repeated Gaussian blur work during blending
|
||||
tile_masks = []
|
||||
for idx, (tx, ty) in enumerate(all_tiles):
|
||||
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
|
||||
|
||||
def _process_master_static_initial_tiles(
|
||||
self,
|
||||
multi_job_id,
|
||||
total_tiles,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
width,
|
||||
height,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
processed_count = 0
|
||||
consecutive_no_tile = 0
|
||||
max_consecutive_no_tile = 2
|
||||
|
||||
while processed_count < total_tiles:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
tile_id = run_async_in_server_loop(context.job_state.next_tile_index(multi_job_id), timeout=5.0)
|
||||
if tile_id is None:
|
||||
tile_idx = run_async_in_server_loop(
|
||||
self._get_next_tile_index(multi_job_id),
|
||||
timeout=5.0
|
||||
)
|
||||
if tile_idx is not None:
|
||||
consecutive_no_tile = 0
|
||||
tile_id = tile_idx
|
||||
processed_count += self._master_process_one_tile(
|
||||
tile_id,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
log(f"USDU Dist: Tiles progress {processed_count}/{total_tiles} (tile {tile_id})")
|
||||
else:
|
||||
consecutive_no_tile += 1
|
||||
if consecutive_no_tile >= max_consecutive_no_tile:
|
||||
debug_log(f"Master processed {processed_count} tiles, moving to collection phase")
|
||||
break
|
||||
time.sleep(0.1)
|
||||
continue
|
||||
|
||||
consecutive_no_tile = 0
|
||||
processed_count += self._master_process_one_tile(
|
||||
tile_id,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
width,
|
||||
height,
|
||||
mode_context=context,
|
||||
)
|
||||
log(f"USDU Dist: Tiles progress {processed_count}/{total_tiles} (tile {tile_id})")
|
||||
return processed_count
|
||||
|
||||
def _collect_remaining_static_tiles(
|
||||
self,
|
||||
multi_job_id,
|
||||
total_tiles,
|
||||
master_processed_count,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
width,
|
||||
height,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
master_processed_count = processed_count
|
||||
|
||||
# Continue processing any remaining tiles while collecting worker results
|
||||
remaining_tiles = total_tiles - master_processed_count
|
||||
if remaining_tiles <= 0:
|
||||
return run_async_in_server_loop(context.job_state.all_completed_tasks(multi_job_id), timeout=5.0)
|
||||
if remaining_tiles > 0:
|
||||
debug_log(f"Master waiting for {remaining_tiles} tiles from workers")
|
||||
|
||||
# Collect worker results using async operations
|
||||
try:
|
||||
# Wait until either all tasks are collected or there are no active workers left
|
||||
collected_tasks = run_async_in_server_loop(
|
||||
self._async_collect_and_monitor_static(multi_job_id, total_tiles, expected_total=total_tiles),
|
||||
timeout=None
|
||||
)
|
||||
except comfy.model_management.InterruptProcessingException:
|
||||
# Clean up job on interruption
|
||||
run_async_in_server_loop(_cleanup_job(multi_job_id), timeout=5.0)
|
||||
raise
|
||||
|
||||
# Check if we need to process any remaining tasks locally after collection
|
||||
completed_count = len(collected_tasks)
|
||||
if completed_count < total_tiles:
|
||||
log(f"Processing remaining {total_tiles - completed_count} tasks locally after worker failures")
|
||||
|
||||
# Process any remaining pending tasks (batched-per-tile)
|
||||
while True:
|
||||
# Check for user interruption
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
|
||||
debug_log(f"Master waiting for {remaining_tiles} tiles from workers")
|
||||
collected_tasks = run_async_in_server_loop(
|
||||
self._async_collect_and_monitor_static(
|
||||
multi_job_id,
|
||||
total_tiles,
|
||||
expected_total=total_tiles,
|
||||
mode_context=context,
|
||||
),
|
||||
timeout=None,
|
||||
)
|
||||
# Get next tile_id from pending queue
|
||||
tile_id = run_async_in_server_loop(
|
||||
self._get_next_tile_index(multi_job_id),
|
||||
timeout=5.0
|
||||
)
|
||||
|
||||
completed_count = len(collected_tasks)
|
||||
if completed_count >= total_tiles:
|
||||
return collected_tasks
|
||||
if tile_id is None:
|
||||
break
|
||||
|
||||
log(f"Processing remaining {total_tiles - completed_count} tasks locally after worker failures")
|
||||
while True:
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
tile_id = run_async_in_server_loop(context.job_state.next_tile_index(multi_job_id), timeout=5.0)
|
||||
if tile_id is None:
|
||||
break
|
||||
self._master_process_one_tile(
|
||||
tile_id,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
core_args,
|
||||
width,
|
||||
height,
|
||||
mode_context=context,
|
||||
self._master_process_one_tile(
|
||||
tile_id,
|
||||
all_tiles,
|
||||
upscaled_image,
|
||||
result_images,
|
||||
tile_masks,
|
||||
multi_job_id,
|
||||
batch_size,
|
||||
num_tiles_per_image,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
force_uniform_tiles,
|
||||
model,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
seed,
|
||||
steps,
|
||||
cfg,
|
||||
sampler_name,
|
||||
scheduler,
|
||||
denoise,
|
||||
tiled_decode,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
else:
|
||||
# Master processed all tiles
|
||||
collected_tasks = run_async_in_server_loop(
|
||||
self._get_all_completed_tasks(multi_job_id),
|
||||
timeout=5.0
|
||||
)
|
||||
return collected_tasks
|
||||
|
||||
def _blend_static_collected_tiles(
|
||||
self,
|
||||
collected_tasks,
|
||||
result_images,
|
||||
all_tiles,
|
||||
tile_masks,
|
||||
num_tiles_per_image,
|
||||
batch_size,
|
||||
tile_width,
|
||||
tile_height,
|
||||
padding,
|
||||
mode_context: StaticModeContext | None = None,
|
||||
):
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
|
||||
# Blend worker tiles synchronously in deterministic tile order.
|
||||
def _sort_key(item):
|
||||
global_idx, tile_data = item
|
||||
batch_idx = tile_data.get('batch_idx', global_idx // num_tiles_per_image)
|
||||
@@ -571,115 +526,45 @@ class StaticModeMixin:
|
||||
return (tile_idx, batch_idx, global_idx)
|
||||
|
||||
for global_idx, tile_data in sorted(collected_tasks.items(), key=_sort_key):
|
||||
# Skip tiles that don't have tensor data (already processed)
|
||||
if 'tensor' not in tile_data and 'image' not in tile_data:
|
||||
continue
|
||||
|
||||
batch_idx = tile_data.get('batch_idx', global_idx // num_tiles_per_image)
|
||||
tile_idx = tile_data.get('tile_idx', global_idx % num_tiles_per_image)
|
||||
|
||||
if batch_idx >= batch_size:
|
||||
continue
|
||||
|
||||
|
||||
# Blend tile synchronously
|
||||
x = tile_data.get('x', 0)
|
||||
y = tile_data.get('y', 0)
|
||||
tile_pil = tile_data['image'] if 'image' in tile_data else tensor_to_pil(tile_data['tensor'], 0)
|
||||
# Prefer PIL image if present to avoid reconversion
|
||||
if 'image' in tile_data:
|
||||
tile_pil = tile_data['image']
|
||||
else:
|
||||
tile_tensor = tile_data['tensor']
|
||||
tile_pil = tensor_to_pil(tile_tensor, 0)
|
||||
orig_x, orig_y = all_tiles[tile_idx]
|
||||
tile_mask = tile_masks[tile_idx]
|
||||
extracted_width = tile_data.get('extracted_width', tile_width + 2 * padding)
|
||||
extracted_height = tile_data.get('extracted_height', tile_height + 2 * padding)
|
||||
result_images[batch_idx] = context.tile_ops.blend_tile(
|
||||
result_images[batch_idx],
|
||||
tile_pil,
|
||||
x,
|
||||
y,
|
||||
(extracted_width, extracted_height),
|
||||
tile_mask,
|
||||
padding,
|
||||
)
|
||||
|
||||
def _result_images_to_tensor(self, result_images, batch_size, upscaled_image):
|
||||
if batch_size == 1:
|
||||
result_tensor = pil_to_tensor(result_images[0])
|
||||
else:
|
||||
result_tensor = torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
|
||||
if upscaled_image.is_cuda:
|
||||
result_tensor = result_tensor.cuda()
|
||||
return result_tensor
|
||||
|
||||
def _process_master_static_sync(self, upscaled_image, core_args,
|
||||
tile_width, tile_height, padding, mask_blur,
|
||||
force_uniform_tiles, multi_job_id, enabled_workers,
|
||||
all_tiles, num_tiles_per_image,
|
||||
mode_context: StaticModeContext | None = None):
|
||||
"""Static mode master processing with optional dynamic queue pulling."""
|
||||
context = mode_context or self._build_static_mode_context()
|
||||
batch_size = upscaled_image.shape[0]
|
||||
total_tiles = batch_size * num_tiles_per_image
|
||||
result_images, tile_masks, width, height = self._init_master_static_state(
|
||||
upscaled_image=upscaled_image,
|
||||
multi_job_id=multi_job_id,
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
enabled_workers=enabled_workers,
|
||||
all_tiles=all_tiles,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
mask_blur=mask_blur,
|
||||
mode_context=context,
|
||||
)
|
||||
|
||||
result_images[batch_idx] = self.blend_tile(result_images[batch_idx], tile_pil,
|
||||
x, y, (extracted_width, extracted_height), tile_mask, padding)
|
||||
|
||||
try:
|
||||
master_processed_count = self._process_master_static_initial_tiles(
|
||||
multi_job_id=multi_job_id,
|
||||
total_tiles=total_tiles,
|
||||
all_tiles=all_tiles,
|
||||
upscaled_image=upscaled_image,
|
||||
result_images=result_images,
|
||||
tile_masks=tile_masks,
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
core_args=core_args,
|
||||
width=width,
|
||||
height=height,
|
||||
mode_context=context,
|
||||
)
|
||||
|
||||
collected_tasks = self._collect_remaining_static_tiles(
|
||||
multi_job_id=multi_job_id,
|
||||
total_tiles=total_tiles,
|
||||
master_processed_count=master_processed_count,
|
||||
all_tiles=all_tiles,
|
||||
upscaled_image=upscaled_image,
|
||||
result_images=result_images,
|
||||
tile_masks=tile_masks,
|
||||
batch_size=batch_size,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
force_uniform_tiles=force_uniform_tiles,
|
||||
core_args=core_args,
|
||||
width=width,
|
||||
height=height,
|
||||
mode_context=context,
|
||||
)
|
||||
|
||||
self._blend_static_collected_tiles(
|
||||
collected_tasks=collected_tasks,
|
||||
result_images=result_images,
|
||||
all_tiles=all_tiles,
|
||||
tile_masks=tile_masks,
|
||||
num_tiles_per_image=num_tiles_per_image,
|
||||
batch_size=batch_size,
|
||||
tile_width=tile_width,
|
||||
tile_height=tile_height,
|
||||
padding=padding,
|
||||
mode_context=context,
|
||||
)
|
||||
|
||||
result_tensor = self._result_images_to_tensor(result_images, batch_size, upscaled_image)
|
||||
# Convert back to tensor
|
||||
if batch_size == 1:
|
||||
result_tensor = pil_to_tensor(result_images[0])
|
||||
else:
|
||||
result_tensors = [pil_to_tensor(img) for img in result_images]
|
||||
result_tensor = torch.cat(result_tensors, dim=0)
|
||||
|
||||
if upscaled_image.is_cuda:
|
||||
result_tensor = result_tensor.cuda()
|
||||
|
||||
log(f"UltimateSDUpscale Master - Job {multi_job_id} complete")
|
||||
return (result_tensor,)
|
||||
finally:
|
||||
run_async_in_server_loop(cleanup_job(multi_job_id), timeout=5.0)
|
||||
# Cleanup (async operation) - always execute
|
||||
run_async_in_server_loop(_cleanup_job(multi_job_id), timeout=5.0)
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
import io
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def parse_tiles_from_form(data: Mapping[str, Any]) -> list[dict[str, Any]]:
|
||||
def _parse_tiles_from_form(data):
|
||||
"""Parse tiles submitted via multipart/form-data into a list of tile dicts."""
|
||||
try:
|
||||
padding = int(data.get('padding', 0)) if data.get('padding') is not None else 0
|
||||
@@ -53,18 +51,14 @@ def parse_tiles_from_form(data: Mapping[str, Any]) -> list[dict[str, Any]]:
|
||||
if 'batch_idx' in meta:
|
||||
try:
|
||||
tile_info['batch_idx'] = int(meta['batch_idx'])
|
||||
except Exception as exc:
|
||||
raise ValueError(f"Invalid batch_idx for tile {i}: {meta.get('batch_idx')} ({exc})")
|
||||
except Exception:
|
||||
pass
|
||||
if 'global_idx' in meta:
|
||||
try:
|
||||
tile_info['global_idx'] = int(meta['global_idx'])
|
||||
except Exception as exc:
|
||||
raise ValueError(f"Invalid global_idx for tile {i}: {meta.get('global_idx')} ({exc})")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
tiles.append(tile_info)
|
||||
|
||||
return tiles
|
||||
|
||||
|
||||
# Backward compatibility for existing imports.
|
||||
_parse_tiles_from_form = parse_tiles_from_form
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class UpscaleCoreArgs:
|
||||
"""Shared denoise/sampling arguments for USDU tile processing."""
|
||||
|
||||
model: Any
|
||||
positive: Any
|
||||
negative: Any
|
||||
vae: Any
|
||||
seed: int
|
||||
steps: int
|
||||
cfg: float
|
||||
sampler_name: str
|
||||
scheduler: str
|
||||
denoise: float
|
||||
tiled_decode: bool
|
||||
+119
-144
@@ -4,8 +4,8 @@ import server
|
||||
from ..utils.constants import DYNAMIC_MODE_MAX_POLL_TIMEOUT, HEARTBEAT_INTERVAL
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.config import get_worker_timeout_seconds
|
||||
from .job_store import ensure_tile_jobs_initialized, mark_task_completed
|
||||
from .job_timeout import check_and_requeue_timed_out_workers
|
||||
from .job_store import ensure_tile_jobs_initialized, _mark_task_completed
|
||||
from .job_timeout import _check_and_requeue_timed_out_workers
|
||||
from .job_models import BaseJobState, ImageJobState, TileJobState
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ class ResultCollectorMixin:
|
||||
Expected co-mixins/attributes:
|
||||
- JobStateMixin methods for queue/task access.
|
||||
- `self._check_and_requeue_timed_out_workers(...)` coroutine.
|
||||
- `self.async_yield(...)` optional helper from WorkerCommsMixin.
|
||||
- `self._async_yield(...)` optional helper from WorkerCommsMixin.
|
||||
"""
|
||||
|
||||
def _log_worker_timeout_status(self, job_data, current_time: float, multi_job_id: str) -> list[str]:
|
||||
@@ -33,191 +33,166 @@ class ResultCollectorMixin:
|
||||
)
|
||||
return list(worker_status.keys())
|
||||
|
||||
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
|
||||
"""Default timeout requeue hook; override in host mixins when needed."""
|
||||
return await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
|
||||
|
||||
def _get_job_data_snapshot(self, prompt_server, multi_job_id):
|
||||
"""Get a snapshot of job data for timeout logging (non-async helper)."""
|
||||
current_job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
|
||||
if isinstance(current_job_data, BaseJobState):
|
||||
return current_job_data
|
||||
return None
|
||||
|
||||
async def _async_collect_worker_tiles(self, multi_job_id, num_workers):
|
||||
"""Collect tiles from workers in static mode."""
|
||||
async def _async_collect_results(self, multi_job_id, num_workers, mode='static',
|
||||
remaining_to_collect=None, batch_size=None):
|
||||
"""Unified async helper to collect results from workers (tiles or images)."""
|
||||
# Get the already initialized queue
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_tile_jobs:
|
||||
raise RuntimeError(f"Job queue not initialized for {multi_job_id}")
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, TileJobState):
|
||||
raise RuntimeError(
|
||||
f"Mode mismatch: expected static, got {getattr(job_data, 'mode', 'unknown')}"
|
||||
)
|
||||
q = job_data.queue
|
||||
expected_count = len(job_data.completed_tasks) + job_data.pending_tasks.qsize()
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} tiles from {num_workers} workers")
|
||||
|
||||
if mode == 'dynamic':
|
||||
if not isinstance(job_data, ImageJobState):
|
||||
raise RuntimeError(
|
||||
f"Mode mismatch: expected dynamic, got {getattr(job_data, 'mode', 'unknown')}"
|
||||
)
|
||||
q = job_data.queue
|
||||
completed_images = job_data.completed_images
|
||||
expected_count = remaining_to_collect or batch_size
|
||||
elif mode == 'static':
|
||||
if not isinstance(job_data, TileJobState):
|
||||
raise RuntimeError(
|
||||
f"Mode mismatch: expected static, got {getattr(job_data, 'mode', 'unknown')}"
|
||||
)
|
||||
q = job_data.queue
|
||||
expected_count = len(job_data.completed_tasks) + job_data.pending_tasks.qsize()
|
||||
else:
|
||||
raise RuntimeError(f"Unsupported mode: {mode}")
|
||||
|
||||
item_type = "images" if mode == 'dynamic' else "tiles"
|
||||
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} {item_type} from {num_workers} workers")
|
||||
|
||||
collected_results = {}
|
||||
workers_done = set()
|
||||
timeout = float(get_worker_timeout_seconds())
|
||||
wait_started_at = time.time()
|
||||
|
||||
while len(workers_done) < num_workers:
|
||||
if comfy.model_management.processing_interrupted():
|
||||
log("Processing interrupted by user")
|
||||
raise comfy.model_management.InterruptProcessingException()
|
||||
|
||||
job_data_snapshot = None
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
job_data_snapshot = self._get_job_data_snapshot(prompt_server, multi_job_id)
|
||||
|
||||
try:
|
||||
result = await asyncio.wait_for(q.get(), timeout=timeout)
|
||||
worker_id = result['worker_id']
|
||||
is_last = result.get('is_last', False)
|
||||
|
||||
tiles = result.get('tiles', [])
|
||||
debug_log(
|
||||
f"UltimateSDUpscale Master - Received batch of {len(tiles)} tiles from worker "
|
||||
f"'{worker_id}' (is_last={is_last})"
|
||||
)
|
||||
|
||||
for tile_data in tiles:
|
||||
if 'batch_idx' not in tile_data:
|
||||
log("UltimateSDUpscale Master - Missing batch_idx in tile data, skipping")
|
||||
continue
|
||||
|
||||
tile_idx = tile_data['tile_idx']
|
||||
key = tile_data.get('global_idx', tile_idx)
|
||||
entry = {
|
||||
**tile_data,
|
||||
'tile_idx': tile_idx,
|
||||
'worker_id': worker_id,
|
||||
'global_idx': key,
|
||||
}
|
||||
collected_results[entry['global_idx']] = entry
|
||||
|
||||
if is_last:
|
||||
workers_done.add(worker_id)
|
||||
debug_log(f"UltimateSDUpscale Master - Worker {worker_id} completed")
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
current_time = time.time()
|
||||
waiting_workers = type(self)._log_worker_timeout_status(
|
||||
self, job_data_snapshot, current_time, multi_job_id,
|
||||
)
|
||||
elapsed = current_time - wait_started_at
|
||||
log(
|
||||
f"UltimateSDUpscale Master - Heartbeat timeout waiting for tiles; "
|
||||
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
|
||||
)
|
||||
break
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {len(collected_results)} tiles from {len(workers_done)} workers")
|
||||
|
||||
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]
|
||||
|
||||
return collected_results
|
||||
|
||||
async def collect_dynamic_images(
|
||||
self,
|
||||
multi_job_id,
|
||||
remaining_to_collect,
|
||||
num_workers,
|
||||
batch_size,
|
||||
master_processed_count,
|
||||
):
|
||||
"""Collect remaining processed images from workers in dynamic mode."""
|
||||
_ = master_processed_count
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
if multi_job_id not in prompt_server.distributed_pending_tile_jobs:
|
||||
raise RuntimeError(f"Job queue not initialized for {multi_job_id}")
|
||||
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
|
||||
if not isinstance(job_data, ImageJobState):
|
||||
raise RuntimeError(
|
||||
f"Mode mismatch: expected dynamic, got {getattr(job_data, 'mode', 'unknown')}"
|
||||
)
|
||||
q = job_data.queue
|
||||
completed_images = job_data.completed_images
|
||||
expected_count = remaining_to_collect or batch_size
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} images from {num_workers} workers")
|
||||
|
||||
workers_done = set()
|
||||
# Unify collector/upscaler wait behavior with the UI worker timeout
|
||||
timeout = float(get_worker_timeout_seconds())
|
||||
last_heartbeat_check = time.time()
|
||||
wait_started_at = time.time()
|
||||
collected_count = 0
|
||||
|
||||
|
||||
while len(workers_done) < num_workers:
|
||||
# Check for user interruption
|
||||
if comfy.model_management.processing_interrupted():
|
||||
log("Processing interrupted by user")
|
||||
raise comfy.model_management.InterruptProcessingException()
|
||||
|
||||
if remaining_to_collect and collected_count >= remaining_to_collect:
|
||||
|
||||
# For dynamic mode with remaining_to_collect, check if we've collected enough
|
||||
if mode == 'dynamic' and remaining_to_collect and collected_count >= remaining_to_collect:
|
||||
break
|
||||
|
||||
job_data_snapshot = None
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
job_data_snapshot = self._get_job_data_snapshot(prompt_server, multi_job_id)
|
||||
current_job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
|
||||
if isinstance(current_job_data, BaseJobState):
|
||||
job_data_snapshot = current_job_data
|
||||
|
||||
try:
|
||||
wait_timeout = min(DYNAMIC_MODE_MAX_POLL_TIMEOUT, timeout)
|
||||
# Shorter poll for dynamic mode, but never exceed the configured timeout
|
||||
wait_timeout = (min(DYNAMIC_MODE_MAX_POLL_TIMEOUT, timeout) if mode == 'dynamic' else timeout)
|
||||
result = await asyncio.wait_for(q.get(), timeout=wait_timeout)
|
||||
worker_id = result['worker_id']
|
||||
is_last = result.get('is_last', False)
|
||||
|
||||
if mode == 'static':
|
||||
# Handle tiles
|
||||
tiles = result.get('tiles', [])
|
||||
debug_log(
|
||||
f"UltimateSDUpscale Master - Received batch of {len(tiles)} tiles from worker "
|
||||
f"'{worker_id}' (is_last={is_last})"
|
||||
)
|
||||
|
||||
if 'image_idx' in result and 'image' in result:
|
||||
image_idx = result['image_idx']
|
||||
image_pil = result['image']
|
||||
completed_images[image_idx] = image_pil
|
||||
collected_count += 1
|
||||
debug_log(f"UltimateSDUpscale Master - Received image {image_idx} from worker {worker_id}")
|
||||
for tile_data in tiles:
|
||||
if 'batch_idx' not in tile_data:
|
||||
log("UltimateSDUpscale Master - Missing batch_idx in tile data, skipping")
|
||||
continue
|
||||
|
||||
tile_idx = tile_data['tile_idx']
|
||||
key = tile_data.get('global_idx', tile_idx)
|
||||
entry = {
|
||||
'tile_idx': tile_idx,
|
||||
'x': tile_data['x'],
|
||||
'y': tile_data['y'],
|
||||
'extracted_width': tile_data['extracted_width'],
|
||||
'extracted_height': tile_data['extracted_height'],
|
||||
'padding': tile_data['padding'],
|
||||
'worker_id': worker_id,
|
||||
'batch_idx': tile_data.get('batch_idx', 0),
|
||||
'global_idx': tile_data.get('global_idx', tile_idx),
|
||||
}
|
||||
if 'image' in tile_data:
|
||||
entry['image'] = tile_data['image']
|
||||
elif 'tensor' in tile_data:
|
||||
entry['tensor'] = tile_data['tensor']
|
||||
collected_results[key] = entry
|
||||
|
||||
elif mode == 'dynamic':
|
||||
# Handle full images
|
||||
if 'image_idx' in result and 'image' in result:
|
||||
image_idx = result['image_idx']
|
||||
image_pil = result['image']
|
||||
completed_images[image_idx] = image_pil
|
||||
collected_results[image_idx] = image_pil
|
||||
collected_count += 1
|
||||
debug_log(f"UltimateSDUpscale Master - Received image {image_idx} from worker {worker_id}")
|
||||
|
||||
if is_last:
|
||||
workers_done.add(worker_id)
|
||||
debug_log(f"UltimateSDUpscale Master - Worker {worker_id} completed")
|
||||
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
current_time = time.time()
|
||||
waiting_workers = type(self)._log_worker_timeout_status(
|
||||
self, job_data_snapshot, current_time, multi_job_id,
|
||||
)
|
||||
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
|
||||
requeued = await type(self)._check_and_requeue_timed_out_workers(
|
||||
self, multi_job_id, batch_size,
|
||||
)
|
||||
if requeued > 0:
|
||||
log(f"UltimateSDUpscale Master - Requeued {requeued} images from timed out workers")
|
||||
last_heartbeat_check = current_time
|
||||
|
||||
if current_time - wait_started_at > timeout:
|
||||
waiting_workers = self._log_worker_timeout_status(job_data_snapshot, current_time, multi_job_id)
|
||||
if mode == 'dynamic':
|
||||
# Check for worker timeouts periodically
|
||||
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
|
||||
# Use the class method to check and requeue
|
||||
requeued = await self._check_and_requeue_timed_out_workers(multi_job_id, batch_size)
|
||||
if requeued > 0:
|
||||
log(f"UltimateSDUpscale Master - Requeued {requeued} images from timed out workers")
|
||||
last_heartbeat_check = current_time
|
||||
|
||||
# Check if we've been waiting too long overall
|
||||
if current_time - wait_started_at > timeout:
|
||||
elapsed = current_time - wait_started_at
|
||||
log(
|
||||
"UltimateSDUpscale Master - Heartbeat timeout while waiting for images; "
|
||||
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
|
||||
)
|
||||
break
|
||||
else:
|
||||
elapsed = current_time - wait_started_at
|
||||
log(
|
||||
"UltimateSDUpscale Master - Heartbeat timeout while waiting for images; "
|
||||
f"UltimateSDUpscale Master - Heartbeat timeout waiting for {item_type}; "
|
||||
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
|
||||
)
|
||||
break
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {collected_count} images from {len(workers_done)} workers")
|
||||
|
||||
|
||||
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {len(collected_results)} {item_type} from {len(workers_done)} workers")
|
||||
|
||||
# Clean up job queue
|
||||
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]
|
||||
|
||||
return collected_results if mode == 'static' else completed_images
|
||||
|
||||
return completed_images
|
||||
async def _async_collect_worker_tiles(self, multi_job_id, num_workers):
|
||||
"""Async helper to collect tiles from workers."""
|
||||
return await self._async_collect_results(multi_job_id, num_workers, mode='static')
|
||||
|
||||
async def mark_image_completed(self, multi_job_id, image_idx, image_pil):
|
||||
async def _mark_image_completed(self, multi_job_id, image_idx, image_pil):
|
||||
"""Mark an image as completed in the job data."""
|
||||
await mark_task_completed(multi_job_id, image_idx, {'image': image_pil})
|
||||
# Mark the image as completed with the image data
|
||||
await _mark_task_completed(multi_job_id, image_idx, {'image': image_pil})
|
||||
prompt_server = ensure_tile_jobs_initialized()
|
||||
async with prompt_server.distributed_tile_jobs_lock:
|
||||
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
|
||||
if isinstance(job_data, ImageJobState):
|
||||
job_data.completed_images[image_idx] = image_pil
|
||||
|
||||
async def _async_collect_dynamic_images(self, multi_job_id, remaining_to_collect, num_workers, batch_size, master_processed_count):
|
||||
"""Collect remaining processed images from workers."""
|
||||
return await self._async_collect_results(multi_job_id, num_workers, mode='dynamic',
|
||||
remaining_to_collect=remaining_to_collect,
|
||||
batch_size=batch_size)
|
||||
|
||||
@@ -371,10 +371,6 @@ class TileOpsMixin:
|
||||
|
||||
return positive_sliced, negative_sliced
|
||||
|
||||
def slice_conditioning(self, positive, negative, batch_idx):
|
||||
"""Public conditioning-slice API."""
|
||||
return self._slice_conditioning(positive, negative, batch_idx)
|
||||
|
||||
def _process_and_blend_tile(self, tile_idx, tile_pos, upscaled_image, result_image,
|
||||
model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height,
|
||||
@@ -403,59 +399,6 @@ class TileOpsMixin:
|
||||
|
||||
return result_image
|
||||
|
||||
def process_and_blend_tile(
|
||||
self,
|
||||
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: int = 0,
|
||||
):
|
||||
"""Public tile-processing API."""
|
||||
return self._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,
|
||||
)
|
||||
|
||||
def _process_single_tile(self, global_idx, num_tiles_per_image, upscaled_image, all_tiles,
|
||||
model, positive, negative, vae, seed, steps, cfg, sampler_name,
|
||||
scheduler, denoise, tiled_decode, tile_width, tile_height, padding,
|
||||
|
||||
@@ -1,56 +0,0 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from .processing_args import UpscaleCoreArgs
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TileBatchArgs:
|
||||
"""Tile/canvas parameters layered on top of shared core processing args."""
|
||||
|
||||
core: UpscaleCoreArgs
|
||||
tile_width: int
|
||||
tile_height: int
|
||||
padding: int
|
||||
force_uniform_tiles: bool
|
||||
width: int
|
||||
height: int
|
||||
|
||||
|
||||
def extract_and_process_tile_batch(
|
||||
*,
|
||||
node: Any,
|
||||
upscaled_image: Any,
|
||||
tx: int,
|
||||
ty: int,
|
||||
args: TileBatchArgs,
|
||||
) -> tuple[Any, int, int, int, int]:
|
||||
"""Extract one tile position for the whole batch and process it."""
|
||||
tile_batch, x1, y1, ew, eh = node.extract_batch_tile_with_padding(
|
||||
upscaled_image,
|
||||
tx,
|
||||
ty,
|
||||
args.tile_width,
|
||||
args.tile_height,
|
||||
args.padding,
|
||||
args.force_uniform_tiles,
|
||||
)
|
||||
region = (x1, y1, x1 + ew, y1 + eh)
|
||||
core = args.core
|
||||
processed_batch = node.process_tiles_batch(
|
||||
tile_batch,
|
||||
core.model,
|
||||
core.positive,
|
||||
core.negative,
|
||||
core.vae,
|
||||
core.seed,
|
||||
core.steps,
|
||||
core.cfg,
|
||||
core.sampler_name,
|
||||
core.scheduler,
|
||||
core.denoise,
|
||||
core.tiled_decode,
|
||||
region,
|
||||
(args.width, args.height),
|
||||
)
|
||||
return processed_batch, x1, y1, ew, eh
|
||||
+91
-209
@@ -1,84 +1,20 @@
|
||||
import asyncio, io, json, time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
import aiohttp
|
||||
from PIL import Image
|
||||
from ..utils.logging import debug_log, log
|
||||
from ..utils.auth import distributed_auth_headers
|
||||
from ..utils.config import load_config
|
||||
from ..utils.network import get_client_session
|
||||
from ..utils.constants import TILE_SEND_TIMEOUT
|
||||
from ..utils.usdu_management import MAX_PAYLOAD_SIZE, send_heartbeat_to_master
|
||||
from ..utils.usdu_managment import MAX_PAYLOAD_SIZE, _send_heartbeat_to_master
|
||||
from ..utils.image import tensor_to_pil
|
||||
|
||||
|
||||
WorkAssignmentKind = Literal["image", "tile", "none"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WorkAssignment:
|
||||
"""Canonical typed representation of a single work-item assignment."""
|
||||
|
||||
kind: WorkAssignmentKind
|
||||
task_idx: int | None
|
||||
estimated_remaining: int = 0
|
||||
batched_static: bool = False
|
||||
|
||||
|
||||
class WorkerCommsMixin:
|
||||
@staticmethod
|
||||
def _master_auth_headers() -> dict[str, str]:
|
||||
return distributed_auth_headers(load_config())
|
||||
async def _send_heartbeat_to_master(self, multi_job_id, master_url, worker_id):
|
||||
"""Proxy heartbeat helper used by worker processing mixins."""
|
||||
await _send_heartbeat_to_master(multi_job_id, master_url, worker_id)
|
||||
|
||||
async def _post_with_retry(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
build_form: callable,
|
||||
max_retries: int = 5,
|
||||
initial_delay: float = 0.5,
|
||||
max_delay: float = 5.0,
|
||||
error_context: str = "",
|
||||
) -> None:
|
||||
"""POST form data with exponential backoff. build_form is called fresh each attempt."""
|
||||
retry_delay = initial_delay
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
session = await get_client_session()
|
||||
async with session.post(
|
||||
url,
|
||||
data=build_form(),
|
||||
headers=self._master_auth_headers(),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
return
|
||||
except Exception as e:
|
||||
if attempt < max_retries - 1:
|
||||
debug_log(f"Retry {attempt + 1}/{max_retries} after error: {e}")
|
||||
await asyncio.sleep(retry_delay)
|
||||
retry_delay = min(retry_delay * 2, max_delay)
|
||||
else:
|
||||
log(f"{error_context}: {e}")
|
||||
raise
|
||||
|
||||
async def send_heartbeat(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
worker_id: str,
|
||||
) -> None:
|
||||
"""Send worker heartbeat to the master."""
|
||||
await send_heartbeat_to_master(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:
|
||||
async def send_tiles_batch_to_master(self, processed_tiles, multi_job_id, master_url,
|
||||
padding, worker_id, is_final_flush=False):
|
||||
"""Send all processed tiles to master, chunked if large."""
|
||||
if not processed_tiles:
|
||||
if is_final_flush:
|
||||
@@ -114,8 +50,12 @@ class WorkerCommsMixin:
|
||||
i = 0
|
||||
chunk_index = 0
|
||||
while i < total_tiles:
|
||||
data = aiohttp.FormData()
|
||||
data.add_field('multi_job_id', multi_job_id)
|
||||
data.add_field('worker_id', str(worker_id))
|
||||
data.add_field('padding', str(padding))
|
||||
|
||||
metadata = []
|
||||
chunk_images: list[tuple[int, int, bytes]] = []
|
||||
used = 0
|
||||
j = i
|
||||
while j < total_tiles:
|
||||
@@ -127,7 +67,7 @@ class WorkerCommsMixin:
|
||||
break
|
||||
# Accept this tile in this chunk
|
||||
metadata.append(meta)
|
||||
chunk_images.append((j - i, j, img_bytes))
|
||||
data.add_field(f'tile_{j - i}', io.BytesIO(img_bytes), filename=f'tile_{j}.png', content_type='image/png')
|
||||
used += len(img_bytes) + overhead
|
||||
j += 1
|
||||
|
||||
@@ -136,34 +76,32 @@ class WorkerCommsMixin:
|
||||
# Single oversized tile, send anyway
|
||||
meta = encoded[j]['meta']
|
||||
metadata.append(meta)
|
||||
chunk_images.append((0, j, encoded[j]['bytes']))
|
||||
data.add_field('tile_0', io.BytesIO(encoded[j]['bytes']), filename=f'tile_{j}.png', content_type='image/png')
|
||||
j += 1
|
||||
|
||||
chunk_size = j - i
|
||||
is_chunk_last = (j >= total_tiles)
|
||||
data.add_field('is_last', str(bool(is_final_flush and is_chunk_last)))
|
||||
data.add_field('batch_size', str(chunk_size))
|
||||
data.add_field('tiles_metadata', json.dumps(metadata), content_type='application/json')
|
||||
|
||||
def _build_chunk_form() -> aiohttp.FormData:
|
||||
data = aiohttp.FormData()
|
||||
data.add_field('multi_job_id', multi_job_id)
|
||||
data.add_field('worker_id', str(worker_id))
|
||||
data.add_field('padding', str(padding))
|
||||
data.add_field('is_last', str(bool(is_final_flush and is_chunk_last)))
|
||||
data.add_field('batch_size', str(chunk_size))
|
||||
data.add_field('tiles_metadata', json.dumps(metadata), content_type='application/json')
|
||||
for relative_idx, source_idx, img_bytes in chunk_images:
|
||||
data.add_field(
|
||||
f'tile_{relative_idx}',
|
||||
io.BytesIO(img_bytes),
|
||||
filename=f'tile_{source_idx}.png',
|
||||
content_type='image/png',
|
||||
)
|
||||
return data
|
||||
|
||||
await self._post_with_retry(
|
||||
f"{master_url}/distributed/submit_tiles",
|
||||
build_form=_build_chunk_form,
|
||||
error_context=f"UltimateSDUpscale Worker - Failed to send chunk {chunk_index}",
|
||||
)
|
||||
# Retry logic with exponential backoff
|
||||
max_retries = 5
|
||||
retry_delay = 0.5
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/submit_tiles"
|
||||
async with session.post(url, data=data) as response:
|
||||
response.raise_for_status()
|
||||
break
|
||||
except Exception as e:
|
||||
if attempt < max_retries - 1:
|
||||
await asyncio.sleep(retry_delay)
|
||||
retry_delay = min(retry_delay * 2, 5.0)
|
||||
else:
|
||||
log(f"UltimateSDUpscale Worker - Failed to send chunk {chunk_index} after {max_retries} attempts: {e}")
|
||||
raise
|
||||
|
||||
debug_log(f"Worker[{worker_id[:8]}] - Sent chunk {chunk_index} ({chunk_size} tiles, ~{used/1e6:.2f} MB)")
|
||||
chunk_index += 1
|
||||
@@ -179,7 +117,7 @@ class WorkerCommsMixin:
|
||||
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/submit_tiles"
|
||||
async with session.post(url, data=data, headers=self._master_auth_headers()) as response:
|
||||
async with session.post(url, data=data) as response:
|
||||
response.raise_for_status()
|
||||
debug_log(f"Worker {worker_id} sent static completion signal")
|
||||
|
||||
@@ -206,7 +144,7 @@ class WorkerCommsMixin:
|
||||
async with session.post(url, json={
|
||||
'worker_id': str(worker_id),
|
||||
'multi_job_id': multi_job_id
|
||||
}, headers=self._master_auth_headers()) as response:
|
||||
}) as response:
|
||||
if response.status == 200:
|
||||
return await response.json()
|
||||
if response.status == 404:
|
||||
@@ -230,92 +168,66 @@ class WorkerCommsMixin:
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _parse_work_assignment(data: dict[str, Any] | None) -> WorkAssignment:
|
||||
"""Normalize assignment payloads into a single discriminated contract."""
|
||||
if not data:
|
||||
return WorkAssignment(kind="none", task_idx=None)
|
||||
|
||||
kind_value = data.get("kind")
|
||||
if kind_value is None:
|
||||
# Backward-compatible parsing for older masters.
|
||||
if "image_idx" in data:
|
||||
kind_value = "image"
|
||||
data = {
|
||||
"kind": "image",
|
||||
"task_idx": data.get("image_idx"),
|
||||
"estimated_remaining": data.get("estimated_remaining", 0),
|
||||
}
|
||||
elif "tile_idx" in data:
|
||||
kind_value = "tile"
|
||||
data = {
|
||||
"kind": "tile",
|
||||
"task_idx": data.get("tile_idx"),
|
||||
"estimated_remaining": data.get("estimated_remaining", 0),
|
||||
"batched_static": data.get("batched_static", False),
|
||||
}
|
||||
else:
|
||||
kind_value = "none"
|
||||
|
||||
if kind_value not in {"image", "tile", "none"}:
|
||||
return WorkAssignment(kind="none", task_idx=None)
|
||||
|
||||
task_idx_raw = data.get("task_idx")
|
||||
if task_idx_raw is None:
|
||||
task_idx = None
|
||||
else:
|
||||
try:
|
||||
task_idx = int(task_idx_raw)
|
||||
except (TypeError, ValueError):
|
||||
task_idx = None
|
||||
|
||||
try:
|
||||
estimated_remaining = int(data.get("estimated_remaining", 0) or 0)
|
||||
except (TypeError, ValueError):
|
||||
estimated_remaining = 0
|
||||
|
||||
return WorkAssignment(
|
||||
kind=kind_value,
|
||||
task_idx=task_idx,
|
||||
estimated_remaining=estimated_remaining,
|
||||
batched_static=bool(data.get("batched_static", False)),
|
||||
)
|
||||
|
||||
async def request_assignment(self, multi_job_id, master_url, worker_id) -> WorkAssignment:
|
||||
"""Request one assignment and parse into the canonical discriminated contract."""
|
||||
async def _request_image_from_master(self, multi_job_id, master_url, worker_id):
|
||||
"""Request an image index to process from master in dynamic mode."""
|
||||
data = await self._request_work_item_from_master(multi_job_id, master_url, worker_id)
|
||||
return self._parse_work_assignment(data)
|
||||
if not data:
|
||||
return None, 0
|
||||
image_idx = data.get('image_idx')
|
||||
estimated_remaining = data.get('estimated_remaining', 0)
|
||||
return image_idx, estimated_remaining
|
||||
|
||||
async def send_full_image(self, image_pil, image_idx, multi_job_id,
|
||||
master_url, worker_id, is_last):
|
||||
async def _request_tile_from_master(self, multi_job_id, master_url, worker_id):
|
||||
"""Request a tile index to process from master in static mode (reusing dynamic infrastructure)."""
|
||||
data = await self._request_work_item_from_master(multi_job_id, master_url, worker_id)
|
||||
if not data:
|
||||
return None, 0, False
|
||||
tile_idx = data.get('tile_idx')
|
||||
estimated_remaining = data.get('estimated_remaining', 0)
|
||||
batched_static = data.get('batched_static', False)
|
||||
return tile_idx, estimated_remaining, batched_static
|
||||
|
||||
async def _send_full_image_to_master(self, image_pil, image_idx, multi_job_id,
|
||||
master_url, worker_id, is_last):
|
||||
"""Send a processed full image back to master in dynamic mode."""
|
||||
# Serialize image to PNG
|
||||
byte_io = io.BytesIO()
|
||||
image_pil.save(byte_io, format='PNG', compress_level=0)
|
||||
image_bytes = byte_io.getvalue()
|
||||
|
||||
def _build_image_form() -> aiohttp.FormData:
|
||||
data = aiohttp.FormData()
|
||||
data.add_field('multi_job_id', multi_job_id)
|
||||
data.add_field('worker_id', str(worker_id))
|
||||
data.add_field('image_idx', str(image_idx))
|
||||
data.add_field('is_last', str(is_last))
|
||||
data.add_field(
|
||||
'full_image',
|
||||
io.BytesIO(image_bytes),
|
||||
filename=f'image_{image_idx}.png',
|
||||
content_type='image/png',
|
||||
)
|
||||
return data
|
||||
byte_io.seek(0)
|
||||
|
||||
await self._post_with_retry(
|
||||
f"{master_url}/distributed/submit_image",
|
||||
build_form=_build_image_form,
|
||||
error_context=f"Failed to send image {image_idx}",
|
||||
)
|
||||
debug_log(f"Successfully sent image {image_idx} to master")
|
||||
# Prepare form data
|
||||
data = aiohttp.FormData()
|
||||
data.add_field('multi_job_id', multi_job_id)
|
||||
data.add_field('worker_id', str(worker_id))
|
||||
data.add_field('image_idx', str(image_idx))
|
||||
data.add_field('is_last', str(is_last))
|
||||
data.add_field('full_image', byte_io, filename=f'image_{image_idx}.png',
|
||||
content_type='image/png')
|
||||
|
||||
# Retry logic
|
||||
max_retries = 5
|
||||
retry_delay = 0.5
|
||||
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/submit_image"
|
||||
|
||||
async with session.post(url, data=data) as response:
|
||||
response.raise_for_status()
|
||||
debug_log(f"Successfully sent image {image_idx} to master")
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
if attempt < max_retries - 1:
|
||||
debug_log(f"Retry {attempt + 1}/{max_retries} after error: {e}")
|
||||
await asyncio.sleep(retry_delay)
|
||||
retry_delay *= 2
|
||||
else:
|
||||
log(f"Failed to send image {image_idx} after {max_retries} attempts: {e}")
|
||||
raise
|
||||
|
||||
async def send_worker_complete_signal(self, multi_job_id, master_url, worker_id):
|
||||
async def _send_worker_complete_signal(self, multi_job_id, master_url, worker_id):
|
||||
"""Send completion signal to master in dynamic mode."""
|
||||
# Send a dummy request with is_last=True
|
||||
data = aiohttp.FormData()
|
||||
@@ -327,16 +239,16 @@ class WorkerCommsMixin:
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/submit_image"
|
||||
|
||||
async with session.post(url, data=data, headers=self._master_auth_headers()) as response:
|
||||
async with session.post(url, data=data) as response:
|
||||
response.raise_for_status()
|
||||
debug_log(f"Worker {worker_id} sent completion signal")
|
||||
|
||||
async def check_job_status(self, multi_job_id, master_url):
|
||||
async def _check_job_status(self, multi_job_id, master_url):
|
||||
"""Check if job is ready on the master."""
|
||||
try:
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/job_status?multi_job_id={multi_job_id}"
|
||||
async with session.get(url, headers=self._master_auth_headers()) as response:
|
||||
async with session.get(url) as response:
|
||||
if response.status == 200:
|
||||
data = await response.json()
|
||||
return data.get('ready', False)
|
||||
@@ -345,36 +257,6 @@ class WorkerCommsMixin:
|
||||
debug_log(f"Job status check failed: {e}")
|
||||
return False
|
||||
|
||||
async def init_dynamic_job_on_master(
|
||||
self,
|
||||
multi_job_id: str,
|
||||
master_url: str,
|
||||
batch_size: int,
|
||||
enabled_workers: list[str],
|
||||
) -> bool:
|
||||
"""Tell the master to create the dynamic job queue (idempotent)."""
|
||||
try:
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/init_dynamic_job"
|
||||
async with session.post(
|
||||
url,
|
||||
json={
|
||||
"multi_job_id": multi_job_id,
|
||||
"batch_size": batch_size,
|
||||
"enabled_workers": enabled_workers,
|
||||
},
|
||||
headers=self._master_auth_headers(),
|
||||
) as response:
|
||||
if response.status == 200:
|
||||
debug_log(f"Worker initialized dynamic job {multi_job_id} on master (batch_size={batch_size})")
|
||||
return True
|
||||
text = await response.text()
|
||||
debug_log(f"init_dynamic_job failed ({response.status}): {text}")
|
||||
return False
|
||||
except Exception as e:
|
||||
debug_log(f"init_dynamic_job request failed: {e}")
|
||||
return False
|
||||
|
||||
async def async_yield(self):
|
||||
async def _async_yield(self):
|
||||
"""Simple async yield to allow event loop processing."""
|
||||
await asyncio.sleep(0)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
"""
|
||||
Utility modules for ComfyUI-Distributed extension.
|
||||
"""
|
||||
|
||||
# Make utils importable as a package
|
||||
+17
-14
@@ -6,6 +6,7 @@ import threading
|
||||
import time
|
||||
import uuid
|
||||
import execution
|
||||
import server
|
||||
from typing import Optional, Any, Coroutine
|
||||
from .network import get_server_loop
|
||||
|
||||
@@ -42,11 +43,10 @@ def run_async_in_server_loop(coro: Coroutine, timeout: Optional[float] = None) -
|
||||
|
||||
# Schedule on server's event loop
|
||||
loop = get_server_loop()
|
||||
task_future = asyncio.run_coroutine_threadsafe(wrapper(), loop)
|
||||
asyncio.run_coroutine_threadsafe(wrapper(), loop)
|
||||
|
||||
# Wait for completion
|
||||
if not event.wait(timeout):
|
||||
task_future.cancel()
|
||||
raise TimeoutError(f"Async operation timed out after {timeout} seconds")
|
||||
|
||||
if error:
|
||||
@@ -54,10 +54,7 @@ def run_async_in_server_loop(coro: Coroutine, timeout: Optional[float] = None) -
|
||||
return result
|
||||
|
||||
|
||||
def _prompt_server_instance():
|
||||
import server
|
||||
|
||||
return server.PromptServer.instance
|
||||
prompt_server = server.PromptServer.instance
|
||||
|
||||
|
||||
def _summarize_node_errors(node_errors: dict) -> str:
|
||||
@@ -109,12 +106,12 @@ class PromptValidationError(RuntimeError):
|
||||
|
||||
|
||||
async def queue_prompt_payload(
|
||||
prompt_obj: dict[str, Any],
|
||||
workflow_meta: dict[str, Any] | None = None,
|
||||
client_id: str | None = None,
|
||||
) -> str:
|
||||
prompt_obj,
|
||||
workflow_meta=None,
|
||||
client_id=None,
|
||||
include_queue_metadata=False,
|
||||
):
|
||||
"""Validate and queue a prompt via ComfyUI's prompt queue."""
|
||||
prompt_server = _prompt_server_instance()
|
||||
payload = {"prompt": prompt_obj}
|
||||
payload = prompt_server.trigger_on_prompt(payload)
|
||||
prompt = payload["prompt"]
|
||||
@@ -126,13 +123,11 @@ async def queue_prompt_payload(
|
||||
node_errors = valid[3] if len(valid) > 3 else {}
|
||||
raise PromptValidationError(error_payload, node_errors)
|
||||
|
||||
extra_data = {}
|
||||
extra_data = {"create_time": int(time.time() * 1000)}
|
||||
if workflow_meta:
|
||||
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
|
||||
if client_id:
|
||||
extra_data["client_id"] = client_id
|
||||
# Keep parity with ComfyUI /prompt endpoint so Jobs API metadata stays valid.
|
||||
extra_data.setdefault("create_time", int(time.time() * 1000))
|
||||
|
||||
sensitive = {}
|
||||
for key in getattr(execution, "SENSITIVE_EXTRA_DATA_KEYS", []):
|
||||
@@ -143,4 +138,12 @@ async def queue_prompt_payload(
|
||||
prompt_server.number = number + 1
|
||||
prompt_queue_item = (number, prompt_id, prompt, extra_data, valid[2], sensitive)
|
||||
prompt_server.prompt_queue.put(prompt_queue_item)
|
||||
|
||||
if include_queue_metadata:
|
||||
return {
|
||||
"prompt_id": prompt_id,
|
||||
"number": number,
|
||||
"node_errors": {},
|
||||
}
|
||||
|
||||
return prompt_id
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import base64
|
||||
import binascii
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -14,7 +13,7 @@ MAX_AUDIO_PAYLOAD_BYTES = int(
|
||||
)
|
||||
|
||||
|
||||
def encode_audio_payload(audio_payload: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
def encode_audio_payload(audio_payload):
|
||||
"""Serialize an AUDIO dict into JSON-safe canonical envelope payload."""
|
||||
if not isinstance(audio_payload, dict):
|
||||
return None
|
||||
@@ -44,7 +43,7 @@ def encode_audio_payload(audio_payload: dict[str, Any] | None) -> dict[str, Any]
|
||||
}
|
||||
|
||||
|
||||
def decode_audio_payload(audio_payload: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
def decode_audio_payload(audio_payload):
|
||||
"""Decode canonical envelope audio payload into an AUDIO dict."""
|
||||
if audio_payload is None:
|
||||
return None
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user