From ec9941b5fccc08d715fbf9ae1a61a4dd9fe7e12d Mon Sep 17 00:00:00 2001 From: facok <128763816+facok@users.noreply.github.com> Date: Sat, 21 Mar 2026 21:30:22 +0800 Subject: [PATCH] Use model's latent_format for space conversion instead of hardcoded FLUX constants Replaces hardcoded SCALE_FACTOR/SHIFT_FACTOR with model.latent_format.process_in/out so LCS works with any model (FLUX, LTXV, SD, etc). Adds LTXAV packed tensor unpack/repack support for audio+video models. All hooks (color, tone, sharpness, observer) now use the model-aware conversion from args["model"]. --- core/sampling.py | 75 +++++++++++++++++++++++++++++++++++++++++++++- nodes/intervene.py | 64 ++++++++++++++++++++++----------------- nodes/observe.py | 24 ++++++++++----- nodes/sharpen.py | 19 ++++++++---- 4 files changed, 141 insertions(+), 41 deletions(-) diff --git a/core/sampling.py b/core/sampling.py index 8e8ccfc..52e1e47 100644 --- a/core/sampling.py +++ b/core/sampling.py @@ -2,7 +2,7 @@ import torch -# FLUX VAE process_in ↔ raw space conversion constants +# Legacy FLUX constants — kept for backward compatibility with observe.py preview SCALE_FACTOR = 0.3611 SHIFT_FACTOR = 0.1159 @@ -19,3 +19,76 @@ def find_step_index(sigma, sigmas): if len(matched) > 0: return matched[0].item() return (sigmas_f - sigma_val).abs().argmin().item() + + +def denoised_to_raw(denoised, model): + """Convert denoised tensor from process_in space to raw VAE space. + + Uses the model's latent_format.process_out (inverse of process_in). + Works for any model: FLUX (scale+shift), LTXV (identity), SD (scale), etc. + """ + return model.latent_format.process_out(denoised) + + +def raw_to_denoised(raw, model): + """Convert raw VAE space tensor back to process_in space. + + Uses the model's latent_format.process_in. + """ + return model.latent_format.process_in(raw) + + +def unpack_video_if_needed(denoised, args): + """Unpack LTXAV-style packed latents if detected. + + LTXAV packs video [B,128,F,H,W] + audio [B,ch,T,freq] into [B,1,flat]. + Returns (tensor_to_process, pack_info) where pack_info is None for + non-packed formats or a dict for repacking. + """ + # Detect packed format: shape [B, 1, flat] with very large last dim + if denoised.ndim == 3 and denoised.shape[1] == 1: + # Try to find latent_shapes from cond data + cond = args.get("cond") + latent_shapes = _extract_latent_shapes(cond) + if latent_shapes is not None and len(latent_shapes) > 1: + import comfy.utils + tensors = comfy.utils.unpack_latents(denoised, latent_shapes) + # tensors[0] = video [B, 128, F, H, W], tensors[1] = audio [B, ch, T, freq] + return tensors[0], {"packed": True, "latent_shapes": latent_shapes, + "other_tensors": tensors[1:], "original": denoised} + return denoised, None + + +def repack_video_if_needed(modified, original_denoised, pack_info): + """Repack video tensor back into LTXAV packed format if it was unpacked. + + modified: the video tensor after intervention [B, 128, F, H, W] + original_denoised: the original packed tensor (for audio portion) + pack_info: from unpack_video_if_needed + """ + if pack_info is None: + return modified + import comfy.utils + all_tensors = [modified] + pack_info["other_tensors"] + packed, _ = comfy.utils.pack_latents(all_tensors) + return packed + + +def _extract_latent_shapes(cond): + """Try to extract latent_shapes from conditioning data. + + After convert_cond, cond is a list of dicts with 'model_conds' containing + CONDConstant-wrapped values like 'latent_shapes'. + """ + if cond is None: + return None + for c in cond: + if isinstance(c, dict): + model_conds = c.get('model_conds', {}) + if 'latent_shapes' in model_conds: + ls = model_conds['latent_shapes'] + # CONDConstant wraps the value in .cond + if hasattr(ls, 'cond'): + return ls.cond + return ls + return None diff --git a/nodes/intervene.py b/nodes/intervene.py index 5eb7821..dcd5561 100644 --- a/nodes/intervene.py +++ b/nodes/intervene.py @@ -6,7 +6,7 @@ from comfy_api.latest import io from ..core.lcs_data import LCSData from ..core.patchify import patchify, unpatchify -from ..core.sampling import SCALE_FACTOR, SHIFT_FACTOR, find_step_index +from ..core.sampling import find_step_index, denoised_to_raw, raw_to_denoised, unpack_video_if_needed, repack_video_if_needed from ..core.timestep import get_alpha_beta, get_alpha_beta_t50, normalize_to_t50, denormalize_from_t50 from ..core.color_space import hex_to_hsl, encode_hsl_to_lcs, decode_lcs_to_hsl, _hue_lerp @@ -23,8 +23,9 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step, """ def post_cfg_fn(args): """Post-CFG hook: project to LCS, apply color intervention, reconstruct.""" - denoised = args["denoised"] # [B, 16, H, W] in process_in space + denoised = args["denoised"] sigma = args["sigma"] + model = args["model"] # Determine current step index sigmas = args["model_options"]["transformer_options"]["sample_sigmas"] @@ -34,19 +35,22 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step, if step_index < start_step or step_index > end_step: return denoised + # Unpack LTXAV packed format if needed + working, pack_info = unpack_video_if_needed(denoised, args) + sigma_val = float(sigma.flatten()[0]) - device = denoised.device - dtype = denoised.dtype + device = working.device + dtype = working.dtype # Move LCS data to device ld = lcs_data.to(device, dtype) - B_mat = ld.basis # [64, 3] - mu = ld.mean # [64] - anchor_lcs = ld.anchor_lcs # [8, 3] - anchor_angles = ld.anchor_angles # [6] + B_mat = ld.basis + mu = ld.mean + anchor_lcs = ld.anchor_lcs + anchor_angles = ld.anchor_angles # Convert from process_in to raw VAE space - raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W] + raw = denoised_to_raw(working, model) # Patchify patches, h_len, w_len, extra_shape = patchify(raw) @@ -154,15 +158,16 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step, new_projection = denormalize_from_t50(new_c_norm, alpha_t, beta_t, alpha_50, beta_50) # Reconstruct patches - patches_new = new_projection @ B_mat.T + mu + residual # [B, L, 64] + patches_new = new_projection @ B_mat.T + mu + residual # Unpatchify raw_new = unpatchify(patches_new, h_len, w_len, extra_shape) # Convert back to process_in space - modified = (raw_new - SHIFT_FACTOR) * SCALE_FACTOR + modified = raw_to_denoised(raw_new, model).to(dtype) - return modified.to(dtype) + # Repack if LTXAV + return repack_video_if_needed(modified, denoised, pack_info) return post_cfg_fn @@ -309,8 +314,9 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature def post_cfg_fn(args): """Post-CFG hook: project to LCS, adjust contrast/brightness/saturation, reconstruct.""" - denoised = args["denoised"] # [B, 16, H, W] in process_in space + denoised = args["denoised"] sigma = args["sigma"] + model = args["model"] # Determine current step index sigmas = args["model_options"]["transformer_options"]["sample_sigmas"] @@ -319,17 +325,20 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature if step_index < start_step or step_index > end_step: return denoised + # Unpack LTXAV packed format if needed + working, pack_info = unpack_video_if_needed(denoised, args) + sigma_val = float(sigma.flatten()[0]) - device = denoised.device - dtype = denoised.dtype + device = working.device + dtype = working.dtype ld = lcs_data.to(device, dtype) - B_mat = ld.basis # [64, 3] - mu = ld.mean # [64] - anchor_lcs = ld.anchor_lcs # [8, 3] + B_mat = ld.basis + mu = ld.mean + anchor_lcs = ld.anchor_lcs # Convert from process_in to raw VAE space - raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W] + raw = denoised_to_raw(working, model) # Patchify patches, h_len, w_len, extra_shape = patchify(raw) @@ -337,11 +346,11 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature return denoised # Incompatible latent format # Project to LCS - projection = (patches - mu) @ B_mat # [B, L, 3] + projection = (patches - mu) @ B_mat - # Compute residual (61D orthogonal complement) - reconstruction = projection @ B_mat.T + mu # [B, L, 64] - residual = patches - reconstruction # [B, L, 64] + # Compute residual (orthogonal complement) + reconstruction = projection @ B_mat.T + mu + residual = patches - reconstruction # Get timestep statistics alpha_t, beta_t = get_alpha_beta(sigma_val, device=device) @@ -350,7 +359,7 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature alpha_50, beta_50 = alpha_50.to(dtype), beta_50.to(dtype) # Normalize to t=50 - c_norm = normalize_to_t50(projection, alpha_t, beta_t, alpha_50, beta_50) # [B, L, 3] + c_norm = normalize_to_t50(projection, alpha_t, beta_t, alpha_50, beta_50) # Achromatic axis: black → white in LCS anchor space black = anchor_lcs[6] # [3] @@ -406,15 +415,16 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature new_projection = denormalize_from_t50(new_c_norm, alpha_t, beta_t, alpha_50, beta_50) # Reconstruct patches - patches_new = new_projection @ B_mat.T + mu + residual # [B, L, 64] + patches_new = new_projection @ B_mat.T + mu + residual # Unpatchify raw_new = unpatchify(patches_new, h_len, w_len, extra_shape) # Convert back to process_in space - modified = (raw_new - SHIFT_FACTOR) * SCALE_FACTOR + modified = raw_to_denoised(raw_new, model).to(dtype) - return modified.to(dtype) + # Repack if LTXAV + return repack_video_if_needed(modified, denoised, pack_info) return post_cfg_fn diff --git a/nodes/observe.py b/nodes/observe.py index 8428013..bc6f7ef 100644 --- a/nodes/observe.py +++ b/nodes/observe.py @@ -12,25 +12,31 @@ from ..core.lcs_data import LCSData from ..core.patchify import patchify from ..core.timestep import get_alpha_beta, get_alpha_beta_t50, normalize_to_t50 from ..core.color_space import decode_lcs_to_hsl, hsl_to_rgb +from ..core.sampling import denoised_to_raw, unpack_video_if_needed LCS_DATA = io.Custom("LCS_DATA") -# FLUX VAE constants -SCALE_FACTOR = 0.3611 -SHIFT_FACTOR = 0.1159 +# FLUX VAE constants — fallback for LCSPreviewColors which has no model access +_FLUX_SCALE_FACTOR = 0.3611 +_FLUX_SHIFT_FACTOR = 0.1159 -def _latent_to_color_preview(samples, lcs_data, sigma, upscale=8): +def _latent_to_color_preview(samples, lcs_data, sigma, upscale=8, model=None): """Convert latent tensor to LCS color preview image. - samples: [B, 16, H, W] in process_in space + samples: [B, C, H, W] or [B, C, T, H, W] in process_in space + model: if provided, uses model.latent_format for space conversion; + otherwise falls back to FLUX constants. Returns: [B, H_up, W_up, 3] float32 in [0,1] """ device = samples.device dtype = samples.dtype ld = lcs_data.to(device, dtype) - raw = samples / SCALE_FACTOR + SHIFT_FACTOR + if model is not None: + raw = denoised_to_raw(samples, model) + else: + raw = samples / _FLUX_SCALE_FACTOR + _FLUX_SHIFT_FACTOR patches, h_len, w_len, _ = patchify(raw) if patches is None: # Incompatible latent format — return black image @@ -130,11 +136,15 @@ class LCSStepObserver(io.ComfyNode): """Post-CFG hook: generate color preview and save to temp directory.""" denoised = args["denoised"] sigma = args["sigma"] + model = args["model"] sigma_val = float(sigma.flatten()[0]) + # Unpack LTXAV packed format if needed + working, _ = unpack_video_if_needed(denoised, args) + # Generate color preview for first batch item preview = _latent_to_color_preview( - denoised[:1], lcs_data, sigma_val, upscale=4 + working[:1], lcs_data, sigma_val, upscale=4, model=model ) # Save to temp directory diff --git a/nodes/sharpen.py b/nodes/sharpen.py index 03d1bef..d9aaeae 100644 --- a/nodes/sharpen.py +++ b/nodes/sharpen.py @@ -10,7 +10,7 @@ from safetensors.torch import save_file, load_file from ..core.sharpness import SharpnessData, calibrate_sharpness from ..core.calibration import vae_fingerprint from ..core.patchify import patchify, unpatchify -from ..core.sampling import SCALE_FACTOR, SHIFT_FACTOR, find_step_index +from ..core.sampling import find_step_index, denoised_to_raw, raw_to_denoised, unpack_video_if_needed, repack_video_if_needed SHARPNESS_DATA = io.Custom("SHARPNESS_DATA") LCS_DATA = io.Custom("LCS_DATA") @@ -115,8 +115,9 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask): edit_vec = (strength * sharpness_data.sign) * pc1_dir # [64], on CPU def post_cfg_fn(args): - denoised = args["denoised"] # [B, 16, H, W] in process_in space + denoised = args["denoised"] sigma = args["sigma"] + model = args["model"] # Step gating sigmas = args["model_options"]["transformer_options"]["sample_sigmas"] @@ -125,14 +126,17 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask): if step_index < start_step or step_index > end_step: return denoised - device = denoised.device - dtype = denoised.dtype + # Unpack LTXAV packed format if needed + working, pack_info = unpack_video_if_needed(denoised, args) + + device = working.device + dtype = working.dtype # Move edit vector to device/dtype (short-circuits if already there) ev = edit_vec.to(device=device, dtype=dtype) # Convert from process_in to raw VAE space - raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W] + raw = denoised_to_raw(working, model) # Patchify patches, h_len, w_len, extra_shape = patchify(raw) @@ -152,7 +156,10 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask): raw_new = unpatchify(patches_new, h_len, w_len, extra_shape) # Convert back to process_in space - return ((raw_new - SHIFT_FACTOR) * SCALE_FACTOR).to(dtype) + modified = raw_to_denoised(raw_new, model).to(dtype) + + # Repack if LTXAV + return repack_video_if_needed(modified, denoised, pack_info) return post_cfg_fn