Compare commits

...
45 Commits
Author SHA1 Message Date
Robert Wojciechowski 2fec61d567 Update video-upscaler-runpod-preset.md 2025-09-15 11:53:49 +10:00
Robert Wojciechowski 528cd2e895 Update pyproject.toml 2025-09-15 11:51:51 +10:00
Robert Wojciechowski 0572e084e3 Update video-upscaler-runpod-preset.md 2025-09-15 11:51:28 +10:00
Robert Wojciechowski bc9bba0c76 Update video-upscaler-runpod-preset.md 2025-09-15 11:50:33 +10:00
Robert Wojciechowski 701c6de9eb Update README.md 2025-09-15 11:44:08 +10:00
Robert Wojciechowski c846aac0e6 Update video-upscaler-runpod-preset.md 2025-09-15 11:41:32 +10:00
Robert Wojciechowski 0af3df071d Create video-upscaler-runpod-preset.md 2025-09-15 11:32:29 +10:00
robertvoy f7c7b960c7 change upscaler in video upscaler workflow 2025-09-14 09:54:10 +10:00
robertvoy c45b50830e Merge branch 'main' of https://github.com/robertvoy/ComfyUI-Distributed 2025-09-13 14:42:41 +10:00
robertvoy 213ac38227 Add new /prompt liveness probe (USDU + collector). when worker timeout value in exceeded, the master will probe the workers. 2025-09-13 14:42:33 +10:00
robertvoy 9ae354a788 upscaler workflow change 2025-09-13 12:21:26 +10:00
Robert Wojciechowski c1928ebd09 Update README.md 2025-09-13 10:41:52 +10:00
Robert Wojciechowski 98a63938c8 Update README.md 2025-09-13 10:40:03 +10:00
Robert Wojciechowski 27b73fd280 Update README.md 2025-09-13 10:38:55 +10:00
robertvoy 48c8792af2 video upscaler workflow 2025-09-13 10:34:06 +10:00
robertvoy 4b8836cd03 Unify worker timeout via UI across collector/upscaler (90s); add interruptible sliced waits + cleanup to collector 2025-09-13 10:26:03 +10:00
robertvoy b69fc40fbd updated video upscaler workflow 2025-09-12 09:58:53 +10:00
Robert Wojciechowski e5eb2c8b86 Delete workflows/distributed-upscale-batch.json 2025-09-12 09:49:48 +10:00
Robert Wojciechowski b90c3e42ac Update README.md 2025-09-12 09:49:27 +10:00
robertvoy 187b38d4f5 worker timeout in settings, fixed requeue video tiles 2025-09-11 16:40:19 +10:00
robertvoy 39fb4db486 warning for non-WAN compatible frame count and WanVideo Closest Frame Count 2025-09-11 14:26:53 +10:00
robertvoy 763e238120 fixed: if batch =! (4*n)+1 2025-09-11 12:51:13 +10:00
robertvoy 857aa80174 optimisations 2025-09-10 10:34:00 +10:00
robertvoy 287e4c05fa speed optimisation 2025-09-10 10:15:35 +10:00
robertvoy 58f26a106f clean up 2025-09-10 09:38:40 +10:00
robertvoy 138cac3eb7 remove static distribution 2025-09-09 19:52:34 +10:00
robertvoy 0e07868b98 initial 2025-09-09 19:18:57 +10:00
Robert Wojciechowski a8be7c7ed9 Update pyproject.toml 2025-08-25 13:38:01 +10:00
Robert Wojciechowski c9ca7f5390 Merge pull request #32 from doubletwisted/fix/pasted-paths
Update executionUtils.js
2025-08-25 13:37:34 +10:00
Dominik Bargiel c860df0aa7 Update executionUtils.js
Previous fix caused problems with loading models that are in subfolders like:
"'FLUX.1\\Shakker-Labs-ControlNet-Union-Pro-2.0.safetensors"

