perf: streamline prompt relay and guide ui

This commit is contained in:
OpenClaw Agent
2026-07-09 15:51:58 +00:00
parent a536ac5a14
commit ab7d14bd89
2 changed files with 172 additions and 152 deletions

View File

@@ -3,22 +3,28 @@ import { app } from "../../scripts/app.js";
// LTX Director Guide is a pure pass-through processor node. // LTX Director Guide is a pure pass-through processor node.
// All configuration (images, insert frames, strengths) comes from // All configuration (images, insert frames, strengths) comes from
// the guide_data output of Prompt Relay Encode (Timeline). // the guide_data output of Prompt Relay Encode (Timeline).
function hideWidget(node, widgetName) {
const widget = node.widgets?.find((entry) => entry.name === widgetName);
if (!widget) return;
widget.hidden = true;
widget.options = widget.options || {};
widget.options.hidden = true;
if (!window.LiteGraph || !window.LiteGraph.vueNodesMode) {
widget.computeSize = () => [0, -4];
widget.draw = () => {};
}
if (widget.element) widget.element.style.display = "none";
}
app.registerExtension({ app.registerExtension({
name: "Comfy.LTXDirectorGuideCS", name: "Comfy.LTXDirectorGuideCS",
async nodeCreated(node) { async nodeCreated(node) {
if (node.comfyClass !== "LTXDirectorGuideCS") return; if (node.comfyClass !== "LTXDirectorGuideCS") return;
// Hide retake_mode widget on LiteGraph as it is dynamically auto-detected from the timeline data. // Hide retake_mode widget on LiteGraph as it is dynamically auto-detected from the timeline data.
const w = node.widgets?.find(x => x.name === "retake_mode"); hideWidget(node, "retake_mode");
if (w) {
w.hidden = true;
if (!w.options) w.options = {};
w.options.hidden = true;
if (!window.LiteGraph || !window.LiteGraph.vueNodesMode) {
w.computeSize = () => [0, -4];
w.draw = () => { };
}
if (w.element) w.element.style.display = "none";
}
}, },
}); });

View File

