Files
pmarmotte2-Comfyui-Sequenti…/video_sequentially_node.py
T

822 lines
30 KiB
Python

import json
import inspect
import math
import os
import subprocess
import tempfile
import threading
import wave
import gc
from typing import Dict, Tuple
import cv2
import numpy as np
import torch
from pyannote.audio import Pipeline
from pyannote.audio.core import task as pyannote_task
import whisper
import folder_paths
from comfy.utils import ProgressBar
VIDEO_EXTENSIONS = ("mp4", "mov", "mkv", "webm", "avi", "m4v")
_STATE_LOCK = threading.Lock()
_STATE_FILE = os.path.join(os.path.dirname(__file__), "turn_state.json")
LLM_MODELS_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "models", "LLM"))
LLM_EMPTY_MSG = "Download any LLM to the models_LLM folder"
TRANSLATION_LANGUAGES = [
"none",
"English",
"French",
"Spanish",
"German",
"Italian",
"Portuguese",
"Dutch",
"Russian",
"Arabic",
"Hindi",
"Japanese",
"Korean",
"Chinese (Simplified)",
]
_TRANSLATION_LOCK = threading.Lock()
_TRANSLATION_CACHE = {"model_id": None, "backend": None, "tokenizer": None, "processor": None, "model": None}
def _log_progress(msg: str) -> None:
print(f"[LoadVideoSequentially] {msg}", flush=True)
def _get_llm_model_choices():
if not os.path.isdir(LLM_MODELS_DIR):
return [LLM_EMPTY_MSG]
choices = []
for root, _, files in os.walk(LLM_MODELS_DIR):
if "config.json" in files:
rel_dir = os.path.relpath(root, LLM_MODELS_DIR).replace("\\", "/")
choices.append(rel_dir)
return sorted(set(choices)) if choices else [LLM_EMPTY_MSG]
def _resolve_llm_model_id(selection: str) -> str:
if not selection or selection == LLM_EMPTY_MSG:
return ""
full = os.path.join(LLM_MODELS_DIR, selection.replace("/", os.sep))
return full if os.path.exists(full) else selection
def _load_translation_model(model_id: str):
from transformers import AutoConfig, AutoModelForCausalLM, AutoModelForImageTextToText, AutoTokenizer, AutoProcessor
with _TRANSLATION_LOCK:
if _TRANSLATION_CACHE["model_id"] == model_id and _TRANSLATION_CACHE["model"] is not None:
return (
_TRANSLATION_CACHE["backend"],
_TRANSLATION_CACHE["tokenizer"],
_TRANSLATION_CACHE["processor"],
_TRANSLATION_CACHE["model"],
)
cfg = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
model_type = str(getattr(cfg, "model_type", "")).lower()
_log_progress(f"Loading translation model: {model_id} (type={model_type})")
dtype = torch.float16 if torch.cuda.is_available() else torch.float32
if "vl" in model_type:
backend = "vl"
processor = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)
tokenizer = None
model = AutoModelForImageTextToText.from_pretrained(
model_id, dtype=dtype, device_map="auto", trust_remote_code=True
)
else:
backend = "causal"
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
processor = None
model = AutoModelForCausalLM.from_pretrained(
model_id, dtype=dtype, device_map="auto", trust_remote_code=True
)
_TRANSLATION_CACHE["model_id"] = model_id
_TRANSLATION_CACHE["backend"] = backend
_TRANSLATION_CACHE["tokenizer"] = tokenizer
_TRANSLATION_CACHE["processor"] = processor
_TRANSLATION_CACHE["model"] = model
return backend, tokenizer, processor, model
def _offload_translation_model() -> None:
with _TRANSLATION_LOCK:
model = _TRANSLATION_CACHE.get("model")
if model is not None:
try:
model.to("cpu")
except Exception:
pass
_TRANSLATION_CACHE["model_id"] = None
_TRANSLATION_CACHE["backend"] = None
_TRANSLATION_CACHE["tokenizer"] = None
_TRANSLATION_CACHE["processor"] = None
_TRANSLATION_CACHE["model"] = None
gc.collect()
if torch.cuda.is_available():
try:
torch.cuda.empty_cache()
if hasattr(torch.cuda, "ipc_collect"):
torch.cuda.ipc_collect()
except Exception:
pass
def _translate_utterance_text(text: str, target_language: str, backend: str, tokenizer, processor, model) -> str:
prompt = (
f"Translate the following utterance to {target_language}. "
"Return only the translation, no metadata, no speaker label, no timestamp.\n\n"
f"Utterance:\n{text}"
)
if backend == "vl":
messages = [{"role": "user", "content": [{"type": "text", "text": prompt}]}]
model_input_text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = processor(text=[model_input_text], return_tensors="pt").to(model.device)
with torch.no_grad():
output = model.generate(**inputs, max_new_tokens=256, do_sample=False, temperature=0.0)
prompt_len = inputs["input_ids"].shape[1]
translated = processor.batch_decode(output[:, prompt_len:], skip_special_tokens=True)[0].strip()
else:
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
with torch.no_grad():
output = model.generate(
**inputs,
max_new_tokens=256,
do_sample=False,
temperature=0.0,
pad_token_id=tokenizer.eos_token_id,
)
translated = tokenizer.decode(output[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip()
return translated if translated else text
def _translate_dialogue_text(dialogue_text: str, target_language: str, llm_model: str) -> str:
if not dialogue_text.strip() or target_language == "none":
return dialogue_text
model_id = _resolve_llm_model_id(llm_model)
if not model_id:
return dialogue_text
try:
backend, tokenizer, processor, model = _load_translation_model(model_id)
out_lines = []
lines = dialogue_text.splitlines()
_log_progress(f"Translating {len(lines)} line(s) to {target_language}")
for line in lines:
if ": " not in line:
out_lines.append(line)
continue
# Preserve timeline + speaker label exactly, translate only utterance text.
prefix, utterance = line.split(": ", 1)
translated = _translate_utterance_text(utterance.strip(), target_language, backend, tokenizer, processor, model)
out_lines.append(f"{prefix}: {translated}")
return "\n".join(out_lines)
finally:
_log_progress("Offloading translation model from GPU/VRAM")
_offload_translation_model()
def _resolve_ffmpeg() -> str:
try:
import imageio_ffmpeg
return imageio_ffmpeg.get_ffmpeg_exe()
except Exception:
return "ffmpeg"
FFMPEG_BIN = _resolve_ffmpeg()
def _resolve_ffprobe() -> str:
try:
import imageio_ffmpeg
ffmpeg_exe = imageio_ffmpeg.get_ffmpeg_exe()
ffmpeg_dir = os.path.dirname(ffmpeg_exe)
ffprobe_name = "ffprobe.exe" if os.name == "nt" else "ffprobe"
ffprobe_path = os.path.join(ffmpeg_dir, ffprobe_name)
if os.path.isfile(ffprobe_path):
return ffprobe_path
except Exception:
pass
return "ffprobe"
FFPROBE_BIN = _resolve_ffprobe()
def _load_state() -> Dict[str, int]:
if not os.path.exists(_STATE_FILE):
return {}
try:
with open(_STATE_FILE, "r", encoding="utf-8") as f:
data = json.load(f)
if isinstance(data, dict):
return {str(k): int(v) for k, v in data.items()}
except Exception:
pass
return {}
def _save_state(state: Dict[str, int]) -> None:
tmp_file = _STATE_FILE + ".tmp"
with open(tmp_file, "w", encoding="utf-8") as f:
json.dump(state, f, indent=2)
os.replace(tmp_file, _STATE_FILE)
def _next_turn(unique_id: str, initial_turn: int, reset_sequence: bool) -> int:
key = str(unique_id)
with _STATE_LOCK:
state = _load_state()
# Always honor the turn provided by the widget for this execution.
# This ensures manual edits take effect immediately.
current = max(1, int(initial_turn))
state[key] = current + 1
_save_state(state)
return current
def _extract_audio_segment(video_path: str, start_sec: float, duration_sec: float) -> Dict[str, torch.Tensor]:
probe_cmd = [
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"default=nokey=1:noprint_wrappers=1",
video_path,
]
probe = subprocess.run(probe_cmd, capture_output=True)
if probe.returncode != 0:
sample_rate = 44100
channels = 2
else:
lines = probe.stdout.decode("utf-8", errors="ignore").strip().splitlines()
if len(lines) >= 2:
sample_rate = int(lines[0])
channels = int(lines[1])
else:
sample_rate = 44100
channels = 2
if channels <= 0:
channels = 1
cmd = [
FFMPEG_BIN,
"-v",
"error",
"-ss",
str(start_sec),
"-i",
video_path,
"-t",
str(duration_sec),
"-vn",
"-ac",
str(channels),
"-ar",
str(sample_rate),
"-f",
"f32le",
"-acodec",
"pcm_f32le",
"-",
]
proc = subprocess.run(cmd, capture_output=True)
if proc.returncode != 0:
stderr = proc.stderr.decode("utf-8", errors="ignore")
raise RuntimeError(f"ffmpeg audio extraction failed: {stderr}")
audio_np = np.frombuffer(proc.stdout, dtype=np.float32)
if audio_np.size == 0:
waveform = torch.zeros((1, max(1, channels), 0), dtype=torch.float32)
return {"waveform": waveform, "sample_rate": sample_rate}
frame_count = audio_np.size // channels
# Copy avoids non-writable NumPy buffer warning when converting to torch tensor.
audio_np = audio_np[: frame_count * channels].copy().reshape(frame_count, channels)
waveform = torch.from_numpy(audio_np).transpose(0, 1).unsqueeze(0).contiguous()
return {"waveform": waveform, "sample_rate": sample_rate}
def _to_mono_float32_np(waveform: torch.Tensor) -> np.ndarray:
arr = waveform.detach().cpu().numpy() if isinstance(waveform, torch.Tensor) else np.asarray(waveform)
arr = np.squeeze(arr)
if arr.ndim == 1:
mono = arr
elif arr.ndim == 2:
if arr.shape[0] <= 8 and arr.shape[1] > arr.shape[0]:
mono = arr.mean(axis=0)
else:
mono = arr.mean(axis=1)
else:
raise ValueError(f"Unsupported waveform shape: {arr.shape}")
mono = np.asarray(mono, dtype=np.float32)
if mono.size == 0:
return mono
max_abs = np.max(np.abs(mono))
if max_abs > 1.0:
mono = mono / 32768.0
return np.clip(mono, -1.0, 1.0)
def _is_effectively_silent(audio: np.ndarray, rms_threshold: float = 1e-4) -> bool:
if audio.size == 0:
return True
rms = float(np.sqrt(np.mean(np.square(audio), dtype=np.float64)))
return rms < rms_threshold
def _write_wav(path: str, mono_audio: np.ndarray, sample_rate: int) -> None:
pcm16 = (mono_audio * 32767.0).astype(np.int16)
with wave.open(path, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2)
wf.setframerate(sample_rate)
wf.writeframes(pcm16.tobytes())
def _configure_torch_safe_globals() -> None:
serialization = getattr(torch, "serialization", None)
torch_version_mod = getattr(torch, "torch_version", None)
torch_version_cls = getattr(torch_version_mod, "TorchVersion", None)
add_safe_globals = getattr(serialization, "add_safe_globals", None)
if add_safe_globals is None:
return
safe = []
if torch_version_cls is not None:
safe.append(torch_version_cls)
for name in dir(pyannote_task):
value = getattr(pyannote_task, name, None)
if isinstance(value, type) and getattr(value, "__module__", "") == pyannote_task.__name__:
safe.append(value)
if safe:
add_safe_globals(safe)
def _load_pipeline_with_inspect_guard(token: str):
original_stack = inspect.stack
def _safe_stack(context=0):
return []
inspect.stack = _safe_stack
try:
hf_kwargs = {}
if token:
params = inspect.signature(Pipeline.from_pretrained).parameters
if "use_auth_token" in params:
hf_kwargs["use_auth_token"] = token
elif "token" in params:
hf_kwargs["token"] = token
elif "auth_token" in params:
hf_kwargs["auth_token"] = token
else:
os.environ.setdefault("HUGGINGFACE_HUB_TOKEN", token)
return Pipeline.from_pretrained("pyannote/speaker-diarization-3.1", **hf_kwargs)
finally:
inspect.stack = original_stack
def _format_mmss(seconds: float) -> str:
total = max(0, int(seconds))
mins = total // 60
secs = total % 60
return f"{mins}:{secs:02d}"
def _pick_speaker_for_segment(seg_start: float, seg_end: float, diar_turns):
best_speaker, best_overlap = None, 0.0
best_start, best_end = None, None
for start, end, speaker in diar_turns:
overlap = max(0.0, min(seg_end, end) - max(seg_start, start))
if overlap > best_overlap:
best_overlap, best_speaker = overlap, speaker
best_start, best_end = start, end
return best_speaker, best_overlap, best_start, best_end
def _pick_speaker_for_timepoint(t: float, diar_turns):
for start, end, speaker in diar_turns:
if start <= t <= end:
return speaker
# Fallback to nearest diarized turn if exact containment is not found.
if not diar_turns:
return None
return min(diar_turns, key=lambda x: min(abs(t - x[0]), abs(t - x[1])))[2]
def _iter_diarization_tracks_compat(diarization_obj):
candidate = diarization_obj
if not hasattr(candidate, "itertracks"):
candidate = getattr(diarization_obj, "speaker_diarization", candidate)
if not hasattr(candidate, "itertracks") and isinstance(diarization_obj, dict):
candidate = diarization_obj.get("speaker_diarization", diarization_obj)
if not hasattr(candidate, "itertracks"):
raise TypeError(f"Unsupported diarization output type: {type(diarization_obj).__name__}")
for turn, _, speaker in candidate.itertracks(yield_label=True):
yield turn, speaker
def _run_transcription_only(tmp_wav_path: str, whisper_model: str, default_label: str, merge_consecutive_speaker: bool) -> str:
_log_progress(f"Whisper transcription started (model={whisper_model}, mode=transcription-only)")
device = "cuda" if torch.cuda.is_available() else "cpu"
asr_model = whisper.load_model(whisper_model, device=device)
asr = asr_model.transcribe(
tmp_wav_path,
verbose=False,
condition_on_previous_text=False,
no_speech_threshold=0.7,
logprob_threshold=-1.0,
)
asr_segments = asr.get("segments", []) or []
entries = []
for segment in asr_segments:
seg_start = float(segment.get("start", 0.0))
text = (segment.get("text") or "").strip()
if text:
entries.append((seg_start, f"{default_label} says", text))
if merge_consecutive_speaker and entries:
merged = []
cur_start, cur_speaker, cur_text = entries[0]
for seg_start, speaker_label, text in entries[1:]:
if speaker_label == cur_speaker:
cur_text = f"{cur_text} {text}".strip()
else:
merged.append((cur_start, cur_speaker, cur_text))
cur_start, cur_speaker, cur_text = seg_start, speaker_label, text
merged.append((cur_start, cur_speaker, cur_text))
entries = merged
return "\n".join([f"{_format_mmss(seg_start)} {speaker_label}: {text}" for seg_start, speaker_label, text in entries])
def _run_diarization_text(audio: Dict[str, torch.Tensor], hf_token: str, whisper_model: str, merge_consecutive_speaker: bool, speaker_labels: Dict[str, str]) -> str:
_log_progress("Preparing audio for diarization/transcription")
token = (hf_token or "").strip()
waveform = audio.get("waveform")
sample_rate = int(audio.get("sample_rate", 0))
mono_audio = _to_mono_float32_np(waveform)
if sample_rate <= 0 or _is_effectively_silent(mono_audio):
return ""
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
tmp_wav_path = tmp.name
try:
_write_wav(tmp_wav_path, mono_audio, sample_rate)
default_label = speaker_labels.get("A", "Speaker A")
if not token:
_log_progress("HF token missing, switching to Whisper-only transcription")
return _run_transcription_only(tmp_wav_path, whisper_model, default_label, merge_consecutive_speaker)
_configure_torch_safe_globals()
try:
_log_progress("Loading pyannote speaker diarization pipeline")
pipeline = _load_pipeline_with_inspect_guard(token)
# Help pyannote avoid collapsing to a single speaker when multiple
# speakers are expected from configured positions.
configured_positions = [speaker_labels.get(k, "") for k in ("A", "B", "C", "D")]
expected_speakers = sum(
1
for p in configured_positions
if isinstance(p, str) and p.strip() and "not set" not in p.strip().lower()
)
diar_kwargs = {}
if expected_speakers >= 2:
diar_kwargs = {"min_speakers": 2, "max_speakers": min(4, expected_speakers)}
_log_progress(
f"Running speaker diarization with constraints min={diar_kwargs['min_speakers']} max={diar_kwargs['max_speakers']}"
)
else:
_log_progress("Running speaker diarization (auto speaker count)")
_log_progress("Running speaker diarization")
diar_input = {"waveform": torch.from_numpy(mono_audio).unsqueeze(0), "sample_rate": sample_rate}
try:
diarization = pipeline(diar_input, **diar_kwargs) if diar_kwargs else pipeline(diar_input)
except TypeError:
# Compatibility fallback for pyannote versions that don't accept
# these kwargs in __call__.
diarization = pipeline(diar_input)
except Exception:
_log_progress("Diarization failed, falling back to Whisper-only transcription")
return _run_transcription_only(tmp_wav_path, whisper_model, default_label, merge_consecutive_speaker)
diar_turns = []
speaker_order = []
for turn, speaker in _iter_diarization_tracks_compat(diarization):
speaker = str(speaker)
diar_turns.append((float(turn.start), float(turn.end), speaker))
if speaker not in speaker_order:
speaker_order.append(speaker)
if not diar_turns:
_log_progress("No speaker segments found, falling back to Whisper-only transcription")
return _run_transcription_only(tmp_wav_path, whisper_model, default_label, merge_consecutive_speaker)
speaker_to_letter = {speaker: chr(ord("A") + idx) for idx, speaker in enumerate(speaker_order)}
_log_progress(f"Diarization complete ({len(speaker_order)} speaker(s)); starting Whisper transcription (model={whisper_model})")
device = "cuda" if torch.cuda.is_available() else "cpu"
asr_model = whisper.load_model(whisper_model, device=device)
asr = asr_model.transcribe(
tmp_wav_path,
verbose=False,
condition_on_previous_text=False,
no_speech_threshold=0.7,
logprob_threshold=-1.0,
)
asr_segments = asr.get("segments", []) or []
entries = []
for segment in asr_segments:
seg_start = float(segment.get("start", 0.0))
seg_end = float(segment.get("end", seg_start))
seg_duration = max(0.0, seg_end - seg_start)
text = (segment.get("text") or "").strip()
if not text:
continue
if float(segment.get("no_speech_prob", 0.0)) >= 0.6:
continue
if float(segment.get("avg_logprob", 0.0)) <= -1.2:
continue
if seg_duration < 0.35:
continue
speaker, overlap, spk_start, spk_end = _pick_speaker_for_segment(seg_start, seg_end, diar_turns)
if speaker is None or overlap < 0.2:
continue
letter = speaker_to_letter.get(speaker, "?")
mapped = speaker_labels.get(letter, f"Speaker {letter}")
# Whisper segment start can be early; anchor timestamp to actual diarized speech start.
ts = seg_start
if spk_start is not None:
ts = max(seg_start, float(spk_start))
entries.append((ts, f"{mapped} says", text))
if not entries:
_log_progress("No speaker-matched lines; using Whisper-only transcription fallback")
return _run_transcription_only(tmp_wav_path, whisper_model, default_label, merge_consecutive_speaker)
if merge_consecutive_speaker and entries:
merged = []
cur_start, cur_speaker, cur_text = entries[0]
for seg_start, speaker_label, text in entries[1:]:
if speaker_label == cur_speaker:
cur_text = f"{cur_text} {text}".strip()
else:
merged.append((cur_start, cur_speaker, cur_text))
cur_start, cur_speaker, cur_text = seg_start, speaker_label, text
merged.append((cur_start, cur_speaker, cur_text))
entries = merged
out = "\n".join([f"{_format_mmss(seg_start)} {speaker_label}: {text}" for seg_start, speaker_label, text in entries])
_log_progress(f"Transcript ready ({len(entries)} line(s))")
return out
finally:
try:
os.remove(tmp_wav_path)
except OSError:
pass
def _load_video_frames_segment(
video_path: str, start_sec: float, end_sec: float, select_every_nth: int
) -> Tuple[torch.Tensor, float, int, float]:
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise ValueError(f"Could not open video: {video_path}")
fps = cap.get(cv2.CAP_PROP_FPS)
if fps <= 0:
fps = 24.0
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
duration = (total_frames / fps) if total_frames > 0 else 0.0
start_frame = max(0, int(math.floor(start_sec * fps)))
end_frame = max(start_frame, int(math.floor(end_sec * fps)))
if total_frames > 0:
end_frame = min(end_frame, total_frames)
cap.set(cv2.CAP_PROP_POS_FRAMES, start_frame)
requested_frames = max(0, end_frame - start_frame)
frame_width = max(1, int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)))
frame_height = max(1, int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)))
select_every_nth = max(1, int(select_every_nth))
estimated_kept = max(1, math.ceil(requested_frames / select_every_nth)) if requested_frames > 0 else 1
def _frame_iter():
read_frames = 0
while read_frames < requested_frames:
ok, frame = cap.read()
if not ok:
break
if read_frames % select_every_nth == 0:
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frame = frame.astype(np.float32) / 255.0
yield frame
read_frames += 1
frames_np = np.fromiter(
_frame_iter(),
dtype=np.dtype((np.float32, (frame_height, frame_width, 3))),
count=estimated_kept,
)
cap.release()
if len(frames_np) == 0:
raise ValueError(
f"No frames available for requested segment [{start_sec:.3f}, {end_sec:.3f}] in video of duration {duration:.3f}s"
)
return torch.from_numpy(frames_np), fps, total_frames, duration
class LoadVideoSequentially:
@classmethod
def INPUT_TYPES(cls):
input_dir = folder_paths.get_input_directory()
files = []
if os.path.isdir(input_dir):
for f in os.listdir(input_dir):
full = os.path.join(input_dir, f)
if not os.path.isfile(full):
continue
ext = os.path.splitext(f)[1].lower().lstrip(".")
if ext in VIDEO_EXTENSIONS:
files.append(f)
return {
"required": {
"video": (sorted(files),),
"duration_seconds": ("FLOAT", {"default": 10.0, "min": 0.1, "max": 1000000.0, "step": 0.1}),
"select_every_nth": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1}),
"hf_token": ("STRING", {"multiline": False, "default": ""}),
"whisper_model": (
["tiny", "base", "small", "medium", "large", "turbo"],
{"default": "small"},
),
"merge_consecutive_speaker": ("BOOLEAN", {"default": True}),
"continue_without_hf_token": ("BOOLEAN", {"default": False}),
"translation_language": (TRANSLATION_LANGUAGES, {"default": "none"}),
"llm_model": (_get_llm_model_choices(),),
"turn": ("INT", {"default": 1, "min": 1, "max": 1000000, "step": 1}),
"reset_sequence": ("BOOLEAN", {"default": False}),
"speaker_a_position": ("STRING", {"default": "Speaker A not set"}),
"speaker_b_position": ("STRING", {"default": "Speaker B not set"}),
"speaker_c_position": ("STRING", {"default": "Speaker C not set"}),
"speaker_d_position": ("STRING", {"default": "Speaker D not set"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("IMAGE", "AUDIO", "INT", "FLOAT", "FLOAT", "FLOAT", "STRING", "STRING", "STRING", "STRING", "STRING")
RETURN_NAMES = (
"images",
"audio",
"turn_used",
"start_second",
"end_second",
"frame_rate",
"speaker_a_position",
"speaker_b_position",
"speaker_c_position",
"speaker_d_position",
"dialogue_text",
)
FUNCTION = "load_video_sequentially"
CATEGORY = "video"
def load_video_sequentially(
self,
video,
duration_seconds,
select_every_nth,
hf_token,
whisper_model,
merge_consecutive_speaker,
continue_without_hf_token,
translation_language,
llm_model,
turn,
reset_sequence,
speaker_a_position,
speaker_b_position,
speaker_c_position,
speaker_d_position,
unique_id,
):
pbar = ProgressBar(6)
step = 0
pbar.update_absolute(step, 6)
video_path = folder_paths.get_annotated_filepath(video)
current_turn = _next_turn(unique_id=unique_id, initial_turn=turn, reset_sequence=reset_sequence)
step += 1
pbar.update_absolute(step, 6)
start_second = float((current_turn - 1) * duration_seconds)
end_second = float(start_second + duration_seconds)
images, fps, total_frames, video_duration = _load_video_frames_segment(
video_path, start_second, end_second, int(select_every_nth)
)
step += 1
pbar.update_absolute(step, 6)
_log_progress(f"Segment loaded: turn={current_turn}, range={start_second:.2f}s-{end_second:.2f}s")
audio = _extract_audio_segment(video_path, start_second, duration_seconds)
step += 1
pbar.update_absolute(step, 6)
speaker_labels = {
"A": speaker_a_position,
"B": speaker_b_position,
"C": speaker_c_position,
"D": speaker_d_position,
}
if not str(hf_token or "").strip() and not bool(continue_without_hf_token):
raise RuntimeError(
"HF Token is not present see node readme for instructions. "
"Set 'continue_without_hf_token' to True to proceed without diarization."
)
dialogue_text = _run_diarization_text(audio, hf_token, whisper_model, bool(merge_consecutive_speaker), speaker_labels)
if translation_language != "none":
try:
dialogue_text = _translate_dialogue_text(dialogue_text, translation_language, llm_model)
except Exception as e:
_log_progress(f"Translation skipped due to error: {e}")
step += 1
pbar.update_absolute(step, 6)
preview = {
"filename": video,
"subfolder": "",
"type": "input",
"format": "video/mp4",
"frame_rate": fps,
"fullpath": video_path,
# VHS-style preview params (consumed by /viewvideo-capable preview widgets).
"start_time": start_second,
"frame_load_cap": max(1, int(round(duration_seconds * fps))),
"force_rate": fps,
}
step += 1
pbar.update_absolute(step, 6)
result = (
images,
audio,
current_turn,
start_second,
end_second,
float(fps),
speaker_a_position,
speaker_b_position,
speaker_c_position,
speaker_d_position,
dialogue_text,
)
step += 1
pbar.update_absolute(step, 6)
return {
"ui": {
"gifs": [preview],
"turn_used": [current_turn],
"next_turn": [current_turn + 1],
},
"result": result,
}
@classmethod
def IS_CHANGED(cls, **kwargs):
return float("nan")
@classmethod
def VALIDATE_INPUTS(cls, video, **kwargs):
try:
video_path = folder_paths.get_annotated_filepath(video)
except Exception:
return f"Invalid video: {video}"
if not os.path.isfile(video_path):
return f"Video not found: {video}"
return True
NODE_CLASS_MAPPINGS = {
"LoadVideoSequentially": LoadVideoSequentially,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadVideoSequentially": "Load Video Sequentially",
}