fix: 🐛 Whisper chunks processing
also add support for whisper chunks in TextToImage
This commit is contained in:
+38
-45
@@ -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
|
||||
|
||||
+185
-46
@@ -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__ = [
|
||||
|
||||
Reference in New Issue
Block a user