diff --git a/nodes/audio.py b/nodes/audio.py index 738c179..4677929 100644 --- a/nodes/audio.py +++ b/nodes/audio.py @@ -4,6 +4,10 @@ import torch import torchaudio from comfy.model_management import get_torch_device from huggingface_hub import snapshot_download +from transformers import ( + WhisperForConditionalGeneration, + WhisperProcessor, +) # from transformers import ( # AutoFeatureExtractor, @@ -50,9 +54,13 @@ class MtbAudio: @staticmethod def resample(audio: AudioTensor, common_sample_rate: int) -> AudioTensor: - if audio["sample_rate"] != common_sample_rate: + current_rate = audio["sample_rate"] + if current_rate != common_sample_rate: + log.debug( + f"Resampling audio from {current_rate} to {common_sample_rate}" + ) resampler = torchaudio.transforms.Resample( - orig_freq=audio["sample_rate"], new_freq=common_sample_rate + orig_freq=current_rate, new_freq=common_sample_rate ) return { "sample_rate": common_sample_rate, @@ -93,8 +101,8 @@ class MtbAudio: class WhisperPipeline(TypedDict): """Whisper model pipeline.""" - processor: Any # WhisperProcessor - model: Any # WhisperForConditionalGeneration + processor: WhisperProcessor + model: WhisperForConditionalGeneration class MTB_LoadWhisper: @@ -140,11 +148,6 @@ class MTB_LoadWhisper: def load(self, model_size="tiny", download_missing=False): """Load Whisper model and processor.""" - from transformers import ( - WhisperForConditionalGeneration, - WhisperProcessor, - ) - whisper_dir = get_model_path("whisper") tag = f"whisper-{model_size}" model_dir = whisper_dir / tag @@ -182,7 +185,7 @@ class MTB_LoadWhisper: return ({"processor": processor, "model": model},) -class MTB_AudioToText: +class MTB_AudioToText(MtbAudio): """Transcribe audio to text using Whisper.""" @classmethod @@ -219,13 +222,19 @@ class MTB_AudioToText: CATEGORY = "mtb/audio" def transcribe( - self, pipeline, audio, language="auto", return_timestamps=True + self, + pipeline: WhisperPipeline, + audio: AudioTensor, + language="auto", + return_timestamps=True, ): """Transcribe audio to text using Whisper.""" processor = pipeline["processor"] model = pipeline["model"] device = model.device + audio = self.resample(audio, WHISPER_SAMPLE_RATE) + waveform = audio["waveform"] log.debug(f"Processed waveform shape: {waveform.shape}") @@ -254,6 +263,9 @@ class MTB_AudioToText: all_text = [] chunk_offsets = [] + last_time = 0.0 + accumulated_offset = 0.0 + for chunk_start in range(0, total_samples, chunk_samples): chunk_end = min(chunk_start + chunk_samples, total_samples) chunk_waveform = waveform[chunk_start:chunk_end] @@ -298,25 +310,13 @@ class MTB_AudioToText: if time_str.replace(".", "").isdigit(): time_val = float(time_str) - if len(chunk_offsets) > 1: - if time_val >= 30.0: - current_chunk = ( - chunk_offsets.index(chunk_offset) + 1 - ) - current_offset = chunk_offsets[ - current_chunk - ] - time_val = time_val - 30.0 + # If this timestamp is less than the last one, we've started a new sequence + if time_val < last_time: + accumulated_offset += last_time - adjusted_time = chunk_offset + time_val - - if 0 <= adjusted_time <= total_duration: - adjusted_tokens.append( - f"<|{adjusted_time:.2f}|>" - ) - log.debug( - f"Token {len(all_tokens)}: {time_val} -> {adjusted_time}" - ) + adjusted_time = time_val + accumulated_offset + adjusted_tokens.append(f"<|{adjusted_time:.2f}|>") + last_time = time_val else: adjusted_tokens.append(token) except ValueError: @@ -325,7 +325,6 @@ class MTB_AudioToText: adjusted_tokens.append(token) all_tokens.extend(adjusted_tokens) - chunk_text = processor.batch_decode( predicted_ids, skip_special_tokens=True )[0] @@ -334,6 +333,7 @@ class MTB_AudioToText: detected_language = "en" if language == "auto": try: + log.debug("Detecting language") with torch.no_grad(): first_chunk_features = processor( waveform[:chunk_samples], @@ -341,21 +341,13 @@ class MTB_AudioToText: return_tensors="pt", ).input_features.to(device) - predicted_language = model.detect_language( - first_chunk_features - )[0] - detected_language = predicted_language.argmax(-1).item() - lang_tokens = [ - t[2:-2] # Remove <| and |> - for t in processor.tokenizer.vocab.keys() - if t.startswith("<|") - and t.endswith("|>") - and len(t) == 6 - ] - if detected_language < len(lang_tokens): - detected_language = lang_tokens[detected_language] - else: - detected_language = "en" + predicted_probs = model.detect_language(first_chunk_features)[0] + language_token = processor.tokenizer.convert_ids_to_tokens( + predicted_probs.argmax(-1).item() + ) + detected_language = language_token[2:-2] if language_token.startswith("<|") else "en" + log.debug(f"Detected language: {detected_language}") + except Exception as e: log.warning(f"Language detection failed: {e}") @@ -721,6 +713,7 @@ class MTB_AudioIsolateSpeaker(MtbAudio): for chunk in whisper_data["chunks"]: if not chunk.get("speakers"): continue + speaker_present = target_speaker in chunk["speakers"] if (mode == "isolate" and speaker_present) or ( mode == "mute" and not speaker_present diff --git a/nodes/generate.py b/nodes/generate.py index 8598e3e..ad46e48 100644 --- a/nodes/generate.py +++ b/nodes/generate.py @@ -1,4 +1,8 @@ -from PIL import Image +import io + +import requests +import torch +from PIL import Image, ImageDraw, ImageFont from ..log import log from ..utils import comfy_dir, font_path, pil2tensor @@ -81,10 +85,6 @@ class MTB_UnsplashImage: CATEGORY = "mtb/generate" def do_unsplash_image(self, width, height, random_seed, keyword=None): - import io - - import requests - base_url = "https://source.unsplash.com/random/" if width and height: @@ -213,14 +213,90 @@ by default it fallsback to a default font. "INT", {"default": 100, "min": 1, "max": 100, "step": 1}, ), - } + }, + "optional": { + "whisper_chunks": ("WHISPER_CHUNKS",), + "fps": ( + "INT", + {"default": 24, "min": 1, "max": 60, "step": 1}, + ), + "fade_duration": ( + "FLOAT", + {"default": 0.5, "min": 0.0, "max": 5.0, "step": 0.1}, + ), + }, } RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("image",) FUNCTION = "text_to_image" CATEGORY = "mtb/generate" + def create_animation_frames( + self, + chunks, + base_image, + font, + font_size, + color, + width, + height, + fps, + fade_duration, + ): + """Create animation frames from Whisper chunks.""" + if not chunks or not chunks.get("chunks"): + return [base_image] + + frames = [] + total_duration = chunks["chunks"][-1]["timestamp"][1] + frame_count = int(total_duration * fps) + fade_frames = int(fade_duration * fps) + + for frame_idx in range(frame_count): + time = frame_idx / fps + frame = base_image.copy() + draw = ImageDraw.Draw(frame) + + active_chunks = [] + for chunk in chunks["chunks"]: + start, end = chunk["timestamp"] + if start <= time <= end: + fade_in_alpha = min( + 1.0, (time - start) * fps / fade_frames + ) + fade_out_alpha = min(1.0, (end - time) * fps / fade_frames) + alpha = min(fade_in_alpha, fade_out_alpha) + active_chunks.append((chunk["text"], alpha)) + + y = height // 4 + for text, alpha in active_chunks: + # Create a temporary image for the text with alpha + text_img = Image.new("RGBA", (width, height), (0, 0, 0, 0)) + text_draw = ImageDraw.Draw(text_img) + + text_draw.text( + (width // 2, y), + text, + font=font, + fill=color, + anchor="mm", + ) + + text_img.putalpha( + Image.fromarray( + (torch.ones((height, width)) * (alpha * 255)) + .byte() + .numpy() + ) + ) + + frame = Image.alpha_composite(frame, text_img) + y += font_size * 1.5 + + frames.append(frame) + + return frames + def text_to_image( self, text: str, @@ -238,58 +314,121 @@ by default it fallsback to a default font. h_offset=0, v_offset=0, h_coverage=100, + whisper_chunks=None, + fps=24, + fade_duration=0.5, ): + """Convert text to image, with optional animation support.""" import textwrap - from PIL import Image, ImageDraw, ImageFont + from PIL import ImageColor font_path = self.fonts[font] - - text = ( - text.encode("ascii", "ignore").decode().strip() if trim else text - ) - # Handle word wrapping - if wrap: - wrap_width = (((width / 100) * h_coverage) / font_size) * 2 - lines = textwrap.wrap(text, width=wrap_width) - else: - lines = [text] font = ImageFont.truetype(font_path, size=font_size) - log.debug(f"Lines: {lines}") - img = Image.new("RGBA", (width, height), background) - draw = ImageDraw.Draw(img) - line_height_px = line_height * font_size + try: + if isinstance(color, str): + color = ImageColor.getrgb(color) + if isinstance(background, str): + background = ImageColor.getrgb(background) - # Vertical alignment - if v_align == "top": - y_text = v_offset - elif v_align == "center": - y_text = ((height - (line_height_px * len(lines))) // 2) + v_offset - else: # bottom - y_text = (height - (line_height_px * len(lines))) - v_offset + if len(color) == 3: + color = color + (255,) + if len(background) == 3: + background = background + (255,) + except ValueError as e: + log.error(f"Color parsing error: {e}") + color = (255, 255, 255, 255) + background = (0, 0, 0, 255) - def get_width(line): - if hasattr(font, "getsize"): - return font.getsize(line)[0] + def render_text(text_to_render, alpha=None): + if trim: + text_to_render = ( + text_to_render.encode("ascii", "ignore").decode().strip() + ) + if wrap: + wrap_width = (((width / 100) * h_coverage) / font_size) * 2 + lines = textwrap.wrap(text_to_render, width=wrap_width) else: - return font.getlength(line) + lines = [text_to_render] - # Draw each line of text - for line in lines: - line_width = get_width(line) - # Horizontal alignment - if h_align == "left": - x_text = h_offset - elif h_align == "center": - x_text = ((width - line_width) // 2) + h_offset - else: # right - x_text = (width - line_width) - h_offset + img = Image.new("RGBA", (width, height), (0, 0, 0, 0)) + draw = ImageDraw.Draw(img) - draw.text((x_text, y_text), line, fill=color, font=font) - y_text += line_height_px + line_height_px = line_height * font_size - return (pil2tensor(img),) + if v_align == "top": + y_text = v_offset + elif v_align == "center": + y_text = ( + (height - (line_height_px * len(lines))) // 2 + ) + v_offset + else: + y_text = (height - (line_height_px * len(lines))) - v_offset + + def get_width(line): + if hasattr(font, "getsize"): + return font.getsize(line)[0] + else: + return font.getlength(line) + + for line in lines: + line_width = get_width(line) + if h_align == "left": + x_text = h_offset + elif h_align == "center": + x_text = ((width - line_width) // 2) + h_offset + else: + x_text = (width - line_width) - h_offset + + text_color = color + if alpha is not None: + text_color = tuple( + list(color[:3]) + [int(alpha * color[3])] + ) + + draw.text((x_text, y_text), line, fill=text_color, font=font) + y_text += line_height_px + + return img + + base_img = Image.new("RGBA", (width, height), background) + + if whisper_chunks and whisper_chunks.get("chunks"): + frames = [] + total_duration = whisper_chunks["chunks"][-1]["timestamp"][1] + frame_count = int(total_duration * fps) + fade_frames = int(fade_duration * fps) + + for frame_idx in range(frame_count): + time = frame_idx / fps + frame = base_img.copy() + + active_chunks = [] + for chunk in whisper_chunks["chunks"]: + start, end = chunk["timestamp"] + if start <= time <= end: + fade_in_alpha = min( + 1.0, (time - start) * fps / fade_frames + ) + fade_out_alpha = min( + 1.0, (end - time) * fps / fade_frames + ) + alpha = min(fade_in_alpha, fade_out_alpha) + active_chunks.append((chunk["text"], alpha)) + + for chunk_text, alpha in active_chunks: + chunk_img = render_text(chunk_text, alpha) + frame = Image.alpha_composite(frame, chunk_img) + + frames.append(frame) + + frame_tensors = [pil2tensor(frame) for frame in frames] + return (torch.cat(frame_tensors, dim=0),) + else: + text_img = render_text(text) + result = Image.alpha_composite(base_img, text_img) + return (pil2tensor(result),) __nodes__ = [