From 75a818c720c1ef0646ad5bf03af740e5d479e882 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Fri, 8 Nov 2024 08:33:44 -0500 Subject: [PATCH 1/6] Move mochi latent node to: latent/video. --- comfy_extras/nodes_mochi.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/comfy_extras/nodes_mochi.py b/comfy_extras/nodes_mochi.py index 4cbbea099..1c474faa9 100644 --- a/comfy_extras/nodes_mochi.py +++ b/comfy_extras/nodes_mochi.py @@ -3,9 +3,6 @@ import torch import comfy.model_management class EmptyMochiLatentVideo: - def __init__(self): - self.device = comfy.model_management.intermediate_device() - @classmethod def INPUT_TYPES(s): return {"required": { "width": ("INT", {"default": 848, "min": 16, "max": nodes.MAX_RESOLUTION, "step": 16}), @@ -15,10 +12,10 @@ class EmptyMochiLatentVideo: RETURN_TYPES = ("LATENT",) FUNCTION = "generate" - CATEGORY = "latent/mochi" + CATEGORY = "latent/video" def generate(self, width, height, length, batch_size=1): - latent = torch.zeros([batch_size, 12, ((length - 1) // 6) + 1, height // 8, width // 8], device=self.device) + latent = torch.zeros([batch_size, 12, ((length - 1) // 6) + 1, height // 8, width // 8], device=comfy.model_management.intermediate_device()) return ({"samples":latent}, ) NODE_CLASS_MAPPINGS = { From dd5b57e3d7371cab793beedfdc20e60890b05a0e Mon Sep 17 00:00:00 2001 From: DenOfEquity <166248528+DenOfEquity@users.noreply.github.com> Date: Fri, 8 Nov 2024 23:16:29 +0000 Subject: [PATCH 2/6] fix for SAG with Kohya HRFix/ Deep Shrink (#5546) now works with arbitrary downscale factors --- comfy_extras/nodes_sag.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/comfy_extras/nodes_sag.py b/comfy_extras/nodes_sag.py index 5e15b99e5..3f03533e2 100644 --- a/comfy_extras/nodes_sag.py +++ b/comfy_extras/nodes_sag.py @@ -57,12 +57,17 @@ def create_blur_map(x0, attn, sigma=3.0, threshold=1.0): attn = attn.reshape(b, -1, hw1, hw2) # Global Average Pool mask = attn.mean(1, keepdim=False).sum(1, keepdim=False) > threshold - ratio = 2**(math.ceil(math.sqrt(lh * lw / hw1)) - 1).bit_length() - mid_shape = [math.ceil(lh / ratio), math.ceil(lw / ratio)] + + f = float(lh) / float(lw) + fh = f ** 0.5 + fw = (1/f) ** 0.5 + S = mask.size(1) ** 0.5 + w = int(0.5 + S * fw) + h = int(0.5 + S * fh) # Reshape mask = ( - mask.reshape(b, *mid_shape) + mask.reshape(b, h, w) .unsqueeze(1) .type(attn.dtype) ) From 6ee066a14f345df0620614b783949ecdbef34684 Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sat, 9 Nov 2024 00:13:34 +0000 Subject: [PATCH 3/6] Live terminal output (#5396) * Add /logs/raw and /logs/subscribe for getting logs on frontend Hijacks stderr/stdout to send all output data to the client on flush * Use existing send sync method * Fix get_logs should return string * Fix bug * pass no server * fix tests * Fix output flush on linux --- api_server/routes/internal/internal_routes.py | 29 ++++++++- api_server/services/terminal_service.py | 47 ++++++++++++++ app/logger.py | 64 +++++++++++++++---- server.py | 2 +- .../server/routes/internal_routes_test.py | 6 +- 5 files changed, 131 insertions(+), 17 deletions(-) create mode 100644 api_server/services/terminal_service.py diff --git a/api_server/routes/internal/internal_routes.py b/api_server/routes/internal/internal_routes.py index 63704f13a..8893500aa 100644 --- a/api_server/routes/internal/internal_routes.py +++ b/api_server/routes/internal/internal_routes.py @@ -2,6 +2,7 @@ from aiohttp import web from typing import Optional from folder_paths import models_dir, user_directory, output_directory, folder_names_and_paths from api_server.services.file_service import FileService +from api_server.services.terminal_service import TerminalService import app.logger class InternalRoutes: @@ -11,7 +12,8 @@ class InternalRoutes: Check README.md for more information. ''' - def __init__(self): + + def __init__(self, prompt_server): self.routes: web.RouteTableDef = web.RouteTableDef() self._app: Optional[web.Application] = None self.file_service = FileService({ @@ -19,6 +21,8 @@ class InternalRoutes: "user": user_directory, "output": output_directory }) + self.prompt_server = prompt_server + self.terminal_service = TerminalService(prompt_server) def setup_routes(self): @self.routes.get('/files') @@ -34,7 +38,28 @@ class InternalRoutes: @self.routes.get('/logs') async def get_logs(request): - return web.json_response(app.logger.get_logs()) + return web.json_response("".join([(l["t"] + " - " + l["m"]) for l in app.logger.get_logs()])) + + @self.routes.get('/logs/raw') + async def get_logs(request): + self.terminal_service.update_size() + return web.json_response({ + "entries": list(app.logger.get_logs()), + "size": {"cols": self.terminal_service.cols, "rows": self.terminal_service.rows} + }) + + @self.routes.patch('/logs/subscribe') + async def subscribe_logs(request): + json_data = await request.json() + client_id = json_data["clientId"] + enabled = json_data["enabled"] + if enabled: + self.terminal_service.subscribe(client_id) + else: + self.terminal_service.unsubscribe(client_id) + + return web.Response(status=200) + @self.routes.get('/folder_paths') async def get_folder_paths(request): diff --git a/api_server/services/terminal_service.py b/api_server/services/terminal_service.py new file mode 100644 index 000000000..284afab5a --- /dev/null +++ b/api_server/services/terminal_service.py @@ -0,0 +1,47 @@ +from app.logger import on_flush +import os + + +class TerminalService: + def __init__(self, server): + self.server = server + self.cols = None + self.rows = None + self.subscriptions = set() + on_flush(self.send_messages) + + def update_size(self): + sz = os.get_terminal_size() + changed = False + if sz.columns != self.cols: + self.cols = sz.columns + changed = True + + if sz.lines != self.rows: + self.rows = sz.lines + changed = True + + if changed: + return {"cols": self.cols, "rows": self.rows} + + return None + + def subscribe(self, client_id): + self.subscriptions.add(client_id) + + def unsubscribe(self, client_id): + self.subscriptions.discard(client_id) + + def send_messages(self, entries): + if not len(entries) or not len(self.subscriptions): + return + + new_size = self.update_size() + + for client_id in self.subscriptions.copy(): # prevent: Set changed size during iteration + if client_id not in self.server.sockets: + # Automatically unsub if the socket has disconnected + self.unsubscribe(client_id) + continue + + self.server.send_sync("logs", {"entries": entries, "size": new_size}, client_id) diff --git a/app/logger.py b/app/logger.py index 4ca0ea88e..527be9fe7 100644 --- a/app/logger.py +++ b/app/logger.py @@ -1,20 +1,69 @@ -import logging -from logging.handlers import MemoryHandler from collections import deque +from datetime import datetime +import io +import logging +import sys +import threading logs = None -formatter = logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s") +stdout_interceptor = None +stderr_interceptor = None + + +class LogInterceptor(io.TextIOWrapper): + def __init__(self, stream, *args, **kwargs): + buffer = stream.buffer + encoding = stream.encoding + super().__init__(buffer, *args, **kwargs, encoding=encoding, line_buffering=stream.line_buffering) + self._lock = threading.Lock() + self._flush_callbacks = [] + self._logs_since_flush = [] + + def write(self, data): + entry = {"t": datetime.now().isoformat(), "m": data} + with self._lock: + self._logs_since_flush.append(entry) + + # Simple handling for cr to overwrite the last output if it isnt a full line + # else logs just get full of progress messages + if isinstance(data, str) and data.startswith("\r") and not logs[-1]["m"].endswith("\n"): + logs.pop() + logs.append(entry) + super().write(data) + + def flush(self): + super().flush() + for cb in self._flush_callbacks: + cb(self._logs_since_flush) + self._logs_since_flush = [] + + def on_flush(self, callback): + self._flush_callbacks.append(callback) def get_logs(): - return "\n".join([formatter.format(x) for x in logs]) + return logs +def on_flush(callback): + if stdout_interceptor is not None: + stdout_interceptor.on_flush(callback) + if stderr_interceptor is not None: + stderr_interceptor.on_flush(callback) + def setup_logger(log_level: str = 'INFO', capacity: int = 300): global logs if logs: return + # Override output streams and log to buffer + logs = deque(maxlen=capacity) + + global stdout_interceptor + global stderr_interceptor + stdout_interceptor = sys.stdout = LogInterceptor(sys.stdout) + stderr_interceptor = sys.stderr = LogInterceptor(sys.stderr) + # Setup default global logger logger = logging.getLogger() logger.setLevel(log_level) @@ -22,10 +71,3 @@ def setup_logger(log_level: str = 'INFO', capacity: int = 300): stream_handler = logging.StreamHandler() stream_handler.setFormatter(logging.Formatter("%(message)s")) logger.addHandler(stream_handler) - - # Create a memory handler with a deque as its buffer - logs = deque(maxlen=capacity) - memory_handler = MemoryHandler(capacity, flushLevel=logging.INFO) - memory_handler.buffer = logs - memory_handler.setFormatter(formatter) - logger.addHandler(memory_handler) diff --git a/server.py b/server.py index ada6d90c3..e663095bc 100644 --- a/server.py +++ b/server.py @@ -152,7 +152,7 @@ class PromptServer(): mimetypes.types_map['.js'] = 'application/javascript; charset=utf-8' self.user_manager = UserManager() - self.internal_routes = InternalRoutes() + self.internal_routes = InternalRoutes(self) self.supports = ["custom_nodes_from_web"] self.prompt_queue = None self.loop = loop diff --git a/tests-unit/server/routes/internal_routes_test.py b/tests-unit/server/routes/internal_routes_test.py index 2d2b43bd6..4fe544249 100644 --- a/tests-unit/server/routes/internal_routes_test.py +++ b/tests-unit/server/routes/internal_routes_test.py @@ -8,7 +8,7 @@ from folder_paths import models_dir, user_directory, output_directory @pytest.fixture def internal_routes(): - return InternalRoutes() + return InternalRoutes(None) @pytest.fixture def aiohttp_client_factory(aiohttp_client, internal_routes): @@ -102,7 +102,7 @@ async def test_file_service_initialization(): # Create a mock instance mock_file_service_instance = MagicMock(spec=FileService) MockFileService.return_value = mock_file_service_instance - internal_routes = InternalRoutes() + internal_routes = InternalRoutes(None) # Check if FileService was initialized with the correct parameters MockFileService.assert_called_once_with({ @@ -112,4 +112,4 @@ async def test_file_service_initialization(): }) # Verify that the file_service attribute of InternalRoutes is set - assert internal_routes.file_service == mock_file_service_instance \ No newline at end of file + assert internal_routes.file_service == mock_file_service_instance From 8b90e50979b0d33e1f8d10d5c938361f59f95474 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Sat, 9 Nov 2024 07:10:43 -0500 Subject: [PATCH 4/6] Properly handle and reshape masks when used on 3d latents. --- comfy/sampler_helpers.py | 8 ++------ comfy/utils.py | 21 +++++++++++++++++++++ 2 files changed, 23 insertions(+), 6 deletions(-) diff --git a/comfy/sampler_helpers.py b/comfy/sampler_helpers.py index 4a2ec123b..1879e670a 100644 --- a/comfy/sampler_helpers.py +++ b/comfy/sampler_helpers.py @@ -1,14 +1,10 @@ import torch import comfy.model_management import comfy.conds +import comfy.utils def prepare_mask(noise_mask, shape, device): - """ensures noise mask is of proper dimensions""" - noise_mask = torch.nn.functional.interpolate(noise_mask.reshape((-1, 1, noise_mask.shape[-2], noise_mask.shape[-1])), size=(shape[2], shape[3]), mode="bilinear") - noise_mask = torch.cat([noise_mask] * shape[1], dim=1) - noise_mask = comfy.utils.repeat_to_batch_size(noise_mask, shape[0]) - noise_mask = noise_mask.to(device) - return noise_mask + return comfy.utils.reshape_mask(noise_mask, shape).to(device) def get_models_from_cond(cond, model_type): models = [] diff --git a/comfy/utils.py b/comfy/utils.py index cc92e1115..3c5d06a4f 100644 --- a/comfy/utils.py +++ b/comfy/utils.py @@ -848,3 +848,24 @@ class ProgressBar: def update(self, value): self.update_absolute(self.current + value) + +def reshape_mask(input_mask, output_shape): + dims = len(output_shape) - 2 + + if dims == 1: + scale_mode = "linear" + + if dims == 2: + mask = input_mask.reshape((-1, 1, input_mask.shape[-2], input_mask.shape[-1])) + scale_mode = "bilinear" + + if dims == 3: + if len(input_mask.shape) < 5: + mask = input_mask.reshape((1, 1, -1, input_mask.shape[-2], input_mask.shape[-1])) + scale_mode = "trilinear" + + mask = torch.nn.functional.interpolate(mask, size=output_shape[2:], mode=scale_mode) + if mask.shape[1] < output_shape[1]: + mask = mask.repeat((1, output_shape[1]) + (1,) * dims)[:,:output_shape[1]] + mask = comfy.utils.repeat_to_batch_size(mask, output_shape[0]) + return mask From 9c1ed58ef218c8a20bd862de329d7d350d69a34d Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Sun, 10 Nov 2024 00:10:45 -0500 Subject: [PATCH 5/6] proper fix for sag. --- comfy_extras/nodes_sag.py | 21 ++++++++++++++------- 1 file changed, 14 insertions(+), 7 deletions(-) diff --git a/comfy_extras/nodes_sag.py b/comfy_extras/nodes_sag.py index 3f03533e2..1bd8d7364 100644 --- a/comfy_extras/nodes_sag.py +++ b/comfy_extras/nodes_sag.py @@ -58,16 +58,23 @@ def create_blur_map(x0, attn, sigma=3.0, threshold=1.0): # Global Average Pool mask = attn.mean(1, keepdim=False).sum(1, keepdim=False) > threshold - f = float(lh) / float(lw) - fh = f ** 0.5 - fw = (1/f) ** 0.5 - S = mask.size(1) ** 0.5 - w = int(0.5 + S * fw) - h = int(0.5 + S * fh) + total = mask.shape[-1] + x = round(math.sqrt((lh / lw) * total)) + xx = None + for i in range(0, math.floor(math.sqrt(total) / 2)): + for j in [(x + i), max(1, x - i)]: + if total % j == 0: + xx = j + break + if xx is not None: + break + + x = xx + y = total // x # Reshape mask = ( - mask.reshape(b, h, w) + mask.reshape(b, x, y) .unsqueeze(1) .type(attn.dtype) ) From bdeb1c171cd773a4df895a77fa384b58fe8acac5 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Sun, 10 Nov 2024 03:39:35 -0500 Subject: [PATCH 6/6] Fast previews for mochi. --- comfy/latent_formats.py | 16 +++++++++++++++- latent_preview.py | 7 ++++++- 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/comfy/latent_formats.py b/comfy/latent_formats.py index a48f60c74..f9fd16d8c 100644 --- a/comfy/latent_formats.py +++ b/comfy/latent_formats.py @@ -190,7 +190,21 @@ class Mochi(LatentFormat): 0.9294154431013696, 1.3720942357788521, 0.881393668867029, 0.9168315692124348, 0.9185249279345552, 0.9274757570805041]).view(1, self.latent_channels, 1, 1, 1) - self.latent_rgb_factors = None #TODO + self.latent_rgb_factors =[ + [-0.0069, -0.0045, 0.0018], + [ 0.0154, -0.0692, -0.0274], + [ 0.0333, 0.0019, 0.0206], + [-0.1390, 0.0628, 0.1678], + [-0.0725, 0.0134, -0.1898], + [ 0.0074, -0.0270, -0.0209], + [-0.0176, -0.0277, -0.0221], + [ 0.5294, 0.5204, 0.3852], + [-0.0326, -0.0446, -0.0143], + [-0.0659, 0.0153, -0.0153], + [ 0.0185, -0.0217, 0.0014], + [-0.0396, -0.0495, -0.0281] + ] + self.latent_rgb_factors_bias = [-0.0940, -0.1418, -0.1453] self.taesd_decoder_name = None #TODO def process_in(self, latent): diff --git a/latent_preview.py b/latent_preview.py index ae9211a27..d60e68d55 100644 --- a/latent_preview.py +++ b/latent_preview.py @@ -47,7 +47,12 @@ class Latent2RGBPreviewer(LatentPreviewer): if self.latent_rgb_factors_bias is not None: self.latent_rgb_factors_bias = self.latent_rgb_factors_bias.to(dtype=x0.dtype, device=x0.device) - latent_image = torch.nn.functional.linear(x0[0].permute(1, 2, 0), self.latent_rgb_factors, bias=self.latent_rgb_factors_bias) + if x0.ndim == 5: + x0 = x0[0, :, 0] + else: + x0 = x0[0] + + latent_image = torch.nn.functional.linear(x0.movedim(0, -1), self.latent_rgb_factors, bias=self.latent_rgb_factors_bias) # latent_image = x0[0].permute(1, 2, 0) @ self.latent_rgb_factors return preview_to_image(latent_image)