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"].
This commit is contained in:
+74
-1
@@ -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
|
||||
|
||||
+37
-27
@@ -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
|
||||
|
||||
|
||||
+17
-7
@@ -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
|
||||
|
||||
+13
-6
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user