import logging import math import os import torch 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): """Gaussian penalty matrix [Lq, Lk] for video cross-attention (integer frame indexing).""" offset = torch.zeros(Lq, Lk, device=device, dtype=dtype) query_frames = (torch.arange(Lq, device=device, dtype=torch.float32) / float(tokens_per_frame)).floor() for seg in q_token_idx: local, cost = _build_segment_cost(query_frames, seg, offset.dtype) offset[:, local.to(device=device)] = cost return offset def build_temporal_cost_scaled(q_token_idx, Lq, Lk, device, dtype, latent_frames, is_audio=False): """Penalty matrix for queries that don't map to integer frames (e.g. LTXAV audio tokens).""" offset = torch.zeros(Lq, Lk, device=device, dtype=dtype) query_frames = torch.arange(Lq, device=device, dtype=torch.float32) * latent_frames / Lq for seg in q_token_idx: if is_audio: active_seg = { "local_token_idx": seg["local_token_idx"], "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: active_seg = seg local, cost = _build_segment_cost(query_frames, active_seg, offset.dtype) offset[:, local.to(device=device)] = cost return offset def debug_log(msg): if not _DEBUG_PROMPT_RELAY: return try: with open(_DEBUG_LOG_PATH, "a", encoding="utf-8") as f: f.write(msg + "\n") except Exception: pass def create_mask_fn(q_token_idx, fallback_tokens_per_frame, latent_frames): """Closure: mask_fn(Lq, Lk, dtype, device, transformer_options) -> additive mask or None. Takes shapes/dtype/device instead of tensors so callers can compute the mask without first materializing q/k projections — required so PromptRelay can wrap an existing cross-attn forward (e.g. KJNodes NAG) instead of replacing it. """ cache = {} 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): 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', [])}") if Lq == Lk: debug_log("mask_fn: Lq == Lk, returning None") return None # Only apply on conditional pass — not unconditional (negative prompt) cond_or_uncond = transformer_options.get("cond_or_uncond", []) if 1 in cond_or_uncond and 0 not in cond_or_uncond: debug_log("mask_fn: unconditional pass, returning None") return None grid_sizes = transformer_options.get("grid_sizes", None) attn_type = transformer_options.get("promptrelay_attn_type", "attn2") is_audio = (attn_type == "audio_attn2") if is_audio: mode = "scaled" video_lq = -1 else: video_tpf = _get_video_query_layout(Lq, grid_sizes) video_lq = latent_frames * video_tpf # Skip cross-modal attention — text keys are padded to a fixed length ≥ max_token_idx and != video_lq if Lk == video_lq or Lk < max_token_idx: debug_log(f"mask_fn: Lk == video_lq ({Lk == video_lq}) or Lk < max_token_idx ({Lk < max_token_idx}), returning None") return None mode = "video" if Lq == video_lq else "scaled" key = (Lq, Lk, mode, device) if key not in cache: if mode == "video": cost = build_temporal_cost(q_token_idx, Lq, Lk, device, dtype, video_tpf) else: cost = build_temporal_cost_scaled(q_token_idx, Lq, Lk, device, dtype, latent_frames, is_audio=is_audio) log.info( "[PromptRelay] Built penalty matrix (%s): Lq=%d, Lk=%d, nonzero=%d/%d", mode, Lq, Lk, (cost > 0).sum().item(), cost.numel(), ) debug_log(f"Built penalty matrix ({mode}): Lq={Lq}, Lk={Lk}, key={key}") cache[key] = -cost return cache[key].to(dtype) return mask_fn def build_segments(token_ranges, segment_lengths, epsilon=1e-3, relay_options=None): """Per-segment metadata for the temporal penalty. relay_options (optional dict) overrides per-stream knobs: video_strength, video_window_scale, audio_epsilon, audio_strength, audio_window_scale Audio knobs only affect architectures whose cross-attention takes the scaled (non-integer-frame) path — currently LTX audio_attn2. """ # Paper uses constant sigma = 1/ln(1/epsilon) regardless of segment length sigma = 1.0 / math.log(1.0 / epsilon) if 0 < epsilon < 1 else 0.1448 opts = relay_options or {} v_strength = opts.get("video_strength", 1.0) v_window_scale = opts.get("video_window_scale", 1.0) a_epsilon = opts.get("audio_epsilon") a_strength = opts.get("audio_strength", 1.0) a_window_scale = opts.get("audio_window_scale", 1.0) if a_epsilon is not None and 0 < a_epsilon < 1: sigma_audio = 1.0 / math.log(1.0 / a_epsilon) else: sigma_audio = sigma if relay_options: log.info( "[PromptRelay] Advanced options active — video: strength=%.3f window_scale=%.3f | " "audio: epsilon=%s strength=%.3f window_scale=%.3f", v_strength, v_window_scale, f"{a_epsilon:.4f}" if a_epsilon is not None else "inherit", a_strength, a_window_scale, ) q_token_idx = [] frame_cursor = 0 for (tok_start, tok_end), L in zip(token_ranges, segment_lengths): if L <= 0: frame_cursor += L continue midpoint = (2 * frame_cursor + L) // 2 base_window = max(L // 2 - 2, 0) q_token_idx.append({ "local_token_idx": torch.arange(tok_start, tok_end), "midpoint": midpoint, "window": max(base_window * v_window_scale, 0.0), "sigma": sigma, "strength": v_strength, "window_audio": max(base_window * a_window_scale, 0.0), "sigma_audio": sigma_audio, "strength_audio": a_strength, }) frame_cursor += L return q_token_idx def get_raw_tokenizer(clip): """Extract the raw SPiece/HF tokenizer from a ComfyUI CLIP object.""" tokenizer_wrapper = clip.tokenizer for attr_name in dir(tokenizer_wrapper): if attr_name.startswith("_"): continue inner = getattr(tokenizer_wrapper, attr_name, None) if inner is not None and hasattr(inner, "tokenizer"): return inner.tokenizer raise RuntimeError( f"Could not find raw tokenizer on CLIP object. " f"Known attributes: {[a for a in dir(tokenizer_wrapper) if not a.startswith('_')]}" ) 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): """Tokenize global + space-prefixed locals; return (full_prompt, per-local token ranges). Uses incremental tokenization to avoid SentencePiece context-dependency issues. """ prefixed_locals = [" " + lp for lp in local_prompts] full_prompt = global_prompt + "".join(prefixed_locals) eos_adj = 1 if _tokenizer_has_eos(raw_tokenizer) else 0 prev_len = _tokenized_length(raw_tokenizer, global_prompt, eos_adj) token_ranges = [] built = global_prompt for plp in prefixed_locals: built += plp cur_len = _tokenized_length(raw_tokenizer, built, eos_adj) if cur_len <= prev_len: raise ValueError(f"Local prompt produced no tokens: '{plp.strip()}'") token_ranges.append((prev_len, cur_len)) prev_len = cur_len return full_prompt, token_ranges def distribute_segment_lengths(num_segments, latent_frames, specified_lengths=None): """Validate or auto-distribute segment frame counts, capped to fit within latent_frames.""" if specified_lengths: if len(specified_lengths) != num_segments: raise ValueError( f"Number of segment_lengths ({len(specified_lengths)}) " f"must match number of local prompts ({num_segments})" ) lengths = specified_lengths else: # ceil division — matches reference implementation step = -(-latent_frames // num_segments) lengths = [step] * num_segments effective = [] cursor = 0 for L in lengths: end = min(cursor + L, latent_frames) effective.append(max(end - cursor, 0)) cursor = end return effective