Now images and wideos are treated separetlly.
2025-08-24 19:46:40 +02:00
Robert Wojciechowski 05662eb7eb Update pyproject.toml 2025-08-16 10:49:56 +10:00
Robert Wojciechowski 8441c92d95 Merge pull request #29 from doubletwisted/fix/pasted-paths
fix: normalize pasted image paths and subfolder extraction
2025-08-16 10:49:27 +10:00
doubletwisted 4b513d49c3 fix: normalize pasted image paths and subfolder extraction (Windows-safe, upstream-only) 2025-08-15 21:45:19 +02:00
Robert Wojciechowski f91d2d579d Update pyproject.toml 2025-08-13 10:00:36 +10:00
Robert Wojciechowski 1ff863fb5d Add files via upload 2025-08-13 10:00:03 +10:00
Robert Wojciechowski 74bba258da Update worker-setup-guides.md 2025-08-09 12:40:53 +10:00
Robert Wojciechowski bb559e8ce0 Merge pull request #27 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-08-06 06:53:17 +10:00
Robert Wojciechowski c032de224b Update README.md 2025-08-05 07:58:40 +10:00
Robert Wojciechowski 68f8627b0b Update README.md 2025-08-04 11:09:36 +10:00
Robert Wojciechowski 0e3c04fba7 Update README.md 2025-08-04 11:02:55 +10:00
robertvoy e076cf3455 Merge branch 'main' of https://github.com/robertvoy/ComfyUI-Distributed 2025-08-03 18:38:06 +10:00
robertvoy 5f3d358f61 wan 2.2 workflow 2025-08-03 18:38:00 +10:00
Robert Wojciechowski 1a6bff3ab1 Update pyproject.toml 2025-08-01 15:15:35 +10:00
Robert Wojciechowski d7aec0b0fe Update pyproject.toml 2025-08-01 15:13:11 +10:00
snomiao 26aaa1f1c8 chore(publish): update GitHub Actions workflow for node publishing
- Add permissions for issue writing
- Restrict job execution to specific repository owner
- Update action version for node publishing to v1
2025-07-10 00:44:29 +00:00
19 changed files with 3119 additions and 3624 deletions
+5 -1
View File
@@ -7,14 +7,18 @@ on:
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'robertvoy' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+40 -5
View File
@@ -1,5 +1,5 @@
<div align="center">
<img width="320" src="https://github.com/user-attachments/assets/533bb98d-0c4a-499f-9bca-5c937e361087" />
<img width="250" src="https://github.com/user-attachments/assets/533bb98d-0c4a-499f-9bca-5c937e361087" />
<br><br>
<a href="https://www.youtube.com/watch?v=p6eE3IlAbOs"><img src="https://img.shields.io/badge/Video_Tutorial-grey?style=flat&logo=youtube&logoColor=white" alt="Video Tutorial"></a>
<a href="/docs/worker-setup-guides.md"><img src="https://img.shields.io/badge/Setup_Guides-grey?style=flat&logo=gitbook&logoColor=white" alt="Setup Guides"></a>
@@ -72,6 +72,14 @@ ComfyUI Distributed supports three types of workers:
---
## Official Sponsor
[<img width="1500" height="339" src="https://github.com/user-attachments/assets/c5f75e1f-3e19-4c57-b05d-151311cd1cf0" />](https://get.runpod.io/0bw29uf3ug0p)
Join Runpod with [this link](https://get.runpod.io/0bw29uf3ug0p) and unlock a special bonus.
---
## Workflow Examples
### Basic Parallel Generation
@@ -103,7 +111,7 @@ 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 Upscaling
### Distributed Image Upscaling
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
![Clipboard Image (3)](https://github.com/user-attachments/assets/ffb57a0d-7b75-4497-96d2-875d60865a1a)
@@ -114,9 +122,24 @@ Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, w
2. Upscale with ESRGAN or similar
3. Connect to **Ultimate SD Upscale Distributed**
4. Configure tile settings
> If your GPUs are similar, set `static_distribution` to true; otherwise, false
5. Enable workers for faster processing
### Distributed Video Upscaling
Accelerate Ultimate SD Upscaler by distributing video tiles across multiple workers, with speed scaling as you add more GPUs.
![Video Upscaler workflow](https://github.com/user-attachments/assets/3c3d61b1-0b5f-422e-8c58-7c1555fed765)
> [Download workflow](/workflows/distributed-upscale-video.json)
1. Load your video
2. Optional: upscale with ESRGAN or similar
3. Connect to **Ultimate SD Upscale Distributed**
4. Configure tile settings
5. Use RES4LYF (bong/res2) to get better results
6. Enable workers for faster processing
> You can run this workflow entirely on Runpod with minimal setup. [Check out the guide here.](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/video-upscaler-runpod-preset.md)
---
## FAQ
@@ -161,6 +184,18 @@ This software is provided "as is" without any warranties, express or implied, in
## Support the Project
<img width="200" src="https://github.com/user-attachments/assets/84291921-c44e-4556-94f2-a3b16500f4f9" />
<img width="200" align="right" src="https://github.com/user-attachments/assets/84291921-c44e-4556-94f2-a3b16500f4f9" />
If my custom nodes have added value to your workflow, consider fueling future development with a coffee!
Your support helps keep this project thriving.
Buy me a coffee at: https://buymeacoffee.com/robertvoy
If my custom nodes have added value to your workflow, consider fueling future development with a coffee! Your support helps keep this project thriving. Buy me a coffee at: https://buymeacoffee.com/robertvoy
+1 -1
View File
@@ -62,4 +62,4 @@ __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())}")
debug_log(f"Available nodes: {list(NODE_CLASS_MAPPINGS.keys())}")
+160 -105
View File
@@ -9,6 +9,7 @@ import aiohttp
from aiohttp import web
import io
import server
import comfy.model_management
import subprocess
import platform
import time
@@ -22,7 +23,7 @@ from comfy.utils import ProgressBar
# Import shared utilities
from .utils.logging import debug_log, log
from .utils.config import CONFIG_FILE, get_default_config, load_config, save_config, ensure_config_exists
from .utils.config import CONFIG_FILE, get_default_config, load_config, save_config, ensure_config_exists, get_worker_timeout_seconds
from .utils.image import tensor_to_pil, pil_to_tensor, ensure_contiguous
from .utils.process import is_process_alive, terminate_process, get_python_executable
from .utils.network import handle_api_error, get_server_port, get_server_loop, get_client_session, cleanup_client_session
@@ -1912,8 +1913,10 @@ class DistributedCollectorNode:
collected_count = 0
workers_done = set()
# Use a reasonable timeout for the first image
timeout = WORKER_JOB_TIMEOUT
# Use unified worker timeout from config/UI with simple sliced waits
base_timeout = float(get_worker_timeout_seconds())
slice_timeout = min(0.5, base_timeout) # small per-wait slice to recheck interrupt
last_activity = time.time()
# Get queue size before starting
@@ -1924,116 +1927,168 @@ class DistributedCollectorNode:
# NEW: Initialize progress bar for workers (total = num_workers)
p = ProgressBar(num_workers)
while len(workers_done) < num_workers:
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=timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
# Check if batch mode
tensors = result.get('tensors', [])
indices = result.get('indices', []) # Get the indices
if tensors:
# Batch mode
debug_log(f"Master - Got batch from worker {worker_id}, size={len(tensors)}, is_last={is_last}")
try:
while len(workers_done) < num_workers:
# Check for user interruption to abort collection promptly
comfy.model_management.throw_exception_if_processing_interrupted()
try:
# Get the queue again each time to ensure we have the right reference
async with prompt_server.distributed_jobs_lock:
q = prompt_server.distributed_pending_jobs[multi_job_id]
current_size = q.qsize()
if worker_id not in worker_images:
worker_images[worker_id] = {}
result = await asyncio.wait_for(q.get(), timeout=slice_timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
# Use actual indices if available, otherwise fall back to sequential
if indices:
for i, tensor in enumerate(tensors):
actual_idx = indices[i]
worker_images[worker_id][actual_idx] = tensor
# Check if batch mode
tensors = result.get('tensors', [])
indices = result.get('indices', []) # Get the indices
if tensors:
# Batch mode
debug_log(f"Master - Got batch from worker {worker_id}, size={len(tensors)}, is_last={is_last}")
if worker_id not in worker_images:
worker_images[worker_id] = {}
# Use actual indices if available, otherwise fall back to sequential
if indices:
for i, tensor in enumerate(tensors):
actual_idx = indices[i]
worker_images[worker_id][actual_idx] = tensor
else:
# Fallback for backward compatibility
for idx, tensor in enumerate(tensors):
worker_images[worker_id][idx] = tensor
collected_count += len(tensors)
else:
# Fallback for backward compatibility
for idx, tensor in enumerate(tensors):
worker_images[worker_id][idx] = tensor
# Single image mode (backward compat)
image_index = result['image_index']
tensor = result['tensor']
collected_count += len(tensors)
else:
# Single image mode (backward compat)
image_index = result['image_index']
tensor = result['tensor']
debug_log(f"Master - Got single result from worker {worker_id}, image {image_index}, is_last={is_last}")
if worker_id not in worker_images:
worker_images[worker_id] = {}
worker_images[worker_id][image_index] = tensor
collected_count += 1
# Once we start receiving images, use shorter timeout
timeout = WORKER_JOB_TIMEOUT
if is_last:
workers_done.add(worker_id)
p.update(1) # +1 per completed worker
else:
pass # Continue waiting for more results
except asyncio.TimeoutError:
missing_workers = set(str(w) for w in enabled_workers) - workers_done
log(f"Master - Timeout. Still waiting for workers: {list(missing_workers)}")
# 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()
debug_log(f"Master - Got single result from worker {worker_id}, image {image_index}, is_last={is_last}")
# Try to drain any remaining items
remaining_items = []
while not final_q.empty():
if worker_id not in worker_images:
worker_images[worker_id] = {}
worker_images[worker_id][image_index] = tensor
collected_count += 1
# Record activity and refresh timeout baseline
last_activity = time.time()
base_timeout = float(get_worker_timeout_seconds())
if is_last:
workers_done.add(worker_id)
p.update(1) # +1 per completed worker
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
log(f"Master - Timeout. Still waiting for workers: {list(missing_workers)}")
# Probe missing workers' /prompt endpoints to check if they are actively processing
any_busy = False
try:
cfg = load_config()
cfg_workers = cfg.get('workers', [])
session = await get_client_session()
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
host = wrec.get('host') or 'localhost'
port = int(wrec.get('port', 8188))
url = f"http://{host}:{port}/prompt"
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)
# Check if batch mode
tensors = item.get('tensors', [])
if tensors:
# Batch mode
if worker_id not in worker_images:
worker_images[worker_id] = {}
async with session.get(url, timeout=aiohttp.ClientTimeout(total=2.0)) as resp:
status = resp.status
q = None
if status == 200:
try:
payload = await resp.json()
q = int(payload.get('exec_info', {}).get('queue_remaining', 0))
except Exception:
q = 0
debug_log(f"Collector probe: worker {wid} status={status} queue_remaining={q}")
if status == 200 and q and q > 0:
any_busy = True
log(f"Master - Probe grace: worker {wid} appears busy (queue_remaining={q}). 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)
for idx, tensor in enumerate(tensors):
worker_images[worker_id][idx] = tensor
# Check if batch mode
tensors = item.get('tensors', [])
if tensors:
# Batch mode
if worker_id not in worker_images:
worker_images[worker_id] = {}
for idx, tensor in enumerate(tensors):
worker_images[worker_id][idx] = tensor
collected_count += len(tensors)
else:
# Single image mode
image_index = item['image_index']
tensor = item['tensor']
if worker_id not in worker_images:
worker_images[worker_id] = {}
worker_images[worker_id][image_index] = tensor
collected_count += 1
collected_count += len(tensors)
else:
# Single image mode
image_index = item['image_index']
tensor = item['tensor']
if worker_id not in worker_images:
worker_images[worker_id] = {}
worker_images[worker_id][image_index] = tensor
collected_count += 1
if is_last:
workers_done.add(worker_id)
p.update(1) # +1 here too
else:
log(f"Master - Queue {multi_job_id} no longer exists!")
break
if is_last:
workers_done.add(worker_id)
p.update(1) # +1 here too
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())
+624 -878
View File
File diff suppressed because it is too large Load Diff
+19
View File
@@ -0,0 +1,19 @@
![Clipboard Image](https://github.com/user-attachments/assets/5dc5224f-3f47-442c-b94a-116afeb28132)
**Accelerated Creative Video Upscaler On Runpod:**
1. Use the [ComfyUI Distributed Pod](https://console.runpod.io/deploy?template=m21ynvo8yo&ref=0bw29uf3ug0p) template.
2. Filter instances by CUDA 12.8 (add filter in Additional Filters at the top of the page).
3. Choose 4x 5090s
4. Press Edit Template to configure the pod's Environment Variables:
- CIVITAI_API_TOKEN: Not necessary for this workflow.
- HF_API_TOKEN: [get your token here](https://huggingface.co/settings/tokens)
- SAGE_ATTENTION: optional optimisation (set to true/false). Recommended for this workflow.
- PRESET_VIDEO_UPSCALER: set to true. This will download everything you need.
5. Deploy your pod.
6. Once pod setup is complete, connect to ComfyUI running on your pod.
7. In ComfyUI, open the GPU panel on the left.
> If you set SAGE_ATTENTION to true, add "--use-sage-attention" to Extra Args on the workers.
8. Launch the workers.
9. Upload video, add prompt and run workflow.
10. Right-click the Video Combine node and click Save Preview to save the video.
+1 -1
View File
@@ -75,7 +75,7 @@
1. Register a [Runpod](https://get.runpod.io/0bw29uf3ug0p) account.
2. On Runpod, go to Storage > New Network Volume and create a volume that will store the models you need. Start with 40 GB, you can always add more later. Learn more [about Network Volumes](https://docs.runpod.io/pods/storage/create-network-volumes).
3. Use the [ComfyUI Distributed Pod](https://console.runpod.io/deploy?template=m21ynvo8yo&ref=ak218p52) template.
3. Use the [ComfyUI Distributed Pod](https://console.runpod.io/deploy?template=m21ynvo8yo&ref=0bw29uf3ug0p) template.
4. Make sure your Network Volume is mounted and choose a suitable GPU.
> ⚠️ To use the ComfyUI Distributed Pod template, you will need to filter instances by CUDA 12.8 (add filter in Additional Filters).
6. Press Edit Template to configure the pod's Environment Variables:
+3 -3
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-Distributed"
description = "ComfyUI extension that enables multi-GPU processing locally and remotely "
version = "1.0.5"
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
version = "1.1.0"
license = {file = "LICENSE"}
dependencies = []
@@ -12,4 +12,4 @@ Repository = "https://github.com/robertvoy/ComfyUI-Distributed"
[tool.comfy]
PublisherId = "robertvoy"
DisplayName = "ComfyUI-Distributed"
Icon = ""
Icon = "https://raw.githubusercontent.com/robertvoy/ComfyUI-Distributed/refs/heads/main/web/distributed-logo-icon.png"
+22 -1
View File
@@ -5,6 +5,9 @@ import os
import json
from .logging import log
# Import defaults for timeout fallbacks
from .constants import HEARTBEAT_TIMEOUT
CONFIG_FILE = os.path.join(os.path.dirname(os.path.dirname(__file__)), "gpu_config.json")
def get_default_config():
@@ -47,4 +50,22 @@ def ensure_config_exists():
from .logging import debug_log
debug_log("Created default config file")
else:
log("Could not create default config file")
log("Could not create default config file")
def get_worker_timeout_seconds(default: int = HEARTBEAT_TIMEOUT) -> int:
"""Return the unified worker timeout (seconds).
Priority:
1) UI-configured setting `settings.worker_timeout_seconds`
2) Fallback to provided `default` (defaults to HEARTBEAT_TIMEOUT which itself
can be overridden via the COMFYUI_HEARTBEAT_TIMEOUT env var)
This value should be used anywhere we consider a worker "timed out" from the
master's perspective (e.g., collector waits, upscaler result collection).
"""
try:
cfg = load_config()
val = int(cfg.get('settings', {}).get('worker_timeout_seconds', default))
return max(1, val)
except Exception:
return max(1, int(default))
+1 -1
View File
@@ -38,4 +38,4 @@ MEMORY_CLEAR_DELAY = 0.5
MAX_BATCH = int(os.environ.get('COMFYUI_MAX_BATCH', '20')) # Maximum items per batch to prevent timeouts/OOM (~100MB chunks for 512x512 PNGs)
# Heartbeat monitoring
HEARTBEAT_TIMEOUT = int(os.environ.get('COMFYUI_HEARTBEAT_TIMEOUT', '60')) # Worker heartbeat timeout in seconds
HEARTBEAT_TIMEOUT = int(os.environ.get('COMFYUI_HEARTBEAT_TIMEOUT', '60')) # Worker heartbeat timeout in seconds (default 60s)
+277 -411
View File
@@ -4,22 +4,99 @@ import json
import copy
import os
import io
from aiohttp import web
from aiohttp import web, ClientTimeout
import server
from PIL import Image
import numpy as np
import torch
# Import from other utilities
from .logging import debug_log, log
from .network import handle_api_error, get_client_session
from .image import tensor_to_pil
# We avoid converting to tensors on the master for tiles; blending uses PIL
# 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)))
# Import HEARTBEAT_TIMEOUT from constants
from .constants import HEARTBEAT_TIMEOUT
from .config import load_config
def _parse_tiles_from_form(data):
"""Parse tiles submitted via multipart/form-data into a list of tile dicts.
Expects the following fields in the aiohttp form data:
- 'tiles_metadata': JSON list with per-tile metadata items containing at least
'tile_idx', 'x', 'y', 'extracted_width', 'extracted_height'. Optional
'batch_idx' and 'global_idx' are included when available.
- 'tile_{i}': image bytes for each tile described in tiles_metadata (PNG).
- 'padding': integer padding used during extraction (optional; defaults 0).
Returns: list of dicts with keys: 'image', 'tile_idx', 'x', 'y',
'extracted_width', 'extracted_height', and optional 'batch_idx', 'global_idx',
plus 'padding'.
"""
try:
# Parse padding if present
padding = int(data.get('padding', 0)) if data.get('padding') is not None else 0
except Exception:
padding = 0
# Parse tiles metadata (JSON list)
meta_raw = data.get('tiles_metadata')
if meta_raw is None:
raise ValueError("Missing tiles_metadata")
try:
metadata = json.loads(meta_raw)
except Exception as e:
raise ValueError(f"Invalid tiles_metadata JSON: {e}")
if not isinstance(metadata, list):
raise ValueError("tiles_metadata must be a list")
tiles = []
# Iterate over metadata items and corresponding uploaded files tile_0, tile_1, ...
for i, meta in enumerate(metadata):
file_field = data.get(f'tile_{i}')
if file_field is None or not hasattr(file_field, 'file'):
raise ValueError(f"Missing tile data for index {i}")
# Read image bytes and decode to PIL
raw = file_field.file.read()
try:
img = Image.open(io.BytesIO(raw)).convert("RGB")
except Exception as e:
raise ValueError(f"Invalid image data for tile {i}: {e}")
# Build tile dictionary (store PIL only; master blends via PIL)
try:
tile_info = {
'image': img,
'tile_idx': int(meta.get('tile_idx', i)),
'x': int(meta.get('x', 0)),
'y': int(meta.get('y', 0)),
'extracted_width': int(meta.get('extracted_width', img.width)),
'extracted_height': int(meta.get('extracted_height', img.height)),
'padding': int(padding),
}
except Exception as e:
raise ValueError(f"Invalid metadata values for tile {i}: {e}")
# Optional fields
if 'batch_idx' in meta:
try:
tile_info['batch_idx'] = int(meta['batch_idx'])
except Exception:
pass
if 'global_idx' in meta:
try:
tile_info['global_idx'] = int(meta['global_idx'])
except Exception:
pass
tiles.append(tile_info)
return tiles
# Unified Job Data Structure Keys
@@ -36,7 +113,46 @@ JOB_NUM_TILES_PER_IMAGE = 'num_tiles_per_image' # For static
TASK_TYPE_TILE = 'tile'
TASK_TYPE_IMAGE = 'image'
async def _init_job_queue(multi_job_id, mode, batch_size=None, num_tiles_per_image=None, all_indices=None, enabled_workers=None, task_assignments=None):
from typing import List, Optional
async def init_dynamic_job(multi_job_id: str, batch_size: int, enabled_workers: List[str], all_indices: Optional[List[int]] = None):
"""Initialize queue for dynamic mode (per-image), with collector fields.
- Creates JOB_PENDING_TASKS with image indices
- Adds 'completed_images' dict and 'pending_images' alias used by collectors
"""
await _init_job_queue(
multi_job_id,
'dynamic',
batch_size=batch_size,
all_indices=all_indices or list(range(batch_size)),
enabled_workers=enabled_workers,
)
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
job_data['completed_images'] = {}
job_data['pending_images'] = job_data[JOB_PENDING_TASKS]
debug_log(f"Job {multi_job_id} initialized with {batch_size} images")
async def init_static_job_batched(multi_job_id: str, batch_size: int, num_tiles_per_image: int, enabled_workers: List[str]):
"""Initialize queue for static mode (batched-per-tile).
- Populates JOB_PENDING_TASKS with tile ids [0..num_tiles_per_image-1]
"""
await _init_job_queue(
multi_job_id,
'static',
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
enabled_workers=enabled_workers,
batched_static=True,
)
# Initialization handled by master; avoid duplicate init logs here
async def _init_job_queue(multi_job_id, mode, batch_size=None, num_tiles_per_image=None, all_indices=None, enabled_workers=None, task_assignments=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:
@@ -56,17 +172,22 @@ async def _init_job_queue(multi_job_id, mode, batch_size=None, num_tiles_per_ima
if mode == 'dynamic':
job_data[JOB_BATCH_SIZE] = batch_size
pending_queue = job_data[JOB_PENDING_TASKS]
for i in all_indices or range(batch_size):
for i in (all_indices or range(batch_size)):
await pending_queue.put(i)
debug_log(f"Initialized dynamic queue with {batch_size} pending images")
debug_log(f"Initialized image queue with {batch_size} pending items")
elif mode == 'static':
job_data[JOB_NUM_TILES_PER_IMAGE] = num_tiles_per_image
# For static with dynamic distribution, populate all tile indices in pending queue
job_data[JOB_BATCH_SIZE] = batch_size
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[JOB_PENDING_TASKS]
total_tiles = batch_size * num_tiles_per_image
for i in range(total_tiles):
await pending_queue.put(i)
debug_log(f"Initialized static queue with {total_tiles} pending tiles for dynamic distribution")
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 = batch_size * num_tiles_per_image
for i in range(total_tiles):
await pending_queue.put(i)
# Keep backward compatibility - if task assignments provided, still track them
if task_assignments and enabled_workers:
@@ -80,37 +201,13 @@ async def _init_job_queue(multi_job_id, mode, batch_size=None, num_tiles_per_ima
prompt_server.distributed_pending_tile_jobs[multi_job_id] = job_data
def _distribute_tasks(items: list, num_participants: int) -> list[list[any]]:
"""Distribute a list of items among N participants (master + workers)."""
if num_participants <= 1:
return [items]
items_per_participant = len(items) // num_participants
remainder = len(items) % num_participants
assignments = []
start_idx = 0
for i in range(num_participants):
count = items_per_participant + (1 if i < remainder else 0)
end_idx = start_idx + count
assignments.append(items[start_idx:end_idx])
start_idx = end_idx
return assignments
async def _get_next_task(multi_job_id):
"""Get next task from pending queue (generalized for tiles/images)."""
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 not job_data or JOB_PENDING_TASKS not in job_data:
return None
try:
task_id = await asyncio.wait_for(job_data[JOB_PENDING_TASKS].get(), timeout=1.0)
return task_id
except asyncio.TimeoutError:
return None
# Note: legacy task distribution and queue pull helpers removed
async def _drain_results_queue(multi_job_id):
"""Drain pending results from queue and update completed_tasks. Returns count drained."""
"""Drain pending results from queue and update completed_tasks. Returns count drained.
Uses non-blocking get_nowait to avoid await timeouts and reduce latency.
"""
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)
@@ -120,46 +217,47 @@ async def _drain_results_queue(multi_job_id):
completed_tasks = job_data[JOB_COMPLETED_TASKS]
collected = 0
while not q.empty():
while True:
try:
result = await asyncio.wait_for(q.get(), timeout=0.1)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
if 'image_idx' in result and 'image' in result:
task_id = result['image_idx']
if task_id not in completed_tasks:
completed_tasks[task_id] = result['image']
collected += 1
elif 'tiles' in result:
for tile_data in result['tiles']:
task_id = tile_data.get('global_idx', tile_data['tile_idx'])
if task_id not in completed_tasks:
completed_tasks[task_id] = tile_data
collected += 1
elif 'tensor' in result and 'tile_idx' in result: # Single tile backward compat
task_id = result.get('global_idx', result['tile_idx'])
if task_id not in completed_tasks:
completed_tasks[task_id] = {
'tensor': result['tensor'],
'tile_idx': result['tile_idx'],
'x': result['x'],
'y': result['y'],
'extracted_width': result['extracted_width'],
'extracted_height': result['extracted_height'],
'padding': result['padding'],
'batch_idx': result.get('batch_idx', 0),
'global_idx': task_id
}
collected += 1
if is_last:
# Track worker completion
if worker_id in job_data[JOB_WORKER_STATUS]:
del job_data[JOB_WORKER_STATUS][worker_id]
except asyncio.TimeoutError:
result = q.get_nowait()
except asyncio.QueueEmpty:
break
worker_id = result['worker_id']
is_last = result.get('is_last', False)
if 'image_idx' in result and 'image' in result:
task_id = result['image_idx']
if task_id not in completed_tasks:
completed_tasks[task_id] = result['image']
collected += 1
elif 'tiles' in result:
for tile_data in result['tiles']:
task_id = tile_data.get('global_idx', tile_data['tile_idx'])
if task_id not in completed_tasks:
completed_tasks[task_id] = tile_data
collected += 1
elif 'tensor' in result and 'tile_idx' in result: # Single tile backward compat
task_id = result.get('global_idx', result['tile_idx'])
if task_id not in completed_tasks:
completed_tasks[task_id] = {
'tensor': result['tensor'],
'tile_idx': result['tile_idx'],
'x': result['x'],
'y': result['y'],
'extracted_width': result['extracted_width'],
'extracted_height': result['extracted_height'],
'padding': result['padding'],
'batch_idx': result.get('batch_idx', 0),
'global_idx': task_id
}
collected += 1
if is_last:
# Track worker completion
if worker_id in job_data[JOB_WORKER_STATUS]:
del job_data[JOB_WORKER_STATUS][worker_id]
return collected
async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
@@ -174,13 +272,92 @@ async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
requeued_count = 0
completed_tasks = job_data.get(JOB_COMPLETED_TASKS, {})
# Allow override via config setting 'worker_timeout_seconds'
cfg = load_config()
hb_timeout = int(cfg.get('settings', {}).get('worker_timeout_seconds', HEARTBEAT_TIMEOUT))
for worker, last_heartbeat in list(job_data.get(JOB_WORKER_STATUS, {}).items()):
if current_time - last_heartbeat > HEARTBEAT_TIMEOUT:
age = current_time - last_heartbeat
debug_log(f"Timeout check: worker={worker} age={age:.1f}s threshold={hb_timeout}s")
if age > hb_timeout:
# Busy-only grace policy: require positive signal from worker (/prompt)
# We also log assignment state for diagnostics but do not grace on it alone.
assigned = job_data.get(JOB_ASSIGNED_TO_WORKERS, {}).get(worker, [])
incomplete_assigned = 0
try:
if assigned:
batched_static = bool(job_data.get('batched_static', False))
if batched_static:
num_tiles_per_image = job_data.get(JOB_NUM_TILES_PER_IMAGE, 1)
batch_size = job_data.get(JOB_BATCH_SIZE, 1)
for task_id in assigned:
for b in range(batch_size):
gidx = b * num_tiles_per_image + task_id
if gidx not in completed_tasks:
incomplete_assigned += 1
break
else:
for task_id in assigned:
if task_id not in completed_tasks:
incomplete_assigned += 1
debug_log(f"Assigned diagnostics: total_assigned={len(assigned)} incomplete_assigned={incomplete_assigned}")
except Exception as e:
debug_log(f"Assigned diagnostics failed for worker {worker}: {e}")
busy = False
probe_status = None
probe_queue = None
try:
cfg_workers = load_config().get('workers', [])
wrec = next((w for w in cfg_workers if str(w.get('id')) == str(worker)), None)
if wrec:
host = wrec.get('host') or 'localhost'
port = int(wrec.get('port', 8188))
url = f"http://{host}:{port}/prompt"
debug_log(f"Probing worker {worker} at {url}")
session = await get_client_session()
async with session.get(url, timeout=ClientTimeout(total=2.0)) as resp:
probe_status = resp.status
if resp.status == 200:
try:
payload = await resp.json()
probe_queue = int(payload.get('exec_info', {}).get('queue_remaining', 0))
except Exception:
probe_queue = 0
busy = probe_queue is not None and probe_queue > 0
except Exception as e:
debug_log(f"Probe failed for worker {worker}: {e}")
finally:
debug_log(f"Probe diagnostics: http_status={probe_status} queue_remaining={probe_queue}")
if busy:
job_data[JOB_WORKER_STATUS][worker] = current_time
debug_log(f"Heartbeat grace: worker {worker} busy via probe; skipping requeue")
continue
log(f"Worker {worker} timed out")
for task_id in job_data.get(JOB_ASSIGNED_TO_WORKERS, {}).get(worker, []):
if task_id not in completed_tasks:
await job_data[JOB_PENDING_TASKS].put(task_id)
requeued_count += 1
# If batched_static, task_id is a tile_idx; consider it complete only if
# all corresponding global_idx entries are present in completed_tasks.
batched_static = bool(job_data.get('batched_static', False))
if batched_static:
num_tiles_per_image = job_data.get(JOB_NUM_TILES_PER_IMAGE, 1)
batch_size = job_data.get(JOB_BATCH_SIZE, 1)
# Check all global indices for this tile across the batch
all_done = True
for b in range(batch_size):
gidx = b * num_tiles_per_image + task_id
if gidx not in completed_tasks:
all_done = False
break
if not all_done:
await job_data[JOB_PENDING_TASKS].put(task_id)
requeued_count += 1
else:
# Legacy/global-idx mode: task_id is a global index key
if task_id not in completed_tasks:
await job_data[JOB_PENDING_TASKS].put(task_id)
requeued_count += 1
if JOB_WORKER_STATUS in job_data:
del job_data[JOB_WORKER_STATUS][worker]
if JOB_ASSIGNED_TO_WORKERS in job_data:
@@ -226,42 +403,6 @@ async def _cleanup_job(multi_job_id):
# API Endpoints (generalized)
@server.PromptServer.instance.routes.post("/distributed/request_task")
async def request_task_endpoint(request):
try:
data = await request.json()
worker_id = data.get('worker_id')
multi_job_id = data.get('multi_job_id')
if not worker_id or not multi_job_id:
return await handle_api_error(request, "Missing worker_id or multi_job_id", 400)
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:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
mode = job_data.get(JOB_MODE)
pending_queue = job_data.get(JOB_PENDING_TASKS)
if pending_queue:
try:
task_id = await asyncio.wait_for(pending_queue.get(), timeout=0.1)
if JOB_ASSIGNED_TO_WORKERS in job_data and worker_id in job_data[JOB_ASSIGNED_TO_WORKERS]:
job_data[JOB_ASSIGNED_TO_WORKERS][worker_id].append(task_id)
if JOB_WORKER_STATUS in job_data:
job_data[JOB_WORKER_STATUS][worker_id] = time.time()
remaining = pending_queue.qsize()
debug_log(f"Assigned task {task_id} to worker {worker_id} in {mode} mode")
return web.json_response({"task_id": task_id, "estimated_remaining": remaining, "mode": mode})
except asyncio.TimeoutError:
return web.json_response({"task_id": None})
else:
return await handle_api_error(request, "No pending tasks", 400)
else:
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.post("/distributed/heartbeat")
async def heartbeat_endpoint(request):
try:
@@ -314,7 +455,7 @@ async def submit_tiles_endpoint(request):
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'static':
return await handle_api_error(request, "Mode mismatch: expected static mode", 400)
return await handle_api_error(request, "Job not configured for tile submissions", 400)
if JOB_QUEUE in job_data:
await job_data[JOB_QUEUE].put({
'worker_id': worker_id,
@@ -324,115 +465,17 @@ async def submit_tiles_endpoint(request):
debug_log(f"Received completion signal from worker {worker_id}")
return web.json_response({"status": "success"})
# Handle batch tiles with metadata
if batch_size > 0:
padding = int(data.get('padding', 32))
metadata_field = data.get('tiles_metadata')
if metadata_field:
if hasattr(metadata_field, 'file'):
metadata_str = metadata_field.file.read().decode('utf-8')
elif isinstance(metadata_field, (bytes, bytearray)):
metadata_str = metadata_field.decode('utf-8')
else:
metadata_str = str(metadata_field)
metadata = json.loads(metadata_str)
if len(metadata) != batch_size:
return await handle_api_error(request, "Metadata length mismatch", 400)
tile_data_list = []
for i in range(batch_size):
tile_field = data.get(f'tile_{i}')
if tile_field is None:
continue
img_data = tile_field.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
if i < len(metadata):
tile_meta = metadata[i]
tile_idx = tile_meta.get('tile_idx', i)
tile_info = {
'tensor': tensor,
'tile_idx': tile_idx,
'x': tile_meta['x'],
'y': tile_meta['y'],
'extracted_width': tile_meta['extracted_width'],
'extracted_height': tile_meta['extracted_height'],
'padding': padding,
'batch_idx': tile_meta.get('batch_idx', 0),
'global_idx': tile_meta.get('global_idx', tile_idx)
}
tile_data_list.append(tile_info)
tile_data_list.sort(key=lambda x: x['tile_idx'])
tiles.extend(tile_data_list)
else:
# Legacy format
for i in range(batch_size):
tile_field = data.get(f'tile_{i}')
if tile_field is None:
continue
tile_idx = int(data.get(f'tile_{i}_idx', i))
x = int(data.get(f'tile_{i}_x', 0))
y = int(data.get(f'tile_{i}_y', 0))
extracted_width = int(data.get(f'tile_{i}_width', 512))
extracted_height = int(data.get(f'tile_{i}_height', 512))
img_data = tile_field.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
tiles.append({
'tensor': tensor,
'tile_idx': tile_idx,
'x': x,
'y': y,
'extracted_width': extracted_width,
'extracted_height': extracted_height,
'padding': padding
})
else:
# Single tile legacy
image_file = data.get('image')
if not image_file:
return await handle_api_error(request, "Missing image data", 400)
tile_idx = int(data.get('tile_idx', 0))
x = int(data.get('x', 0))
y = int(data.get('y', 0))
extracted_width = int(data.get('extracted_width', 512))
extracted_height = int(data.get('extracted_height', 512))
padding = int(data.get('padding', 32))
img_data = image_file.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
tiles = [{
'tensor': tensor,
'tile_idx': tile_idx,
'x': x,
'y': y,
'extracted_width': extracted_width,
'extracted_height': extracted_height,
'padding': padding,
'batch_idx': 0,
'global_idx': tile_idx
}]
try:
tiles = _parse_tiles_from_form(data)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
# Submit tiles to queue
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'static':
return await handle_api_error(request, "Mode mismatch: expected static mode", 400)
return await handle_api_error(request, "Job not configured for tile submissions", 400)
q = job_data[JOB_QUEUE]
if batch_size > 0 or len(tiles) > 0:
@@ -441,6 +484,7 @@ async def submit_tiles_endpoint(request):
'tiles': tiles,
'is_last': is_last
})
debug_log(f"Received {len(tiles)} tiles from worker {worker_id} (is_last={is_last})")
else:
await q.put({
'worker_id': worker_id,
@@ -484,7 +528,7 @@ async def submit_image_endpoint(request):
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'dynamic':
return await handle_api_error(request, "Mode mismatch: expected dynamic mode", 400)
return await handle_api_error(request, "Job not configured for image submissions", 400)
if JOB_QUEUE in job_data:
await job_data[JOB_QUEUE].put({
'worker_id': worker_id,
@@ -500,7 +544,7 @@ async def submit_image_endpoint(request):
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'dynamic':
return await handle_api_error(request, "Mode mismatch: expected dynamic mode", 400)
return await handle_api_error(request, "Job not configured for image submissions", 400)
if JOB_QUEUE in job_data:
await job_data[JOB_QUEUE].put({
'worker_id': worker_id,
@@ -516,188 +560,7 @@ async def submit_image_endpoint(request):
except Exception as e:
return await handle_api_error(request, e, 500)
# Keep legacy endpoint for backward compatibility
@server.PromptServer.instance.routes.post("/distributed/tile_complete")
async def tile_complete_endpoint(request):
try:
content_length = request.headers.get('content-length')
if content_length and int(content_length) > MAX_PAYLOAD_SIZE:
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
data = await request.post()
multi_job_id = data.get('multi_job_id')
worker_id = data.get('worker_id')
is_last = data.get('is_last', 'False').lower() == 'true'
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()
if 'full_image' in data and 'image_idx' in data:
image_idx = int(data.get('image_idx'))
img_data = data['full_image'].file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
debug_log(f"Received full image {image_idx} from worker {worker_id}")
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'dynamic':
return await handle_api_error(request, "Mode mismatch for image submission", 400)
if JOB_QUEUE in job_data:
await job_data[JOB_QUEUE].put({
'worker_id': worker_id,
'image_idx': image_idx,
'image': img,
'is_last': is_last
})
return web.json_response({"status": "success"})
batch_size = int(data.get('batch_size', 0))
tiles = []
if batch_size == 0 and is_last:
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_QUEUE in job_data:
await job_data[JOB_QUEUE].put({
'worker_id': worker_id,
'is_last': True,
'tiles': []
})
debug_log(f"Received completion signal from worker {worker_id}")
return web.json_response({"status": "success"})
if batch_size > 0:
padding = int(data.get('padding', 32))
metadata_field = data.get('tiles_metadata')
if metadata_field:
if hasattr(metadata_field, 'file'):
metadata_str = metadata_field.file.read().decode('utf-8')
elif isinstance(metadata_field, (bytes, bytearray)):
metadata_str = metadata_field.decode('utf-8')
else:
metadata_str = str(metadata_field)
metadata = json.loads(metadata_str)
if len(metadata) != batch_size:
return await handle_api_error(request, "Metadata length mismatch", 400)
tile_data_list = []
for i in range(batch_size):
tile_field = data.get(f'tile_{i}')
if tile_field is None:
continue
img_data = tile_field.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
if i < len(metadata):
tile_meta = metadata[i]
tile_idx = tile_meta.get('tile_idx', i)
tile_info = {
'tensor': tensor,
'tile_idx': tile_idx,
'x': tile_meta['x'],
'y': tile_meta['y'],
'extracted_width': tile_meta['extracted_width'],
'extracted_height': tile_meta['extracted_height'],
'padding': padding,
'batch_idx': tile_meta.get('batch_idx', 0),
'global_idx': tile_meta.get('global_idx', tile_idx)
}
tile_data_list.append(tile_info)
tile_data_list.sort(key=lambda x: x['tile_idx'])
tiles.extend(tile_data_list)
else:
# Legacy
for i in range(batch_size):
tile_field = data.get(f'tile_{i}')
if tile_field is None:
continue
tile_idx = int(data.get(f'tile_{i}_idx', i))
x = int(data.get(f'tile_{i}_x', 0))
y = int(data.get(f'tile_{i}_y', 0))
extracted_width = int(data.get(f'tile_{i}_width', 512))
extracted_height = int(data.get(f'tile_{i}_height', 512))
img_data = tile_field.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
tiles.append({
'tensor': tensor,
'tile_idx': tile_idx,
'x': x,
'y': y,
'extracted_width': extracted_width,
'extracted_height': extracted_height,
'padding': padding
})
else:
# Single tile legacy
image_file = data.get('image')
if not image_file:
return await handle_api_error(request, "Missing image data", 400)
tile_idx = int(data.get('tile_idx', 0))
x = int(data.get('x', 0))
y = int(data.get('y', 0))
extracted_width = int(data.get('extracted_width', 512))
extracted_height = int(data.get('extracted_height', 512))
padding = int(data.get('padding', 32))
img_data = image_file.file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
img_np = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(img_np)[None,]
tiles = [{
'tensor': tensor,
'tile_idx': tile_idx,
'x': x,
'y': y,
'extracted_width': extracted_width,
'extracted_height': extracted_height,
'padding': padding,
'batch_idx': 0,
'global_idx': tile_idx
}]
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if JOB_MODE in job_data and job_data[JOB_MODE] != 'static':
return await handle_api_error(request, "Mode mismatch for tile submission", 400)
q = job_data[JOB_QUEUE]
if batch_size > 0 or len(tiles) > 0:
await q.put({
'worker_id': worker_id,
'tiles': tiles,
'is_last': is_last
})
else:
await q.put({
'worker_id': worker_id,
'is_last': True,
'tiles': []
})
return web.json_response({"status": "success"})
else:
return await handle_api_error(request, "Job not found", 404)
except Exception as e:
return await handle_api_error(request, e, 500)
# Note: Removed legacy /distributed/tile_complete endpoint. Use /distributed/submit_tiles.
@@ -708,7 +571,8 @@ def clone_control_chain(control, clone_hint=True):
return None
new_control = copy.copy(control) # Shallow copy (shares model)
if clone_hint and hasattr(control, 'cond_hint_original'):
new_control.cond_hint_original = control.cond_hint_original.clone()
hint = getattr(control, 'cond_hint_original', None)
new_control.cond_hint_original = hint.clone() if hint is not None else None
if hasattr(control, 'previous_controlnet'):
new_control.previous_controlnet = clone_control_chain(control.previous_controlnet, clone_hint)
return new_control
@@ -722,10 +586,12 @@ def clone_conditioning(cond_list, clone_hints=True):
if 'control' in new_dict:
new_dict['control'] = clone_control_chain(new_dict['control'], clone_hints)
if 'mask' in new_dict:
new_dict['mask'] = new_dict['mask'].clone()
if new_dict['mask'] is not None:
new_dict['mask'] = new_dict['mask'].clone()
# Handle other potential fields if needed
if 'pooled_output' in new_dict:
new_dict['pooled_output'] = new_dict['pooled_output'].clone()
if new_dict['pooled_output'] is not None:
new_dict['pooled_output'] = new_dict['pooled_output'].clone()
if 'area' in new_dict:
new_dict['area'] = new_dict['area'][:] # Shallow copy list/tuple
new_cond.append([new_emb, new_dict])
@@ -776,7 +642,7 @@ async def request_image_endpoint(request):
elif mode == 'static' and JOB_PENDING_TASKS in job_data:
pending_queue = job_data[JOB_PENDING_TASKS]
else:
return await handle_api_error(request, f"Invalid {mode} mode configuration", 400)
return await handle_api_error(request, "Invalid job configuration", 400)
try:
task_idx = await asyncio.wait_for(pending_queue.get(), timeout=0.1)
@@ -795,7 +661,7 @@ async def request_image_endpoint(request):
return web.json_response({"image_idx": task_idx, "estimated_remaining": remaining})
else: # static
debug_log(f"UltimateSDUpscale API - Assigned tile {task_idx} to worker {worker_id}")
return web.json_response({"tile_idx": task_idx, "estimated_remaining": remaining})
return web.json_response({"tile_idx": task_idx, "estimated_remaining": remaining, "batched_static": job_data.get('batched_static', False)})
except asyncio.TimeoutError:
if mode == 'dynamic':
return web.json_response({"image_idx": None})
+56 -1
View File
@@ -4,7 +4,6 @@ import torch
import torch.nn.functional as F
from torchvision.transforms import GaussianBlur
import math
import os
if (not hasattr(Image, 'Resampling')): # For older versions of Pillow
Image.Resampling = Image
@@ -442,6 +441,61 @@ def crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad
cond_dict["mask"] = torch.cat(masks, dim=0) # (B, H, W)
# Added Flux-Kontext Support crop_reference_latents by TBG ETUR
def crop_reference_latents(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad):
"""
1. Resize each latent to `canvas_size` in latent units.
2. Crop the rectangle `region` (pixel coordinates).
3. Down-sample the crop to latent-space `tile_size`.
Expects a list of BCHW tensors under "reference_latents".
"""
latents = cond_dict.get("reference_latents")
if not isinstance(latents, list):
return # nothing to do
k = 8 # down-sample factor from pixel space → latent space (SD-type models)
W_can_px, H_can_px = canvas_size
# canvas size expressed in latent units
W_can_lat, H_can_lat = W_can_px // k, H_can_px // k
W_tile_px, H_tile_px = tile_size
W_tile_lat, H_tile_lat = max(1, W_tile_px // k), max(1, H_tile_px // k)
x1_px, y1_px, x2_px, y2_px = region
new_latents = []
for t in latents: # (B,C,H_lat_in,W_lat_in)
if t.ndim != 4:
raise ValueError(f"expected BCHW, got {t.shape}")
# 1. Resize to canvas resolution in latent units only if needed
if t.shape[-2:] != (H_can_lat, W_can_lat):
t = F.interpolate(t,
size=(H_can_lat, W_can_lat),
mode="bilinear",
align_corners=False)
# 2. Convert pixel crop → latent slice
w0_lat = int(round(x1_px / k))
w1_lat = int(round(x2_px / k))
h0_lat = int(round(y1_px / k))
h1_lat = int(round(y2_px / k))
cropped = t[:, :, h0_lat:h1_lat, w0_lat:w1_lat] # view
# 3. Down-sample to latent-tile size
cropped = F.interpolate(cropped,
size=(H_tile_lat, W_tile_lat),
mode="bilinear",
align_corners=False)
new_latents.append(cropped)
cond_dict["reference_latents"] = new_latents
def crop_cond(cond, region, init_size, canvas_size, tile_size, w_pad=0, h_pad=0):
cropped = []
@@ -452,5 +506,6 @@ def crop_cond(cond, region, init_size, canvas_size, tile_size, w_pad=0, h_pad=0)
crop_gligen(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_area(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_mask(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
crop_reference_latents(cond_dict, region, init_size, canvas_size, tile_size, w_pad, h_pad)
cropped.append(n)
return cropped
Binary file not shown.

After

Width:  |  Height:  |  Size: 4.7 KiB

+18 -6
View File
@@ -19,13 +19,25 @@ function convertPathsForPlatform(apiPrompt, targetSeparator) {
const isLikelyFilename = (value) => {
return value.match(/\.(ckpt|safetensors|pt|pth|bin|yaml|json|png|jpg|jpeg|webp|gif|bmp|latent|txt|vae|lora|embedding)(\s*\[\w+\])?$/i);
};
const isImageOrVideo = (value) => {
return value.match(/\.(png|jpg|jpeg|webp|gif|bmp|mp4|avi|mov|mkv|webm)(\s*\[\w+\])?$/i);
};
function convert(obj) {
if (typeof obj === 'string') {
// Only convert strings that look like file paths
if ((obj.includes('\\') || obj.includes('/')) && isLikelyFilename(obj)) {
// Replace any path separator with the target one
return obj.replace(/[\\\/]/g, targetSeparator);
const trimmed = obj.trim();
const hasDrive = /^[A-Za-z]:\\\\|^[A-Za-z]:\//.test(trimmed);
const isAbsolute = trimmed.startsWith('/') || trimmed.startsWith('\\\\');
const hasProtocol = /^\w+:\/\//.test(trimmed);
// For annotated relative image/video paths, keep forward slashes
if (!hasDrive && !isAbsolute && !hasProtocol && isImageOrVideo(trimmed)) {
return trimmed.replace(/[\\\\]/g, '/');
}
// Otherwise replace any path separator with the worker's target separator
return trimmed.replace(/[\\\\\/]/g, targetSeparator);
}
return obj;
} else if (Array.isArray(obj)) {
@@ -507,9 +519,9 @@ export async function uploadImagesToWorker(extension, workerUrl, images) {
let cleanName = imageData.name;
let subfolder = '';
// Extract subfolder if present
if (cleanName.includes('/')) {
const parts = cleanName.split('/');
// Extract subfolder if present (handle both slash styles)
if (cleanName.includes('/') || cleanName.includes('\\')) {
const parts = cleanName.replace(/\\/g, '/').split('/');
subfolder = parts.slice(0, -1).join('/');
cleanName = parts[parts.length - 1];
}
@@ -587,4 +599,4 @@ export async function performPreflightCheck(extension, workers) {
});
return activeWorkers;
}
}
+15 -19
View File
@@ -154,11 +154,22 @@ class DistributedExtension {
try {
await this.api.updateSetting(key, value);
const prettyKey = key.replace(/_/g, ' ').replace(/\b\w/g, l => l.toUpperCase());
let detail;
if (key === 'worker_timeout_seconds') {
const secs = parseInt(value, 10);
detail = `Worker Timeout set to ${Number.isFinite(secs) ? secs : value}s`;
} else if (typeof value === 'boolean') {
detail = `${prettyKey} ${value ? 'enabled' : 'disabled'}`;
} else {
detail = `${prettyKey} set to ${value}`;
}
app.extensionManager.toast.add({
severity: "success",
summary: "Setting Updated",
detail: `${key.replace(/_/g, ' ').replace(/\b\w/g, l => l.toUpperCase())} ${value ? 'enabled' : 'disabled'}`,
detail,
life: 2000
});
} catch (error) {
@@ -217,22 +228,7 @@ class DistributedExtension {
this.panelElement = null;
}
updateSummary() {
const summaryEl = document.getElementById('distributed-summary');
if (summaryEl) {
// Count active workers from state
const activeWorkers = this.config.workers.filter(w => {
const status = this.state.getWorkerStatus(w.id);
return status.online;
}).length;
const totalGPUs = activeWorkers + 1;
if (this.isEnabled) {
summaryEl.textContent = `If Collector node is present, total generation = (${totalGPUs} GPUs × Batch Size)`;
} else {
summaryEl.textContent = "Only the master GPU will be used.";
}
}
}
// updateSummary removed
// --- Core Logic & Execution ---
@@ -1333,4 +1329,4 @@ app.registerExtension({
async setup() {
new DistributedExtension();
}
});
});
+71 -31
View File
@@ -148,57 +148,67 @@ export async function renderSidebarContent(extension, el) {
// Settings section
const settingsSection = document.createElement("div");
settingsSection.style.cssText = "border-top: 1px solid #444; padding-top: 10px; margin-bottom: 10px;";
// Top separator only; spacing handled by the clickable toggle area for equal top/bottom spacing
settingsSection.style.cssText = "border-top: 1px solid #444; margin-bottom: 10px;";
// Settings header with toggle
// Settings header with toggle (full-area clickable between separators)
const settingsToggleArea = document.createElement("div");
// Equal spacing above header (to top separator) and below header (to bottom separator)
settingsToggleArea.style.cssText = "padding: 16.5px 0; cursor: pointer; user-select: none;";
const settingsHeader = document.createElement("div");
settingsHeader.style.cssText = "display: flex; align-items: center; justify-content: space-between; cursor: pointer; user-select: none;";
settingsHeader.style.cssText = "display: flex; align-items: center; justify-content: space-between;";
const workerSettingsTitle = document.createElement("h4");
workerSettingsTitle.textContent = "Settings";
workerSettingsTitle.style.cssText = "margin: 0; font-size: 14px;";
const workerSettingsToggle = document.createElement("span");
workerSettingsToggle.textContent = "▶"; // Right arrow when collapsed
workerSettingsToggle.style.cssText = "font-size: 12px; color: #888; transition: all 0.2s ease;";
settingsHeader.appendChild(workerSettingsTitle);
settingsHeader.appendChild(workerSettingsToggle);
settingsToggleArea.appendChild(settingsHeader);
// Hover effect for toggle area
settingsToggleArea.onmouseover = () => { workerSettingsToggle.style.color = "#fff"; };
settingsToggleArea.onmouseout = () => { workerSettingsToggle.style.color = "#888"; };
// Hover effect for header
settingsHeader.onmouseover = () => {
workerSettingsToggle.style.color = "#fff";
};
settingsHeader.onmouseout = () => {
workerSettingsToggle.style.color = "#888";
};
// A small separator shown only when collapsed (to make the section boundary obvious)
const settingsSeparator = document.createElement("div");
// No margin so the bottom spacing is controlled by settingsToggleArea padding-bottom
settingsSeparator.style.cssText = "border-bottom: 1px solid #444; margin: 0;";
// Collapsible settings content
const settingsContent = document.createElement("div");
settingsContent.style.cssText = "max-height: 0; overflow: hidden; opacity: 0; transition: max-height 0.3s ease, opacity 0.3s ease;";
const settingsDiv = document.createElement("div");
settingsDiv.style.cssText = "display: flex; flex-direction: column; gap: 8px; padding-top: 10px;";
settingsDiv.style.cssText = "display: grid; grid-template-columns: 1fr auto; row-gap: 10px; column-gap: 10px; padding-top: 10px; align-items: center;";
// Toggle functionality
let settingsExpanded = false;
settingsHeader.onclick = () => {
settingsToggleArea.onclick = () => {
settingsExpanded = !settingsExpanded;
if (settingsExpanded) {
settingsContent.style.maxHeight = "200px";
settingsContent.style.opacity = "1";
workerSettingsToggle.style.transform = "rotate(90deg)";
settingsSeparator.style.display = "none";
} else {
settingsContent.style.maxHeight = "0";
settingsContent.style.opacity = "0";
workerSettingsToggle.style.transform = "rotate(0deg)";
settingsSeparator.style.display = "block";
}
};
// Debug mode setting
// Section: General
const generalLabel = document.createElement("div");
generalLabel.textContent = "GENERAL";
generalLabel.style.cssText = "grid-column: 1 / span 2; font-size: 11px; color: #888; letter-spacing: 0.06em; padding-top: 2px;";
const debugGroup = document.createElement("div");
debugGroup.style.cssText = "display: flex; align-items: center; gap: 8px;";
debugGroup.style.cssText = "grid-column: 1 / span 2; display: flex; align-items: center; gap: 8px;";
const debugCheckbox = document.createElement("input");
debugCheckbox.type = "checkbox";
debugCheckbox.id = "setting-debug";
@@ -209,13 +219,14 @@ export async function renderSidebarContent(extension, el) {
debugLabel.htmlFor = "setting-debug";
debugLabel.textContent = "Debug Mode";
debugLabel.style.cssText = "font-size: 12px; color: #ccc; cursor: pointer;";
debugLabel.title = "Enable verbose logging in the browser console.";
debugGroup.appendChild(debugCheckbox);
debugGroup.appendChild(debugLabel);
// Auto-launch workers setting
const autoLaunchGroup = document.createElement("div");
autoLaunchGroup.style.cssText = "display: flex; align-items: center; gap: 8px;";
autoLaunchGroup.style.cssText = "grid-column: 1 / span 2; display: flex; align-items: center; gap: 8px;";
const autoLaunchCheckbox = document.createElement("input");
autoLaunchCheckbox.type = "checkbox";
@@ -227,13 +238,14 @@ export async function renderSidebarContent(extension, el) {
autoLaunchLabel.htmlFor = "setting-auto-launch";
autoLaunchLabel.textContent = "Auto-launch Local Workers on Startup";
autoLaunchLabel.style.cssText = "font-size: 12px; color: #ccc; cursor: pointer;";
autoLaunchLabel.title = "Start local worker processes automatically when the master starts.";
autoLaunchGroup.appendChild(autoLaunchCheckbox);
autoLaunchGroup.appendChild(autoLaunchLabel);
// Stop workers on exit setting
// Stop workers on exit setting (under General)
const stopOnExitGroup = document.createElement("div");
stopOnExitGroup.style.cssText = "display: flex; align-items: center; gap: 8px;";
stopOnExitGroup.style.cssText = "grid-column: 1 / span 2; display: flex; align-items: center; gap: 8px;";
const stopOnExitCheckbox = document.createElement("input");
stopOnExitCheckbox.type = "checkbox";
@@ -245,28 +257,56 @@ export async function renderSidebarContent(extension, el) {
stopOnExitLabel.htmlFor = "setting-stop-on-exit";
stopOnExitLabel.textContent = "Stop Local Workers on Master Exit";
stopOnExitLabel.style.cssText = "font-size: 12px; color: #ccc; cursor: pointer;";
stopOnExitLabel.title = "Stop local worker processes automatically when the master exits.";
stopOnExitGroup.appendChild(stopOnExitCheckbox);
stopOnExitGroup.appendChild(stopOnExitLabel);
settingsDiv.appendChild(generalLabel);
settingsDiv.appendChild(debugGroup);
settingsDiv.appendChild(autoLaunchGroup);
settingsDiv.appendChild(stopOnExitGroup);
// Worker Timeout setting (seconds)
// Section: Timeouts
const timeoutsLabel = document.createElement("div");
timeoutsLabel.textContent = "TIMEOUTS";
timeoutsLabel.style.cssText = "grid-column: 1 / span 2; font-size: 11px; color: #888; letter-spacing: 0.06em; padding-top: 4px;";
const timeoutGroup = document.createElement("div");
timeoutGroup.style.cssText = "grid-column: 1 / span 2; display: flex; align-items: center; gap: 6px;";
const timeoutLabel = document.createElement("label");
timeoutLabel.htmlFor = "setting-worker-timeout";
timeoutLabel.textContent = "Worker Timeout";
timeoutLabel.style.cssText = "font-size: 12px; color: #ccc; cursor: default;";
timeoutLabel.title = "Seconds without a heartbeat before a worker is considered timed out and its tasks are requeued. Default 60. For WAN, consider 300–600.";
const timeoutInput = document.createElement("input");
timeoutInput.type = "number";
timeoutInput.id = "setting-worker-timeout";
timeoutInput.min = "10";
timeoutInput.step = "1";
timeoutInput.style.cssText = "width: 80px; padding: 2px 6px; background: #222; color: #ddd; border: 1px solid #333; border-radius: 3px;";
timeoutInput.value = (extension.config?.settings?.worker_timeout_seconds ?? 60);
timeoutInput.onchange = (e) => {
const v = parseInt(e.target.value, 10);
if (!Number.isFinite(v) || v <= 0) return;
extension._updateSetting('worker_timeout_seconds', v);
};
timeoutGroup.appendChild(timeoutLabel);
timeoutGroup.appendChild(timeoutInput);
settingsDiv.appendChild(timeoutsLabel);
settingsDiv.appendChild(timeoutGroup);
settingsContent.appendChild(settingsDiv);
settingsSection.appendChild(settingsHeader);
settingsSection.appendChild(settingsToggleArea);
settingsSection.appendChild(settingsSeparator);
settingsSection.appendChild(settingsContent);
container.appendChild(settingsSection);
const summarySection = document.createElement("div");
summarySection.style.cssText = "border-top: 1px solid #444; padding-top: 10px;";
const summary = document.createElement("div");
summary.id = "distributed-summary";
summary.style.cssText = "font-size: 11px; color: #888;";
summarySection.appendChild(summary);
container.appendChild(summarySection);
el.appendChild(container);
extension.updateSummary();
// Start checking worker statuses immediately in parallel
setTimeout(() => extension.checkAllWorkerStatuses(), 0);
@@ -274,4 +314,4 @@ export async function renderSidebarContent(extension, el) {
// Always reset the rendering flag
extension._isRendering = false;
}
}
}
+5 -3
View File
@@ -105,13 +105,15 @@ export function findImageReferences(extension, apiPrompt) {
if (typeof mediaValue === 'string') {
// Clean special suffixes like [input] or [output]
const cleanValue = mediaValue.replace(/\s*\[\w+\]$/, '').trim();
if (imageExtensions.test(cleanValue)) {
images.set(cleanValue, {
// Normalize to forward slashes so subfolder/filename derivation is consistent on Windows
const normalizedValue = cleanValue.replace(/\\/g, '/');
if (imageExtensions.test(normalizedValue)) {
images.set(normalizedValue, {
nodeId,
nodeType: node.class_type,
inputName: 'image' // Keep as 'image' for compatibility
});
extension.log(`Found media reference: ${cleanValue} in node ${nodeId} (${node.class_type})`, "debug");
extension.log(`Found media reference: ${normalizedValue} in node ${nodeId} (${node.class_type})`, "debug");
}
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff