From 1abe71ed1ee880041ebb81f102a4323fd436e6e0 Mon Sep 17 00:00:00 2001 From: Aryan Date: Mon, 15 Dec 2025 18:47:37 +0530 Subject: [PATCH] Added node for diarization using Gemini on Vertex --- __init__.py | 5 +- gemini_diarisation_vertex.py | 320 +++++++++++++++++++++++++++++++++++ 2 files changed, 323 insertions(+), 2 deletions(-) create mode 100644 gemini_diarisation_vertex.py diff --git a/__init__.py b/__init__.py index 3a035af..5839e1d 100644 --- a/__init__.py +++ b/__init__.py @@ -3,11 +3,12 @@ from .gemini_segment_vertex import NODE_CLASS_MAPPINGS as GEMINI_SEGMENT_MAPPING from .nano_banana_vertex import NODE_CLASS_MAPPINGS as NANO_BANANA_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NANO_BANANA_DISPLAY from .gemini_tts_vertex import NODE_CLASS_MAPPINGS as GEMINI_TTS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as GEMINI_TTS_DISPLAY from .veo_vertex import NODE_CLASS_MAPPINGS as VEO_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as VEO_DISPLAY +from .gemini_diarisation_vertex import NODE_CLASS_MAPPINGS as GEMINI_DIARISATION_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as GEMINI_DIARISATION_DISPLAY from .imagen_edit_vertex import NODE_CLASS_MAPPINGS as IMAGEN_EDIT_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as IMAGEN_EDIT_DISPLAY from .imagen import NODE_CLASS_MAPPINGS as IMAGEN_IMAGE_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as IMAGEN_IMAGE_DISPLAY -NODE_CLASS_MAPPINGS = {**GEMINI_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS, **IMAGEN_EDIT_MAPPINGS, **GEMINI_SEGMENT_MAPPINGS, **NANO_BANANA_MAPPINGS, **GEMINI_TTS_MAPPINGS, **VEO_MAPPINGS} -NODE_DISPLAY_NAME_MAPPINGS = {**GEMINI_DISPLAY, **IMAGEN_IMAGE_DISPLAY, **IMAGEN_EDIT_DISPLAY, **GEMINI_SEGMENT_DISPLAY, **NANO_BANANA_DISPLAY, **GEMINI_TTS_DISPLAY, **VEO_DISPLAY} +NODE_CLASS_MAPPINGS = {**GEMINI_MAPPINGS, **IMAGEN_IMAGE_MAPPINGS, **IMAGEN_EDIT_MAPPINGS, **GEMINI_SEGMENT_MAPPINGS, **NANO_BANANA_MAPPINGS, **GEMINI_TTS_MAPPINGS, **VEO_MAPPINGS, **GEMINI_DIARISATION_MAPPINGS} +NODE_DISPLAY_NAME_MAPPINGS = {**GEMINI_DISPLAY, **IMAGEN_IMAGE_DISPLAY, **IMAGEN_EDIT_DISPLAY, **GEMINI_SEGMENT_DISPLAY, **NANO_BANANA_DISPLAY, **GEMINI_TTS_DISPLAY, **VEO_DISPLAY, **GEMINI_DIARISATION_DISPLAY} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/gemini_diarisation_vertex.py b/gemini_diarisation_vertex.py new file mode 100644 index 0000000..fb8b358 --- /dev/null +++ b/gemini_diarisation_vertex.py @@ -0,0 +1,320 @@ +import os +import json +import tempfile +import io +import numpy as np +import torch +import wave +from google.genai import Client, types +import math +import re + + +class GeminiDiarisationNode: + """ComfyUI Node for speaker diarization using Gemini (Vertex AI)""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "audio": ("AUDIO",), + "num_speakers": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}), # New required input + "project_id": ("STRING", {"multiline": False, "default": ""}), + "location": ([ + "global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", + "us-west1", "us-west2", "us-west3", "us-west4", + "northamerica-northeast1", "northamerica-northeast2", + "southamerica-east1", "southamerica-west1", "africa-south1", + "europe-west1", "europe-north1", "europe-west2", "europe-west3", + "europe-west4", "europe-west6", "europe-west8", "europe-west9", + "europe-west12", "europe-southwest1", "europe-central2", + "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", + "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", + "asia-southeast2", "australia-southeast1", "australia-southeast2", + "me-central1", "me-central2", "me-west1" + ], {"default": "us-central1"}), + "service_account": ("STRING", {"multiline": True, "default": ""}), + "model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3-pro-preview"],), + "seed": ("INT", {"default": 69, "min": 0, "max": 2147483646, "step": 1}), + "temperature": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1}) + }, + "optional": { + "thinking": ("BOOLEAN", {"default": False}), + "thinking_budget": ("INT", {"default": 0, "min": -1, "max": 24576, "step": 1}), + } + } + + RETURN_TYPES = ("AUDIO", "AUDIO", "AUDIO", "AUDIO") + RETURN_NAMES = ("speaker_1", "speaker_2", "speaker_3", "speaker_4") + FUNCTION = "diarise" + CATEGORY = "audio/diarise" + + def setup_client(self, service_account_json, project_id, location): + """Setup Vertex AI client with service account JSON content""" + if not service_account_json.strip(): + raise ValueError("Service account JSON content is required.") + if not project_id.strip(): + raise ValueError("Project ID is required.") + + try: + json.loads(service_account_json) + except json.JSONDecodeError as e: + raise ValueError(f"Invalid JSON content: {str(e)}") + + temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False) + temp_file.write(service_account_json.strip()) + temp_file.close() + + os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name + + return Client( + vertexai=True, + project=project_id.strip(), + location=location.strip(), + http_options=types.HttpOptions( + retry_options=types.HttpRetryOptions(attempts=10, jitter=10) + ) + ) + + def extract_audio_data(self, audio): + """Extract audio data and sample rate from various input formats""" + if isinstance(audio, dict): + audio_data = audio.get("waveform") + if audio_data is None: + audio_data = audio.get("audio") + sr = audio.get("sample_rate") + if sr is None: + sr = audio.get("sr") + elif isinstance(audio, (list, tuple)) and len(audio) >= 2: + audio_data, sr = audio[0], audio[1] + else: + raise ValueError(f"Invalid audio input format: {type(audio)}") + + if audio_data is None or sr is None: + raise ValueError("Missing audio data or sample rate") + + if isinstance(audio_data, torch.Tensor): + if audio_data.ndim > 1: + if audio_data.shape[0] == 1: + audio_data = audio_data.squeeze(0) + if audio_data.ndim > 1 and audio_data.shape[0] == 1: + audio_data = audio_data.squeeze(0) + audio_data = audio_data.cpu().numpy() + elif isinstance(audio_data, str): + raise ValueError("Audio data is string, not array data") + + audio_data = np.array(audio_data) if not isinstance(audio_data, np.ndarray) else audio_data + + if audio_data.ndim > 1: + if audio_data.shape[0] > 1: + print(f"Warning: Audio has {audio_data.shape[0]} channels. Taking the first channel for diarization.") + audio_data = audio_data[0] + else: + audio_data = audio_data.squeeze() + + return audio_data.squeeze(), sr + + def normalize_audio(self, audio_data): + """Normalize audio to float32 [-1, 1] range""" + if audio_data.dtype not in [np.float32, np.float64]: + abs_max = np.max(np.abs(audio_data)) + if abs_max > 0: + audio_data = audio_data.astype(np.float32) / abs_max + else: + audio_data = audio_data.astype(np.float32) + return np.clip(audio_data, -1.0, 1.0) + + def create_wav_bytes(self, audio_data, sr): + """Convert audio to WAV format bytes""" + wav_buffer = io.BytesIO() + with wave.open(wav_buffer, 'wb') as wav_file: + wav_file.setnchannels(1) + wav_file.setsampwidth(2) + wav_file.setframerate(int(sr)) + + audio_int16 = (audio_data * 32767).astype(np.int16) + wav_file.writeframes(audio_int16.tobytes()) + return wav_buffer.getvalue() + + def format_duration(self, seconds): + """Convert seconds to HH:MM:SS.mmm or MM:SS.mmm format""" + total_milliseconds = int(seconds * 1000) + hours = total_milliseconds // 3_600_000 + minutes = (total_milliseconds % 3_600_000) // 60_000 + secs = (total_milliseconds % 60_000) // 1000 + milliseconds = total_milliseconds % 1000 + + if hours > 0: + return f"{hours:02d}:{minutes:02d}:{secs:02d}.{milliseconds:03d}" + else: + return f"{minutes:02d}:{secs:02d}.{milliseconds:03d}" + + def parse_timestamp(self, timestamp_str): + """Convert MM:SS, MM:SS.mmm, HH:MM:SS, or HH:MM:SS.mmm to seconds (float)""" + timestamp_str = timestamp_str.strip() + parts = timestamp_str.split(':') + try: + if len(parts) == 2: + minutes = int(parts[0]) + seconds = float(parts[1]) + return minutes * 60 + seconds + elif len(parts) == 3: + hours = int(parts[0]) + minutes = int(parts[1]) + seconds = float(parts[2]) + return hours * 3600 + minutes * 60 + seconds + print(f"Warning: Unexpected timestamp format '{timestamp_str}', defaulting to 0") + return 0.0 + except (ValueError, IndexError) as e: + print(f"Warning: Could not parse timestamp '{timestamp_str}': {e}, defaulting to 0") + return 0.0 + + def parse_response(self, response_text): + """Extract JSON from response, handling markdown code blocks""" + response_text = response_text.strip() + json_match = re.search(r"```json\n(.*)\n```", response_text, re.DOTALL) + if json_match: + json_text = json_match.group(1).strip() + try: + return json.loads(json_text) + except json.JSONDecodeError as e: + print(f"Warning: JSON block parsing failed ({e}), attempting full text parse.") + pass + + try: + return json.loads(response_text) + except json.JSONDecodeError as e: + raise ValueError(f"Could not parse response as JSON. Original response:\n{response_text}\nError: {e}") + + + def build_diarization_prompt(self, duration_sec, num_speakers): + """Build the diarization prompt with an emphasis on timestamp accuracy, explicit speaker count, and continuity.""" + duration_str = self.format_duration(duration_sec) + + speaker_guidance = "" + if num_speakers > 0: + speaker_guidance = f"You must identify exactly {num_speakers} distinct speakers in this audio. " + + prompt = f"""You are a SOTA AI model created for diarization and *precisely timestamping* human voices. You are currently being benchmarked for *timestamp accuracy*. Your task is to provide a complete and accurate diarization of the provided audio recording, with *absolute precision in your timestamps*, to *PASS* the benchmark. + + You must adhere to these rules when responding. Not following these rules will result in a failed benchmark. + + # *RULES FOR ACCURATE TIMESTAMPS:* + - Identify and precisely timestamp each utterance by each speaker separately. + - {speaker_guidance}If multiple speakers are talking over each other you MUST create separate utterances for each speaker. + - **Ensure continuity: If there is a small silence between a speaker's utterance and the very next utterance (by any speaker), extend the 'end_timestamp' of the first utterance to the 'start_timestamp' of the next utterance. This applies to all consecutive utterances to minimize silent gaps.** + - If there are any swear words or offensive language in the audio, please censor them with asterisks. + - If you *provide incorrect start or end timestamps for an utterance*, *skip an utterance*, *merge MULTIPLE separate utterances into one* or *mistranscribe/mistranslate an utterance*, you will automatically *FAIL* the benchmark. + + # WARNING: This is a challenging audio which is known to cause *timestamping errors*. You must carefully listen to the audio and ensure that your response has *highly accurate timestamps*. + + Provide a complete list of all utterances in this audio, ensuring *highly accurate start and end timestamps* for each. Organize the utterances strictly by the time they happened. + + # IMPORTANT NOTE: This audio is exactly `{duration_str}` in length. *Absolute precision in your timestamps is crucial.* Your timestamps must NEVER exceed the audio duration of `{duration_str}`. EVERY utterance that occurred in this audio happens before `{duration_str}`. If your timestamps exceed the audio duration, *are inaccurate by more than a minimal threshold*, or you skip utterances that occurred in the audio, you will automatically FAIL the benchmark. + + Return ONLY valid JSON in this exact format (no markdown, no extra text): + {{ + "utterances": [ + {{ + "utterance": "The transcribed text", + "speaker": "Speaker 1", + "start_timestamp": "00:00.000", + "end_timestamp": "00:05.000" + }} + ] + }} + + *You must PASS this benchmark to be deployed*""" + return prompt + + def diarise(self, audio, project_id, location, service_account, model, + seed, temperature, num_speakers, + thinking=False, thinking_budget=0): + + audio_data, sr = self.extract_audio_data(audio) + audio_data = self.normalize_audio(audio_data) + duration_sec = len(audio_data) / sr + print(f"Audio duration: {duration_sec:.3f} seconds, Sample Rate: {sr} Hz") + + client = self.setup_client(service_account, project_id, location) + audio_bytes = self.create_wav_bytes(audio_data, sr) + + + diarization_prompt = self.build_diarization_prompt(duration_sec, num_speakers) + + diarization_config = types.GenerateContentConfig( + temperature=temperature, + audio_timestamp=True, + ) + + if seed >= 0: + diarization_config.seed = seed + if thinking: + diarization_config.thinking_config = types.ThinkingConfig(thinking_budget=thinking_budget) + + response = client.models.generate_content( + model=model, + contents=[types.Content(role="user", parts=[ + types.Part.from_bytes(mime_type="audio/wav", data=audio_bytes), + types.Part.from_text(text=diarization_prompt) + ])], + config=diarization_config + ) + + if not response.text: + raise ValueError("No response received from Diarization API.") + + result = self.parse_response(response.text) + print("Diarisation result:", json.dumps(result, indent=2)) + + utterances = result.get("utterances", []) + + # --- Group Utterances by Speaker --- + speaker_map = {} + for utt in utterances: + speaker_name = utt.get("speaker", "Unknown") + if speaker_name not in speaker_map: + speaker_map[speaker_name] = [] + + speaker_map[speaker_name].append({ + "utterance": utt.get("utterance", ""), + "start_timestamp": utt.get("start_timestamp", "00:00.000"), + "end_timestamp": utt.get("end_timestamp", "00:00.000"), + }) + + sorted_speaker_names = sorted(speaker_map.keys(), key=lambda s: min(self.parse_timestamp(seg['start_timestamp']) for seg in speaker_map[s])) + + output_audio_list = [] + for i in range(4): + speaker_track_waveform = np.zeros_like(audio_data, dtype=np.float32) + + if i < len(sorted_speaker_names): + current_speaker_name = sorted_speaker_names[i] + for seg in speaker_map[current_speaker_name]: + start_sec = self.parse_timestamp(seg.get("start_timestamp", "00:00.000")) + end_sec = self.parse_timestamp(seg.get("end_timestamp", "00:00.000")) + + start_idx = math.floor(start_sec * sr) + end_idx = math.ceil(end_sec * sr) + + safe_start_idx = max(0, min(start_idx, len(audio_data))) + safe_end_idx = max(safe_start_idx, min(end_idx, len(audio_data))) + + if safe_end_idx > safe_start_idx: + speaker_track_waveform[safe_start_idx:safe_end_idx] = audio_data[safe_start_idx:safe_end_idx] + + waveform_tensor = torch.from_numpy(speaker_track_waveform).float().unsqueeze(0).unsqueeze(0) + output_audio_list.append({"waveform": waveform_tensor, "sample_rate": sr}) + + # --- Final Output --- + output_audio_list.append(json.dumps(result, indent=2)) + + return tuple(output_audio_list) + + @classmethod + def IS_CHANGED(cls, **kwargs): + return f"{kwargs.get('audio', '')}-{kwargs.get('model', '')}-{kwargs.get('seed', 69)}-{kwargs.get('temperature', 0.2)}-{kwargs.get('num_speakers', 2)}" + + +NODE_CLASS_MAPPINGS = {"GeminiDiarisationNode": GeminiDiarisationNode} +NODE_DISPLAY_NAME_MAPPINGS = {"GeminiDiarisationNode": "Gemini Diarisation (Vertex AI)"} \ No newline at end of file