822 lines
30 KiB
Python
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",
|
|
}
|