From 3bb0fc56c46fee97ed1d7cda97bc2604a785ad52 Mon Sep 17 00:00:00 2001 From: Fillip Date: Tue, 26 May 2026 23:47:41 -0700 Subject: [PATCH] feat: add context window ksampler --- __init__.py | 3 + nodes/ksamplers/FL_KsamplerContextWindow.py | 357 ++++++++++++++++++ pyproject.toml | 2 +- .../ksamplers/FL_KsamplerContextWindow.js | 198 ++++++++++ 4 files changed, 559 insertions(+), 1 deletion(-) create mode 100644 nodes/ksamplers/FL_KsamplerContextWindow.py create mode 100644 web/nodes/ksamplers/FL_KsamplerContextWindow.js diff --git a/__init__.py b/__init__.py index 673c6b8..d4224f8 100644 --- a/__init__.py +++ b/__init__.py @@ -141,6 +141,7 @@ from .nodes.image.FL_SaveWebpImages import FL_SaveWebPImage # KSAMPLERS NODES from .nodes.ksamplers.FL_KsamplerBasic import FL_KsamplerBasic +from .nodes.ksamplers.FL_KsamplerContextWindow import FL_KsamplerContextWindow from .nodes.ksamplers.FL_KsamplerPlus import FL_KsamplerPlus from .nodes.ksamplers.FL_KsamplerPlusV2 import FL_KsamplerPlusV2 from .nodes.ksamplers.FL_KsamplerSigma import FL_KsamplerSigma @@ -310,6 +311,7 @@ NODE_CLASS_MAPPINGS = { "FL_KsamplerPlus": FL_KsamplerPlus, "FL_KsamplerPlusV2": FL_KsamplerPlusV2, "FL_KsamplerBasic": FL_KsamplerBasic, + "FL_KsamplerContextWindow": FL_KsamplerContextWindow, "FL_KsamplerSigma": FL_KsamplerSigma, "FL_KsamplerSEG_Regions": FL_KsamplerSEG_Regions, "FL_KsamplerSEG_Captioner": FL_KsamplerSEG_Captioner, @@ -510,6 +512,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FL_KsamplerPlus": "FL KSampler Plus", "FL_KsamplerPlusV2": "FL KSampler Plus V2", "FL_KsamplerBasic": "FL KSampler Basic", + "FL_KsamplerContextWindow": "FL Context Window KSampler", "FL_KsamplerSigma": "FL KSampler Sigma", "FL_KsamplerSEG_Regions": "FL KSampler SEG Regions", "FL_KsamplerSEG_Captioner": "FL KSampler SEG Captioner", diff --git a/nodes/ksamplers/FL_KsamplerContextWindow.py b/nodes/ksamplers/FL_KsamplerContextWindow.py new file mode 100644 index 0000000..776824b --- /dev/null +++ b/nodes/ksamplers/FL_KsamplerContextWindow.py @@ -0,0 +1,357 @@ +import logging +import time + +import torch + +import comfy.context_windows +import comfy.samplers +from comfy_execution.utils import get_executing_context +from nodes import VAEDecode, VAEEncode, common_ksampler + + +CONTEXT_SCHEDULES = [ + comfy.context_windows.ContextSchedules.STATIC_STANDARD, + comfy.context_windows.ContextSchedules.UNIFORM_STANDARD, + comfy.context_windows.ContextSchedules.UNIFORM_LOOPED, + comfy.context_windows.ContextSchedules.BATCHED, +] + +FUSE_METHODS = comfy.context_windows.ContextFuseMethods.LIST_STATIC + +TEMPORAL_UNITS = [ + "video_frames_4n_plus_1", + "latent_frames", +] + + +class FLSafeIndexListContextHandler(comfy.context_windows.IndexListContextHandler): + def __init__(self, *args, node_id=None, **kwargs): + super().__init__(*args, **kwargs) + self.node_id = str(node_id) if node_id is not None else None + self._total_steps = 1 + self._last_total_windows = 1 + self._last_progress_value = 0 + self._last_event_at = 0.0 + + def set_step(self, timestep: torch.Tensor, model_options: dict[str]): + sample_sigmas = model_options.get("transformer_options", {}).get("sample_sigmas") + if sample_sigmas is None: + return + self._total_steps = max(1, int(sample_sigmas.numel()) - 1) + current_timestep = timestep[0].to(device=sample_sigmas.device, dtype=sample_sigmas.dtype) + mask = torch.isclose(sample_sigmas, current_timestep, rtol=0.0001) + matches = torch.nonzero(mask) + if torch.numel(matches) == 0: + return + self._step = int(matches[0].item()) + + def get_context_windows(self, model, x_in: torch.Tensor, model_options: dict): + context_windows = super().get_context_windows(model, x_in, model_options) + self._last_total_windows = max(1, len(context_windows)) + return context_windows + + def combine_context_window_results( + self, + x_in: torch.Tensor, + sub_conds_out, + sub_conds, + window, + window_idx: int, + total_windows: int, + timestep: torch.Tensor, + conds_final, + counts_final, + biases_final, + ): + result = super().combine_context_window_results( + x_in, + sub_conds_out, + sub_conds, + window, + window_idx, + total_windows, + timestep, + conds_final, + counts_final, + biases_final, + ) + self._emit_progress(window_idx, max(1, total_windows), window) + return result + + def emit_done(self): + if self.node_id is None: + return + max_value = max(1, self._total_steps * self._last_total_windows) + self._send_event( + { + "node": self.node_id, + "status": "done", + "value": max_value, + "max": max_value, + "step": self._total_steps, + "total_steps": self._total_steps, + "window_index": self._last_total_windows, + "total_windows": self._last_total_windows, + "window": [], + } + ) + + def _emit_progress(self, window_idx: int, total_windows: int, window): + if self.node_id is None: + return + + max_value = max(1, self._total_steps * total_windows) + raw_value = min(max_value, self._step * total_windows + window_idx + 1) + value = max(self._last_progress_value, raw_value) + self._last_progress_value = value + + now = time.monotonic() + if now - self._last_event_at < 0.25 and value < max_value: + return + self._last_event_at = now + + self._send_event( + { + "node": self.node_id, + "status": "running", + "value": value, + "max": max_value, + "step": min(self._step + 1, self._total_steps), + "total_steps": self._total_steps, + "window_index": window_idx + 1, + "total_windows": total_windows, + "window": list(getattr(window, "index_list", [])), + } + ) + + @staticmethod + def _send_event(payload: dict): + try: + from server import PromptServer + + PromptServer.instance.send_sync("fl_context_window_progress", payload) + except Exception as e: + logging.debug(f"[FL_KsamplerContextWindow] progress event send failed: {e}") + + +class FL_KsamplerContextWindow: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "latent_image": ("LATENT",), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), + "cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 100.0}), + "sampler_name": (comfy.samplers.KSampler.SAMPLERS,), + "scheduler": (comfy.samplers.KSampler.SCHEDULERS,), + "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "context_length": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4}), + "context_overlap": ("INT", {"default": 30, "min": 0, "max": 10000}), + "context_schedule": (CONTEXT_SCHEDULES, {"default": comfy.context_windows.ContextSchedules.STATIC_STANDARD}), + "context_stride": ("INT", {"default": 1, "min": 1, "max": 10000}), + "fuse_method": (FUSE_METHODS, {"default": comfy.context_windows.ContextFuseMethods.PYRAMID}), + "temporal_unit": (TEMPORAL_UNITS, {"default": "video_frames_4n_plus_1"}), + "closed_loop": ("BOOLEAN", {"default": False, "advanced": True}), + "freenoise": ("BOOLEAN", {"default": False, "advanced": True}), + "causal_window_fix": ("BOOLEAN", {"default": True, "advanced": True}), + "temporal_dim": ("INT", {"default": 2, "min": 0, "max": 5, "advanced": True}), + "cond_retain_index_list": ("STRING", {"default": "", "multiline": False, "advanced": True}), + "split_conds_to_windows": ("BOOLEAN", {"default": False, "advanced": True}), + }, + "optional": { + "vae": ("VAE",), + "image": ("IMAGE",), + }, + "hidden": {"unique_id": "UNIQUE_ID"}, + } + + RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "LATENT", "VAE", "IMAGE", "STRING") + RETURN_NAMES = ("model", "positive", "negative", "latent", "vae", "image", "debug_info") + FUNCTION = "sample" + CATEGORY = "🏵️Fill Nodes/Ksamplers" + + def sample( + self, + model, + positive, + negative, + latent_image, + seed, + steps, + cfg, + sampler_name, + scheduler, + denoise, + context_length, + context_overlap, + context_schedule, + context_stride, + fuse_method, + temporal_unit, + vae=None, + image=None, + closed_loop=False, + freenoise=False, + causal_window_fix=True, + temporal_dim=2, + cond_retain_index_list="", + split_conds_to_windows=False, + unique_id=None, + ): + try: + node_id = unique_id + if node_id is None: + context = get_executing_context() + if context is not None: + node_id = context.node_id + + if image is not None: + if vae is None: + raise ValueError("FL_KsamplerContextWindow: image input requires a VAE.") + latent_image = VAEEncode().encode(vae, image)[0] + + if latent_image is None or "samples" not in latent_image: + raise ValueError("FL_KsamplerContextWindow: latent_image must contain samples.") + + samples = latent_image["samples"] + if not isinstance(samples, torch.Tensor): + raise ValueError("FL_KsamplerContextWindow: nested tensor latents are not supported.") + + if temporal_dim >= samples.ndim: + raise ValueError( + "FL_KsamplerContextWindow: temporal_dim is outside the latent sample shape. " + f"Got temporal_dim={temporal_dim}, shape={tuple(samples.shape)}." + ) + + if samples.ndim == 4 and temporal_dim != 0: + raise ValueError( + "FL_KsamplerContextWindow: expected a 5D video latent [B, C, T, H, W]. " + "For 4D latents, set temporal_dim=0 only if you intentionally want to window over batch." + ) + + latent_context_length, latent_context_overlap = self._convert_context_units( + context_length, + context_overlap, + temporal_unit, + ) + self._validate_context(latent_context_length, latent_context_overlap) + + context_model = model.clone() + context_model.model_options["context_handler"] = FLSafeIndexListContextHandler( + context_schedule=comfy.context_windows.get_matching_context_schedule(context_schedule), + fuse_method=comfy.context_windows.get_matching_fuse_method(fuse_method), + context_length=latent_context_length, + context_overlap=latent_context_overlap, + context_stride=context_stride, + closed_loop=closed_loop, + dim=temporal_dim, + freenoise=freenoise, + cond_retain_index_list=cond_retain_index_list, + split_conds_to_windows=split_conds_to_windows, + causal_window_fix=causal_window_fix, + node_id=node_id, + ) + context_handler = context_model.model_options["context_handler"] + comfy.context_windows.create_prepare_sampling_wrapper(context_model) + if freenoise: + comfy.context_windows.create_sampler_sample_wrapper(context_model) + + sampled = common_ksampler( + context_model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + denoise=denoise, + )[0] + context_handler.emit_done() + + output_image = None + if vae is not None: + output_image = VAEDecode().decode(vae, sampled)[0] + + debug_info = self._debug_info( + samples=samples, + temporal_dim=temporal_dim, + context_length=context_length, + context_overlap=context_overlap, + latent_context_length=latent_context_length, + latent_context_overlap=latent_context_overlap, + context_schedule=context_schedule, + context_stride=context_stride, + closed_loop=closed_loop, + fuse_method=fuse_method, + freenoise=freenoise, + causal_window_fix=causal_window_fix, + temporal_unit=temporal_unit, + ) + + return (model, positive, negative, sampled, vae, output_image, debug_info) + + except Exception as e: + logging.error(f"Error in FL_KsamplerContextWindow: {e}") + raise + + @staticmethod + def _convert_context_units(context_length, context_overlap, temporal_unit): + if temporal_unit == "latent_frames": + return int(context_length), int(context_overlap) + if temporal_unit == "video_frames_4n_plus_1": + latent_length = max(((int(context_length) - 1) // 4) + 1, 1) + latent_overlap = max(((int(context_overlap) - 1) // 4) + 1, 0) if context_overlap > 0 else 0 + return latent_length, latent_overlap + raise ValueError(f"FL_KsamplerContextWindow: unknown temporal_unit '{temporal_unit}'.") + + @staticmethod + def _validate_context(context_length, context_overlap): + if context_length < 1: + raise ValueError("FL_KsamplerContextWindow: context_length must be at least 1.") + if context_overlap < 0: + raise ValueError("FL_KsamplerContextWindow: context_overlap cannot be negative.") + if context_overlap >= context_length: + raise ValueError("FL_KsamplerContextWindow: context_overlap must be smaller than context_length.") + + @staticmethod + def _debug_info( + samples, + temporal_dim, + context_length, + context_overlap, + latent_context_length, + latent_context_overlap, + context_schedule, + context_stride, + closed_loop, + fuse_method, + freenoise, + causal_window_fix, + temporal_unit, + ): + total_temporal = samples.shape[temporal_dim] + context_active = total_temporal > latent_context_length + return ( + "FL Context Window KSampler\n" + f"- latent_shape: {tuple(samples.shape)}\n" + f"- temporal_dim: {temporal_dim}\n" + f"- total_temporal_length: {total_temporal}\n" + f"- temporal_unit: {temporal_unit}\n" + f"- requested_context_length: {context_length}\n" + f"- requested_context_overlap: {context_overlap}\n" + f"- effective_latent_context_length: {latent_context_length}\n" + f"- effective_latent_context_overlap: {latent_context_overlap}\n" + f"- context_schedule: {context_schedule}\n" + f"- context_stride: {context_stride}\n" + f"- closed_loop: {closed_loop}\n" + f"- fuse_method: {fuse_method}\n" + f"- freenoise: {freenoise}\n" + f"- causal_window_fix: {causal_window_fix}\n" + f"- context_active: {context_active}" + ) diff --git a/pyproject.toml b/pyproject.toml index 45cb836..f8db51e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_fill-nodes" description = "Fill-Nodes is a versatile collection of custom nodes for ComfyUI that extends functionality across multiple domains. Features include advanced image processing (pixelation, slicing, masking), visual effects generation (glitch, halftone, pixel art), comprehensive file handling (PDF creation/extraction, Google Drive integration), AI model interfaces (GPT, DALL-E, Hugging Face), utility nodes for workflow enhancement, and specialized tools for video processing, captioning, and batch operations. The pack provides both practical workflow solutions and creative tools within a unified node collection." -version = "2.7.6" +version = "2.7.7" license = "LICENSE" dependencies = ["librosa", "sounddevice", "glitch_this", "PyOpenGL", "glfw", "scipy>=1.13.1", "requests", "aiohttp", "moviepy", "matplotlib", "reportlab", "openai", "PyPDF2", "pdf2image", "PyMuPDF", "reportlab", "PyPDF2", "ollama", "kornia", "opencv-python", "gdown", "open_clip_torch", "google-genai"] diff --git a/web/nodes/ksamplers/FL_KsamplerContextWindow.js b/web/nodes/ksamplers/FL_KsamplerContextWindow.js new file mode 100644 index 0000000..b437bf5 --- /dev/null +++ b/web/nodes/ksamplers/FL_KsamplerContextWindow.js @@ -0,0 +1,198 @@ +import { app } from "../../../../scripts/app.js"; +import { api } from "../../../../scripts/api.js"; + +const STYLES = ` + .flks-context-widget { + background: #17181c; + border: 1px solid #2a2d34; + border-radius: 8px; + color: #f4f4f5; + display: flex; + flex-direction: column; + font-family: Inter, -apple-system, BlinkMacSystemFont, sans-serif; + gap: 8px; + min-height: 104px; + padding: 10px; + box-sizing: border-box; + } + .flks-context-widget * { box-sizing: border-box; } + .flks-context-header { + align-items: center; + display: flex; + justify-content: space-between; + gap: 8px; + } + .flks-context-title { + font-size: 11px; + font-weight: 650; + line-height: 1.2; + } + .flks-context-badge { + background: #06b6d4; + border-radius: 999px; + color: white; + font-size: 10px; + font-variant-numeric: tabular-nums; + font-weight: 700; + line-height: 1; + padding: 4px 7px; + white-space: nowrap; + } + .flks-context-bar { + background: #27272a; + border-radius: 999px; + height: 9px; + overflow: hidden; + width: 100%; + } + .flks-context-fill { + background: linear-gradient(90deg, #06b6d4, #22c55e); + height: 100%; + transition: width 120ms linear; + width: 0%; + } + .flks-context-meta { + color: #cbd5e1; + display: grid; + gap: 4px; + grid-template-columns: 1fr 1fr; + font-size: 10px; + font-variant-numeric: tabular-nums; + line-height: 1.25; + } + .flks-context-window { + color: #94a3b8; + font-size: 10px; + line-height: 1.25; + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; + } +`; + +class ContextWindowProgressWidget { + constructor({ container }) { + this.container = container; + this.injectStyles(); + this.element = document.createElement("div"); + this.element.className = "flks-context-widget"; + this.element.innerHTML = ` +
+ Context Windows + idle +
+
+
+
+
+ step - / - + window - / - +
+
Run to see context-window progress.
+ `; + this.percentEl = this.element.querySelector('[data-role="percent"]'); + this.fillEl = this.element.querySelector('[data-role="fill"]'); + this.stepEl = this.element.querySelector('[data-role="step"]'); + this.windowEl = this.element.querySelector('[data-role="window"]'); + this.indicesEl = this.element.querySelector('[data-role="indices"]'); + this.container.appendChild(this.element); + } + + injectStyles() { + const id = "flks-context-window-styles"; + if (document.getElementById(id)) return; + const style = document.createElement("style"); + style.id = id; + style.textContent = STYLES; + document.head.appendChild(style); + } + + reset() { + this.percentEl.textContent = "0%"; + this.fillEl.style.width = "0%"; + this.stepEl.textContent = "step 0 / -"; + this.windowEl.textContent = "window 0 / -"; + this.indicesEl.textContent = "Waiting for first context window..."; + } + + update(detail) { + const value = Number(detail.value || 0); + const max = Math.max(1, Number(detail.max || 1)); + const pct = Math.max(0, Math.min(100, (value / max) * 100)); + this.percentEl.textContent = detail.status === "done" ? "done" : `${pct.toFixed(1)}%`; + this.fillEl.style.width = `${pct}%`; + this.stepEl.textContent = `step ${detail.step ?? "-"} / ${detail.total_steps ?? "-"}`; + this.windowEl.textContent = `window ${detail.window_index ?? "-"} / ${detail.total_windows ?? "-"}`; + + const indices = Array.isArray(detail.window) ? detail.window : []; + if (indices.length) { + const first = indices[0]; + const last = indices[indices.length - 1]; + this.indicesEl.textContent = `latent frames ${first}-${last} (${indices.length})`; + } else if (detail.status === "done") { + this.indicesEl.textContent = "Sampling completed."; + } + } + + dispose() { + this.element?.remove(); + } +} + +const INSTANCES = new Map(); + +app.registerExtension({ + name: "ComfyUI.FL_KsamplerContextWindow", + nodeCreated(node) { + const comfyClass = (node.constructor && node.constructor.comfyClass) || ""; + if (comfyClass !== "FL_KsamplerContextWindow") return; + + const container = document.createElement("div"); + container.style.width = "100%"; + container.style.minHeight = "104px"; + + const widget = node.addDOMWidget( + "context_progress", + "flks-context-window-progress", + container, + { + getMinHeight: () => 130, + hideOnZoom: false, + serialize: false, + } + ); + + const [oldW, oldH] = node.size; + node.setSize([Math.max(oldW, 330), Math.max(oldH, 730)]); + + setTimeout(() => { + const inst = new ContextWindowProgressWidget({ container }); + INSTANCES.set(node.id, inst); + }, 50); + + widget.onRemove = () => { + const inst = INSTANCES.get(node.id); + if (inst) { + inst.dispose(); + INSTANCES.delete(node.id); + } + }; + }, +}); + +api.addEventListener("executing", (event) => { + const detail = event.detail; + if (!detail || !detail.node) return; + const nodeId = parseInt(detail.node, 10); + const inst = INSTANCES.get(nodeId); + if (inst) inst.reset(); +}); + +api.addEventListener("fl_context_window_progress", (event) => { + const detail = event.detail; + if (!detail) return; + const nodeId = parseInt(detail.node, 10); + const inst = INSTANCES.get(nodeId); + if (!inst) return; + inst.update(detail); +});