Remove character image analysis flow
This commit is contained in:
+33
-362
@@ -366,50 +366,32 @@ def _build_character_tag_groups(characters: list[dict]) -> list[tuple[str, ...]]
|
||||
return groups
|
||||
|
||||
|
||||
def _character_prompt_replacements(characters: list[dict]) -> dict[str, str]:
|
||||
replacements: dict[str, str] = {}
|
||||
for character, tags in zip(characters or [], _build_character_tag_groups(characters)):
|
||||
replacement = character.get("description", "") or ""
|
||||
for tag in tags:
|
||||
replacements[tag] = replacement
|
||||
return replacements
|
||||
def _build_character_reference_block(characters: list[dict] | None = None) -> str:
|
||||
lines: list[str] = []
|
||||
for idx, character in enumerate(characters or []):
|
||||
description = (character.get("description", "") or "").strip()
|
||||
if not description:
|
||||
continue
|
||||
|
||||
alias = _normalize_character_alias(character.get("alias", ""))
|
||||
canonical_tag = f"@char{idx + 1}"
|
||||
ref_label = canonical_tag if not alias else f"{canonical_tag} / @{alias}"
|
||||
lines.append(f"{ref_label} - {description}")
|
||||
|
||||
if not lines:
|
||||
return ""
|
||||
|
||||
return "Character references:\n" + "\n".join(lines)
|
||||
|
||||
|
||||
def _prepend_character_descriptions(global_prompt: str, characters: list[dict] | None = None) -> str:
|
||||
descriptions = [
|
||||
(character.get("description", "") or "").strip()
|
||||
for character in (characters or [])
|
||||
if (character.get("description", "") or "").strip()
|
||||
]
|
||||
if not descriptions:
|
||||
return global_prompt or ""
|
||||
|
||||
prefix = ". ".join(descriptions)
|
||||
if not (global_prompt or "").strip():
|
||||
return prefix
|
||||
return f"{prefix}. {global_prompt.strip()}"
|
||||
|
||||
|
||||
def _preprocess_prompts_with_characters(global_prompt, local_prompts, characters: list[dict] | None = None):
|
||||
"""Invisibly swaps out @characterN/@charN/@alias tags with their character descriptions."""
|
||||
if "@" not in (global_prompt or "") and "@" not in (local_prompts or ""):
|
||||
return global_prompt or "", local_prompts or ""
|
||||
|
||||
replacements = _character_prompt_replacements(characters or [])
|
||||
|
||||
def apply_replacements(text: str) -> str:
|
||||
updated = text or ""
|
||||
for tag, replacement in replacements.items():
|
||||
if tag in updated:
|
||||
updated = updated.replace(tag, replacement)
|
||||
return updated
|
||||
|
||||
gp = apply_replacements(global_prompt or "")
|
||||
if not local_prompts:
|
||||
return gp, ""
|
||||
|
||||
processed_locals = [apply_replacements(part.strip()) for part in local_prompts.split("|")]
|
||||
return gp, " | ".join(processed_locals)
|
||||
def _append_character_references(global_prompt: str, characters: list[dict] | None = None) -> str:
|
||||
prompt = (global_prompt or "").strip()
|
||||
reference_block = _build_character_reference_block(characters)
|
||||
if not reference_block:
|
||||
return prompt
|
||||
if not prompt:
|
||||
return reference_block
|
||||
return f"{prompt}\n\n{reference_block}"
|
||||
|
||||
|
||||
def _load_image_source(b64_or_url: str, filename: str = None, cache: dict | None = None,
|
||||
@@ -512,312 +494,6 @@ async def ltx_director_check_file(request):
|
||||
|
||||
|
||||
# --- Provider defaults shared by the analyze + unload endpoints ---
|
||||
_PROVIDER_DEFAULTS = {
|
||||
"ollama": {"url": "http://127.0.0.1:11434", "model": "huihui_ai/qwen3.5-abliterated:2b"},
|
||||
"lmstudio": {"url": "http://127.0.0.1:1234", "model": ""},
|
||||
"custom": {"url": "", "model": ""},
|
||||
}
|
||||
|
||||
_ANALYZE_SYSTEM_PROMPT = ""
|
||||
|
||||
_ANALYZE_PROMPT = (
|
||||
"Describe the character's physical appearance in two concise sentences. "
|
||||
"Specify their hair color/style, face details, and their clothing type/color. "
|
||||
"Keep the entire response very brief."
|
||||
)
|
||||
|
||||
|
||||
def _resolve_analyze_prompt(data: dict) -> str:
|
||||
prompt = (data.get("prompt") or "").strip()
|
||||
return prompt or _ANALYZE_PROMPT
|
||||
|
||||
|
||||
def _extract_ollama_generated_text(resp_json: dict) -> str:
|
||||
if not isinstance(resp_json, dict):
|
||||
return ""
|
||||
|
||||
candidates = [resp_json.get("response"), resp_json.get("content"), resp_json.get("thinking")]
|
||||
message = resp_json.get("message")
|
||||
if isinstance(message, dict):
|
||||
candidates.extend([
|
||||
message.get("content"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking"),
|
||||
])
|
||||
|
||||
cleaned_candidates = []
|
||||
for candidate in candidates:
|
||||
if isinstance(candidate, str):
|
||||
text = candidate.strip()
|
||||
if "<think>" in text:
|
||||
text = text.split("</think>")[-1].strip()
|
||||
if text:
|
||||
cleaned_candidates.append(text)
|
||||
|
||||
if not cleaned_candidates:
|
||||
return ""
|
||||
|
||||
usable = [text for text in cleaned_candidates if _analysis_text_is_usable(text)]
|
||||
if usable:
|
||||
return max(usable, key=len)
|
||||
|
||||
return max(cleaned_candidates, key=len)
|
||||
|
||||
|
||||
def _collect_ollama_text_candidates(resp_json: dict) -> list[str]:
|
||||
if not isinstance(resp_json, dict):
|
||||
return []
|
||||
|
||||
candidates = [resp_json.get("response"), resp_json.get("content"), resp_json.get("thinking")]
|
||||
message = resp_json.get("message")
|
||||
if isinstance(message, dict):
|
||||
candidates.extend([
|
||||
message.get("content"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking"),
|
||||
])
|
||||
|
||||
cleaned = []
|
||||
for candidate in candidates:
|
||||
if isinstance(candidate, str):
|
||||
text = candidate.strip()
|
||||
if "<think>" in text:
|
||||
text = text.split("</think>")[-1].strip()
|
||||
if text:
|
||||
cleaned.append(text)
|
||||
return cleaned
|
||||
|
||||
|
||||
def _analysis_text_is_usable(text: str) -> bool:
|
||||
text = (text or "").strip()
|
||||
if not text:
|
||||
return False
|
||||
if len(text) < 24:
|
||||
return False
|
||||
if len(text.split()) < 6:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _compress_analysis_image_b64(b64_payload: str, max_dim: int = 768, quality: int = 82) -> str:
|
||||
"""Shrink analysis images so multimodal providers do not burn their full context on pixels."""
|
||||
try:
|
||||
raw = base64.b64decode(_normalise_b64_payload(b64_payload))
|
||||
with Image.open(_io.BytesIO(raw)) as img:
|
||||
img = img.convert("RGB")
|
||||
w, h = img.size
|
||||
if max(w, h) > max_dim:
|
||||
scale = max_dim / float(max(w, h))
|
||||
img = img.resize((max(1, int(round(w * scale))), max(1, int(round(h * scale)))), Image.LANCZOS)
|
||||
out = _io.BytesIO()
|
||||
img.save(out, format="JPEG", quality=quality, optimize=True)
|
||||
return base64.b64encode(out.getvalue()).decode("ascii")
|
||||
except Exception:
|
||||
return _normalise_b64_payload(b64_payload)
|
||||
|
||||
|
||||
def _prepare_analysis_images(cleaned_b64_list: list[str], max_dim: int = 768, quality: int = 82) -> list[str]:
|
||||
return [
|
||||
_compress_analysis_image_b64(b64, max_dim=max_dim, quality=quality)
|
||||
for b64 in cleaned_b64_list
|
||||
]
|
||||
|
||||
|
||||
def _resolve_provider(data):
|
||||
provider = (data.get("provider") or "ollama").lower()
|
||||
defs = _PROVIDER_DEFAULTS.get(provider, _PROVIDER_DEFAULTS["ollama"])
|
||||
base_url = (data.get("base_url") or defs["url"]).rstrip("/")
|
||||
model = data.get("model") or defs["model"]
|
||||
return provider, base_url, model
|
||||
|
||||
|
||||
# --- Character reference analysis endpoint (Ollama / LM Studio / Custom OpenAI-compatible) ---
|
||||
@PromptServer.instance.routes.post("/ltx_director/analyze_character")
|
||||
async def analyze_character_endpoint(request):
|
||||
try:
|
||||
import aiohttp
|
||||
data = await request.json()
|
||||
image_b64 = data.get("image_b64", "")
|
||||
image_debug = data.get("image_debug") or []
|
||||
char_index = int(data.get("char_index", 0))
|
||||
provider, base_url, model_name = _resolve_provider(data)
|
||||
analyze_prompt = _resolve_analyze_prompt(data)
|
||||
|
||||
if provider == "off":
|
||||
return web.json_response({"status": "error", "message": "Analyze is set to Off / Manual."})
|
||||
if not image_b64:
|
||||
return web.json_response({"status": "error", "message": "No image provided for analysis."})
|
||||
|
||||
b64_list = image_b64 if isinstance(image_b64, list) else [image_b64]
|
||||
cleaned_b64_list = []
|
||||
for b64 in b64_list:
|
||||
if "," in b64:
|
||||
b64 = b64.split(",", 1)[1]
|
||||
cleaned_b64_list.append(b64)
|
||||
if not cleaned_b64_list:
|
||||
return web.json_response({"status": "error", "message": "No valid base64 images decoded."})
|
||||
if provider in ("lmstudio", "custom") and not model_name:
|
||||
return web.json_response({
|
||||
"status": "error",
|
||||
"message": f"No model name set for {provider}. Open the gear menu and enter your loaded model's name.",
|
||||
})
|
||||
|
||||
log.info("[LTXDirector] Analyzing Character %d via %s (%s, model '%s')...",
|
||||
char_index + 1, provider, base_url, model_name)
|
||||
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
if provider == "ollama":
|
||||
debug_candidates = []
|
||||
analysis_images = cleaned_b64_list
|
||||
payload = {
|
||||
"model": model_name,
|
||||
"system": _ANALYZE_SYSTEM_PROMPT,
|
||||
"prompt": analyze_prompt,
|
||||
"images": analysis_images,
|
||||
"stream": False,
|
||||
"keep_alive": 0,
|
||||
"options": {
|
||||
"temperature": 0.2,
|
||||
"num_predict": 768,
|
||||
},
|
||||
}
|
||||
async with session.post(f"{base_url}/api/generate", json=payload, timeout=300) as response:
|
||||
if response.status != 200:
|
||||
err_txt = await response.text()
|
||||
if response.status == 400 and "exceeds the available context size" in err_txt:
|
||||
analysis_images = _prepare_analysis_images(cleaned_b64_list, max_dim=512, quality=70)
|
||||
payload["images"] = analysis_images
|
||||
async with session.post(f"{base_url}/api/generate", json=payload, timeout=300) as retry_response:
|
||||
if retry_response.status != 200:
|
||||
retry_err = await retry_response.text()
|
||||
return web.json_response({"status": "error", "message": f"Ollama HTTP {retry_response.status}: {retry_err}"})
|
||||
resp_json = await retry_response.json()
|
||||
debug_candidates.extend(_collect_ollama_text_candidates(resp_json))
|
||||
generated_text = _extract_ollama_generated_text(resp_json)
|
||||
else:
|
||||
return web.json_response({"status": "error", "message": f"Ollama HTTP {response.status}: {err_txt}"})
|
||||
else:
|
||||
resp_json = await response.json()
|
||||
debug_candidates.extend(_collect_ollama_text_candidates(resp_json))
|
||||
generated_text = _extract_ollama_generated_text(resp_json)
|
||||
|
||||
if not _analysis_text_is_usable(generated_text):
|
||||
chat_payload = {
|
||||
"model": model_name,
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": _ANALYZE_SYSTEM_PROMPT,
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": analyze_prompt,
|
||||
"images": analysis_images,
|
||||
},
|
||||
],
|
||||
"stream": False,
|
||||
"keep_alive": 0,
|
||||
"options": {
|
||||
"temperature": 0.2,
|
||||
"num_predict": 768,
|
||||
},
|
||||
}
|
||||
async with session.post(f"{base_url}/api/chat", json=chat_payload, timeout=300) as response:
|
||||
if response.status == 200:
|
||||
resp_json = await response.json()
|
||||
debug_candidates.extend(_collect_ollama_text_candidates(resp_json))
|
||||
chat_text = _extract_ollama_generated_text(resp_json)
|
||||
if _analysis_text_is_usable(chat_text):
|
||||
generated_text = chat_text
|
||||
if not _analysis_text_is_usable(generated_text):
|
||||
return web.json_response({
|
||||
"status": "error",
|
||||
"message": "Ollama returned only a truncated analysis response.",
|
||||
"debug_candidates": debug_candidates,
|
||||
"image_debug": image_debug,
|
||||
"image_lengths": [len(b64) for b64 in cleaned_b64_list],
|
||||
"description": generated_text,
|
||||
})
|
||||
else:
|
||||
# OpenAI-compatible vision chat (LM Studio / Custom).
|
||||
content = [{"type": "text", "text": analyze_prompt}]
|
||||
for b64 in cleaned_b64_list:
|
||||
content.append({"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{b64}"}})
|
||||
payload = {
|
||||
"model": model_name,
|
||||
"messages": [
|
||||
{"role": "system", "content": _ANALYZE_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": content},
|
||||
],
|
||||
"max_tokens": 4096, "stream": False,
|
||||
}
|
||||
async with session.post(f"{base_url}/v1/chat/completions", json=payload, timeout=120) as response:
|
||||
if response.status != 200:
|
||||
err_txt = await response.text()
|
||||
return web.json_response({"status": "error", "message": f"{provider} HTTP {response.status}: {err_txt}"})
|
||||
resp_json = await response.json()
|
||||
try:
|
||||
msg = resp_json["choices"][0]["message"]
|
||||
generated_text = (msg.get("content") or "").strip()
|
||||
# Reasoning/"thinking" models (Gemma, Qwen-thinking, etc.) may leave
|
||||
# content empty and put their output in reasoning_content instead.
|
||||
if not generated_text:
|
||||
generated_text = (msg.get("reasoning_content") or "").strip()
|
||||
except (KeyError, IndexError, TypeError):
|
||||
return web.json_response({"status": "error", "message": f"Unexpected response shape from {provider}."})
|
||||
except aiohttp.ClientConnectorError:
|
||||
return web.json_response({
|
||||
"status": "error",
|
||||
"message": f"Could not connect to {provider} at {base_url}. Make sure the server is running and reachable.",
|
||||
})
|
||||
|
||||
if "<think>" in generated_text:
|
||||
generated_text = generated_text.split("</think>")[-1].strip()
|
||||
|
||||
log.info("[LTXDirector] Analysis complete: %s", generated_text)
|
||||
return web.json_response({"status": "success", "description": generated_text})
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"[LTXDirector] Failed to analyze character: {e}")
|
||||
return web.json_response({"status": "error", "message": str(e)}, status=500)
|
||||
|
||||
|
||||
|
||||
@PromptServer.instance.routes.post("/ltx_director/unload_ollama")
|
||||
async def unload_ollama_endpoint(request):
|
||||
"""Evict the analysis model from VRAM right before a generation run, so the VLM doesn't
|
||||
compete with LTX for VRAM (the cause of intermittent CUDA offload crashes).
|
||||
|
||||
Ollama supports a clean instant unload (keep_alive=0). LM Studio / Custom have no reliable
|
||||
cross-version HTTP unload, so for those this is a graceful no-op — users should set a short
|
||||
JIT / auto-unload TTL in their server instead. Fully tolerant: never raises into the run.
|
||||
"""
|
||||
try:
|
||||
import aiohttp
|
||||
try:
|
||||
data = await request.json()
|
||||
except Exception:
|
||||
data = {}
|
||||
provider, base_url, model_name = _resolve_provider(data)
|
||||
|
||||
if provider == "ollama":
|
||||
payload = {"model": model_name, "keep_alive": 0}
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(f"{base_url}/api/generate", json=payload, timeout=8) as response:
|
||||
await response.text()
|
||||
log.info("[LTXDirector] Asked Ollama to release '%s' from VRAM before the run.", model_name)
|
||||
except Exception:
|
||||
pass # Ollama not running / unreachable -> nothing to free.
|
||||
return web.json_response({"status": "ok", "provider": provider})
|
||||
|
||||
# LM Studio / Custom: no reliable HTTP unload across versions -> graceful no-op.
|
||||
return web.json_response({"status": "ok", "provider": provider, "note": "no-op (set a JIT/TTL unload in your server)"})
|
||||
except Exception as e:
|
||||
return web.json_response({"status": "error", "message": str(e)})
|
||||
|
||||
|
||||
def read_wav_peaks(wav_path):
|
||||
import wave
|
||||
with wave.open(wav_path, 'rb') as w:
|
||||
@@ -1635,6 +1311,8 @@ def _build_reference_mode_outputs(
|
||||
vae,
|
||||
global_prompt: str,
|
||||
local_prompts: str,
|
||||
raw_global_prompt: str,
|
||||
raw_local_prompts: str,
|
||||
segment_lengths: str,
|
||||
duration_frames: int,
|
||||
epsilon: float,
|
||||
@@ -1667,7 +1345,7 @@ def _build_reference_mode_outputs(
|
||||
tensor = _resize_image(tensor, latent_w, latent_h, "stretch to fit", divisible_by)
|
||||
return tensor
|
||||
|
||||
prompt_text = (global_prompt or "") + " " + (local_prompts or "")
|
||||
prompt_text = (raw_global_prompt or "") + " " + (raw_local_prompts or "")
|
||||
tag_groups = _build_character_tag_groups(characters)
|
||||
referenced_slots = [i for i, tags in enumerate(tag_groups) if any(t in prompt_text for t in tags)]
|
||||
selected = []
|
||||
@@ -2103,21 +1781,12 @@ class LTXDirector(io.ComfyNode):
|
||||
image_cache=image_cache,
|
||||
input_dir=input_dir,
|
||||
)
|
||||
global_prompt = _prepend_character_descriptions(global_prompt, characters)
|
||||
raw_global_prompt = global_prompt or ""
|
||||
raw_local_prompts = local_prompts or ""
|
||||
global_prompt = _append_character_references(raw_global_prompt, characters)
|
||||
char_slot_images = [list(character.get("images") or []) for character in characters]
|
||||
char_images = [img for slot_images in char_slot_images for img in slot_images]
|
||||
|
||||
# --- @charN substitution ---
|
||||
# OFF: swap @char tags for their VLM descriptions in the prompt text.
|
||||
# Licon MSR: leave the tags raw — there the reference IMAGE drives identity, and the
|
||||
# tags are used only to pick which character slots feed the slideshow.
|
||||
if reference_mode == "Licon MSR (Prefix)":
|
||||
ref_global, ref_local = global_prompt, local_prompts
|
||||
else:
|
||||
ref_global, ref_local = _preprocess_prompts_with_characters(
|
||||
global_prompt, local_prompts, characters
|
||||
)
|
||||
global_prompt, local_prompts = ref_global, ref_local
|
||||
local_prompts = raw_local_prompts
|
||||
|
||||
guide_data, derived_w, derived_h = _build_guide_data_from_timeline(
|
||||
tdata=tdata,
|
||||
@@ -2168,6 +1837,8 @@ class LTXDirector(io.ComfyNode):
|
||||
vae=vae,
|
||||
global_prompt=global_prompt,
|
||||
local_prompts=local_prompts,
|
||||
raw_global_prompt=raw_global_prompt,
|
||||
raw_local_prompts=raw_local_prompts,
|
||||
segment_lengths=segment_lengths,
|
||||
duration_frames=duration_frames,
|
||||
epsilon=epsilon,
|
||||
|
||||
Reference in New Issue
Block a user