From 821a031bfcca48143d9db3a88208f982e1651fa1 Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Thu, 26 Jun 2025 12:56:38 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20=E2=9C=A8=20yet=20another=20repl=20node?= =?UTF-8?q?=20for=20comfy?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Still unsure how to expose it --- __init__.py | 4 + features.json | 3 + repl.py | 637 ++++++++++++++++++++++++++++++++++++++++++++++++ web/mtb_repl.js | 529 ++++++++++++++++++++++++++++++++++++++++ 4 files changed, 1173 insertions(+) create mode 100644 features.json create mode 100644 repl.py create mode 100644 web/mtb_repl.js diff --git a/__init__.py b/__init__.py index 2d2cca4..d1274bc 100644 --- a/__init__.py +++ b/__init__.py @@ -265,6 +265,10 @@ def register_routes(): from PIL import Image + from .repl import setup_custom_web_routes + + setup_custom_web_routes(PromptServer.instance.app) + with contextlib.suppress(ImportError): from cachetools import TTLCache diff --git a/features.json b/features.json new file mode 100644 index 0000000..9a58c45 --- /dev/null +++ b/features.json @@ -0,0 +1,3 @@ +{ + "use_repl": false +} diff --git a/repl.py b/repl.py new file mode 100644 index 0000000..b85391c --- /dev/null +++ b/repl.py @@ -0,0 +1,637 @@ +import base64 +import code +import io +import re +import sys +from contextlib import redirect_stderr, redirect_stdout + +# import matplotlib.pyplot as plt +import numpy as np +import torch +from aiohttp import web +from PIL import Image +from rich.console import Console +from rich.traceback import Traceback + +from .log import log + +try: + import pyflakes.api + import pyflakes.reporter + + _HAS_LINT = True +except ImportError: + print( + "ComfyREPL: pyflakes not found. Linting will be disabled. Install with 'pip install pyflakes'." + ) + _HAS_LINT = False + +# --- Linting Library --- +# try: +# import ruff +# import ruff.lint +# import ruff.lint.linter +# import ruff.settings +# +# _HAS_LINT = True +# except ImportError: +# print( +# "ComfyREPL: ruff not found. Linting will be disabled. Install with 'pip install ruff'." +# ) +# _HAS_LINT = False + +# --- Audio/Video Libraries --- +try: + import scipy.io.wavfile + + _HAS_SCIPY = True +except ImportError: + print( + "ComfyREPL: SciPy not found. Audio display will be disabled. Install with 'pip install scipy'." + ) + _HAS_SCIPY = False + +try: + import imageio + import imageio.plugins.ffmpeg # Ensure ffmpeg plugin is available + + _HAS_IMAGEIO = True +except ImportError: + print( + "ComfyREPL: Imageio or imageio-ffmpeg not found. Video display will be disabled. Install with 'pip install imageio imageio-ffmpeg'." + ) + _HAS_IMAGEIO = False + + +# --- Audio Display --- +class AudioDisplay: + def __init__(self, samples, sample_rate): + if not _HAS_SCIPY: + raise ImportError("Audio display requires scipy and numpy.") + if not isinstance(samples, (np.ndarray, torch.Tensor)): + raise TypeError( + "Audio samples must be a numpy array or torch tensor." + ) + if isinstance(samples, torch.Tensor): + samples = samples.detach().cpu().numpy() + + # Ensure samples are in a format scipy.io.wavfile can handle (e.g., int16, float32) + if samples.dtype == np.float64: + samples = samples.astype(np.float32) + elif samples.dtype == np.int64: + # Or scale to int32 if range requires + samples = samples.astype(np.int16) + + self.samples = samples + self.sample_rate = sample_rate + + def _to_wav_base64(self): + buffer = io.BytesIO() + try: + scipy.io.wavfile.write(buffer, self.sample_rate, self.samples) + audio_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8") + return audio_base64 + except Exception as e: + return f"
Error encoding audio: {e}
" + + def _repr_html_(self): + base64_data = self._to_wav_base64() + if base64_data.startswith("' + + +def render_audio(samples, sample_rate): + """ + Render audio samples as an HTML audio player. + + Args: + samples (np.ndarray or torch.Tensor): Audio samples. + sample_rate (int): Sample rate in Hz. + + Returns + ------- + AudioDisplay: An object that will render as an HTML audio player. + """ + return AudioDisplay(samples, sample_rate) + + +# --- Display Classes --- +class VideoDisplay: + def __init__(self, frames, fps=24, options=None): + if not _HAS_IMAGEIO: # numpy/PIL/torch needed for frames + raise ImportError( + "Video display requires imageio, imageio-ffmpeg, and image libraries (numpy, Pillow, torch)." + ) + + self.frames = [] + for frame in frames: + if isinstance(frame, Image.Image): + self.frames.append(np.array(frame)) + elif isinstance(frame, np.ndarray): + # Ensure HWC and uint8 + if frame.ndim == 3 and frame.shape[0] in [1, 3, 4]: # CHW + frame = np.transpose(frame, (1, 2, 0)) + if frame.dtype != np.uint8: + frame = ( + (frame * 255).astype(np.uint8) + if frame.max() <= 1.0 + else frame.astype(np.uint8) + ) + self.frames.append(frame) + elif isinstance(frame, torch.Tensor): + np_frame = frame.detach().cpu().numpy() + if np_frame.ndim == 3 and np_frame.shape[0] in [ + 1, + 3, + 4, + ]: # CHW + np_frame = np.transpose(np_frame, (1, 2, 0)) + if np_frame.dtype != np.uint8: + np_frame = ( + (np_frame * 255).astype(np.uint8) + if np_frame.max() <= 1.0 + else np_frame.astype(np.uint8) + ) + self.frames.append(np_frame) + else: + raise TypeError( + f"Unsupported frame type: {type(frame)}. Must be PIL.Image, numpy.ndarray, or torch.Tensor." + ) + + self.fps = fps + self.options = options if options is not None else {} + + def _to_mp4_base64(self): + buffer = io.BytesIO() + try: + # Use imageio to write frames to an in-memory MP4 file + imageio.mimwrite( + buffer, + self.frames, + format="mp4", + fps=self.fps, + codec="libx264", + quality=8, + ) # quality 1-10 + video_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8") + return video_base64 + except Exception as e: + return f"
Error encoding video: {e}
" + + def _repr_html_(self): + base64_data = self._to_mp4_base64() + if base64_data.startswith("' + + +def render_video(batch_tensor_or_array_of_pil_images, fps=24, options=None): + """ + Render video frames as an HTML video player. + + Args: + batch_tensor_or_array_of_pil_images (list of PIL.Image, np.ndarray, or torch.Tensor): + A list of frames, or a single batch tensor/array (B, H, W, C) or (B, C, H, W). + fps (int): Frames per second. + options (dict): Dictionary of HTML