@@ -1,21 +1,28 @@
import logging import logging
import math import math
import os
import torch import torch
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
_DEBUG_PROMPT_RELAY = os.environ.get("PROMPT_RELAY_DEBUG", "").strip().lower() in {"1", "true", "yes", "on"}
_DEBUG_LOG_PATH = os.path.join(os.path.dirname(__file__), "debug_prompt_relay.log")
def _build_segment_cost(query_frames, seg, dtype):
local = seg["local_token_idx"]
distance = (query_frames[:, None] - seg["midpoint"]).abs()
cost = seg["strength"] * (torch.relu(distance - seg["window"]) ** 2) / (2 * seg["sigma"] ** 2)
return local, cost.to(dtype)
def build_temporal_cost(q_token_idx, Lq, Lk, device, dtype, tokens_per_frame): def build_temporal_cost(q_token_idx, Lq, Lk, device, dtype, tokens_per_frame):
"""Gaussian penalty matrix [Lq, Lk] for video cross-attention (integer frame indexing).""" """Gaussian penalty matrix [Lq, Lk] for video cross-attention (integer frame indexing)."""
offset = torch.zeros(Lq, Lk, device=device, dtype=dtype) offset = torch.zeros(Lq, Lk, device=device, dtype=dtype)
query_frames = torch.arange(Lq, device=device, dtype=torch.long) // tokens_per_frame query_frames = (torch.arange(Lq, device=device, dtype=torch.float32) / float(tokens_per_frame)).floor()
for seg in q_token_idx: for seg in q_token_idx:
local = seg["local_token_idx"].to(device=device) local, cost = _build_segment_cost(query_frames, seg, offset.dtype)
d = (query_frames.float()[:, None] - seg["midpoint"]).abs() offset[:, local.to(device=device)] = cost
strength = seg.get("strength", 1.0)
cost = strength * (torch.relu(d - seg["window"]) ** 2) / (2 * seg["sigma"] ** 2)
offset[:, local] = cost.to(offset.dtype)
return offset return offset
@@ -26,27 +33,27 @@ def build_temporal_cost_scaled(q_token_idx, Lq, Lk, device, dtype, latent_frames
query_frames = torch.arange(Lq, device=device, dtype=torch.float32) * latent_frames / Lq query_frames = torch.arange(Lq, device=device, dtype=torch.float32) * latent_frames / Lq
for seg in q_token_idx: for seg in q_token_idx:
local = seg["local_token_idx"].to(device=device)
d = (query_frames[:, None] - seg["midpoint"]).abs()
if is_audio: if is_audio:
sigma_val = seg.get("sigma_audio", seg["sigma"]) active_seg = {
window_val = seg.get("window_audio", seg["window"]) "local_token_idx": seg["local_token_idx"],
strength_val = seg.get("strength_audio", 1.0) "midpoint": seg["midpoint"],
"window": seg.get("window_audio", seg["window"]),
"sigma": seg.get("sigma_audio", seg["sigma"]),
"strength": seg.get("strength_audio", 1.0),
}
else: else:
sigma_val = seg["sigma"] active_seg = seg
window_val = seg["window"] local, cost = _build_segment_cost(query_frames, active_seg, offset.dtype)
strength_val = seg.get("strength", 1.0) offset[:, local.to(device=device)] = cost
cost = strength_val * (torch.relu(d - window_val) ** 2) / (2 * sigma_val ** 2)
offset[:, local] = cost.to(offset.dtype)
return offset return offset
def debug_log(msg): def debug_log(msg):
if not _DEBUG_PROMPT_RELAY:
return
try: try:
import os with open(_DEBUG_LOG_PATH, "a", encoding="utf-8") as f:
log_path = os.path.join(os.path.dirname(__file__), "debug_prompt_relay.log")
with open(log_path, "a", encoding="utf-8") as f:
f.write(msg + "\n") f.write(msg + "\n")
except Exception: except Exception:
pass pass
@@ -61,6 +68,15 @@ def create_mask_fn(q_token_idx, fallback_tokens_per_frame, latent_frames):
""" """
cache = {} cache = {}
max_token_idx = max(int(seg["local_token_idx"].max().item()) for seg in q_token_idx) + 1 max_token_idx = max(int(seg["local_token_idx"].max().item()) for seg in q_token_idx) + 1
latent_frames = max(1, int(latent_frames))
fallback_tokens_per_frame = max(1, int(fallback_tokens_per_frame))
def _get_video_query_layout(Lq, grid_sizes):
if grid_sizes is not None:
return int(grid_sizes[1]) * int(grid_sizes[2])
if Lq % latent_frames == 0:
return Lq // latent_frames
return fallback_tokens_per_frame
def mask_fn(Lq, Lk, dtype, device, transformer_options): def mask_fn(Lq, Lk, dtype, device, transformer_options):
debug_log(f"mask_fn check: Lq={Lq} Lk={Lk} max_token_idx={max_token_idx} cond_or_uncond={transformer_options.get('cond_or_uncond', [])}") debug_log(f"mask_fn check: Lq={Lq} Lk={Lk} max_token_idx={max_token_idx} cond_or_uncond={transformer_options.get('cond_or_uncond', [])}")
@@ -82,13 +98,7 @@ def create_mask_fn(q_token_idx, fallback_tokens_per_frame, latent_frames):
mode = "scaled" mode = "scaled"
video_lq = -1 video_lq = -1
else: else:
if grid_sizes is not None: video_tpf = _get_video_query_layout(Lq, grid_sizes)
video_tpf = int(grid_sizes[1]) * int(grid_sizes[2])
else:
if Lq % latent_frames == 0:
video_tpf = Lq // latent_frames
else:
video_tpf = fallback_tokens_per_frame
video_lq = latent_frames * video_tpf video_lq = latent_frames * video_tpf
# Skip cross-modal attention — text keys are padded to a fixed length ≥ max_token_idx and != video_lq # Skip cross-modal attention — text keys are padded to a fixed length ≥ max_token_idx and != video_lq
@@ -189,6 +199,34 @@ def get_raw_tokenizer(clip):
) )
def _get_tokenizer_input_ids(raw_tokenizer, text):
tokenized = raw_tokenizer(text)
if isinstance(tokenized, dict) and "input_ids" in tokenized:
return tokenized["input_ids"]
if hasattr(tokenized, "input_ids"):
return tokenized.input_ids
if isinstance(tokenized, list):
return tokenized
return []
def _tokenized_length(raw_tokenizer, text, eos_adjustment):
return max(0, len(_get_tokenizer_input_ids(raw_tokenizer, text)) - eos_adjustment)
def _tokenizer_has_eos(raw_tokenizer):
if getattr(raw_tokenizer, "add_eos", False):
return True
try:
ids = _get_tokenizer_input_ids(raw_tokenizer, "test")
if not ids:
return False
eos_id = getattr(raw_tokenizer, "eos_token_id", None)
return (eos_id is not None and ids[-1] == eos_id) or ids[-1] == 1
except Exception:
return False
def map_token_indices(raw_tokenizer, global_prompt, local_prompts): def map_token_indices(raw_tokenizer, global_prompt, local_prompts):
"""Tokenize global + space-prefixed locals; return (full_prompt, per-local token ranges). """Tokenize global + space-prefixed locals; return (full_prompt, per-local token ranges).
@@ -197,38 +235,14 @@ def map_token_indices(raw_tokenizer, global_prompt, local_prompts):
prefixed_locals = [" " + lp for lp in local_prompts] prefixed_locals = [" " + lp for lp in local_prompts]
full_prompt = global_prompt + "".join(prefixed_locals) full_prompt = global_prompt + "".join(prefixed_locals)
# Detect if the tokenizer appends EOS dynamically eos_adj = 1 if _tokenizer_has_eos(raw_tokenizer) else 0
has_eos = getattr(raw_tokenizer, "add_eos", False) prev_len = _tokenized_length(raw_tokenizer, global_prompt, eos_adj)
if not has_eos:
try:
test_res = raw_tokenizer("test")
if isinstance(test_res, dict) and "input_ids" in test_res:
ids = test_res["input_ids"]
elif hasattr(test_res, "input_ids"):
ids = test_res.input_ids
elif isinstance(test_res, list):
ids = test_res
else:
ids = []
if ids:
eos_id = getattr(raw_tokenizer, "eos_token_id", None)
if eos_id is not None and ids[-1] == eos_id:
has_eos = True
elif ids[-1] == 1:
has_eos = True
except Exception:
pass
eos_adj = 1 if has_eos else 0
prev_len = len(raw_tokenizer(global_prompt)["input_ids"]) - eos_adj
token_ranges = [] token_ranges = []
built = global_prompt built = global_prompt
for plp in prefixed_locals: for plp in prefixed_locals:
built += plp built += plp
cur_len = len(raw_tokenizer(built)["input_ids"]) - eos_adj cur_len = _tokenized_length(raw_tokenizer, built, eos_adj)
if cur_len <= prev_len: if cur_len <= prev_len:
raise ValueError(f"Local prompt produced no tokens: '{plp.strip()}'") raise ValueError(f"Local prompt produced no tokens: '{plp.strip()}'")
token_ranges.append((prev_len, cur_len)) token_ranges.append((prev_len, cur_len))