Add GeminiMultimodalLiveLLMService support for VAD Params

This commit is contained in:
Mark Backman
2025-04-11 11:11:34 -04:00
parent 121e6d2157
commit e28c199dd1
2 changed files with 105 additions and 24 deletions

View File

@@ -8,6 +8,7 @@
import base64 import base64
import io import io
import json import json
from enum import Enum
from typing import List, Literal, Optional from typing import List, Literal, Optional
from PIL import Image from PIL import Image
@@ -35,6 +36,38 @@ class Turn(BaseModel):
parts: List[ContentPart] parts: List[ContentPart]
class StartSensitivity(str, Enum):
"""Determines how start of speech is detected."""
UNSPECIFIED = "START_SENSITIVITY_UNSPECIFIED" # Default is HIGH
HIGH = "START_SENSITIVITY_HIGH" # Detect start of speech more often
LOW = "START_SENSITIVITY_LOW" # Detect start of speech less often
class EndSensitivity(str, Enum):
"""Determines how end of speech is detected."""
UNSPECIFIED = "END_SENSITIVITY_UNSPECIFIED" # Default is HIGH
HIGH = "END_SENSITIVITY_HIGH" # End speech more often
LOW = "END_SENSITIVITY_LOW" # End speech less often
class AutomaticActivityDetection(BaseModel):
"""Configures automatic detection of activity."""
disabled: Optional[bool] = None
start_of_speech_sensitivity: Optional[StartSensitivity] = None
prefix_padding_ms: Optional[int] = None
end_of_speech_sensitivity: Optional[EndSensitivity] = None
silence_duration_ms: Optional[int] = None
class RealtimeInputConfig(BaseModel):
"""Configures the realtime input behavior."""
automatic_activity_detection: Optional[AutomaticActivityDetection] = None
class RealtimeInput(BaseModel): class RealtimeInput(BaseModel):
mediaChunks: List[MediaChunk] mediaChunks: List[MediaChunk]
@@ -88,6 +121,7 @@ class Setup(BaseModel):
tools: Optional[List[dict]] = None tools: Optional[List[dict]] = None
generation_config: Optional[dict] = None generation_config: Optional[dict] = None
output_audio_transcription: Optional[AudioTranscriptionConfig] = None output_audio_transcription: Optional[AudioTranscriptionConfig] = None
realtime_input_config: Optional[RealtimeInputConfig] = None
class Config(BaseModel): class Config(BaseModel):

View File

@@ -245,6 +245,16 @@ class GeminiMediaResolution(str, Enum):
HIGH = "MEDIA_RESOLUTION_HIGH" # Zoomed reframing with 256 tokens HIGH = "MEDIA_RESOLUTION_HIGH" # Zoomed reframing with 256 tokens
class GeminiVADParams(BaseModel):
"""Voice Activity Detection parameters."""
disabled: Optional[bool] = Field(default=None)
start_sensitivity: Optional[events.StartSensitivity] = Field(default=None)
end_sensitivity: Optional[events.EndSensitivity] = Field(default=None)
prefix_padding_ms: Optional[int] = Field(default=None)
silence_duration_ms: Optional[int] = Field(default=None)
class InputParams(BaseModel): class InputParams(BaseModel):
frequency_penalty: Optional[float] = Field(default=None, ge=0.0, le=2.0) frequency_penalty: Optional[float] = Field(default=None, ge=0.0, le=2.0)
max_tokens: Optional[int] = Field(default=4096, ge=1) max_tokens: Optional[int] = Field(default=4096, ge=1)
@@ -259,6 +269,7 @@ class InputParams(BaseModel):
media_resolution: Optional[GeminiMediaResolution] = Field( media_resolution: Optional[GeminiMediaResolution] = Field(
default=GeminiMediaResolution.UNSPECIFIED default=GeminiMediaResolution.UNSPECIFIED
) )
vad: Optional[GeminiVADParams] = Field(default=None)
extra: Optional[Dict[str, Any]] = Field(default_factory=dict) extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
@@ -321,6 +332,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
self._language_code = ( self._language_code = (
language_to_gemini_language(params.language) if params.language else "en-US" language_to_gemini_language(params.language) if params.language else "en-US"
) )
self._vad_params = params.vad
self._settings = { self._settings = {
"frequency_penalty": params.frequency_penalty, "frequency_penalty": params.frequency_penalty,
@@ -332,6 +344,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
"modalities": params.modalities, "modalities": params.modalities,
"language": self._language_code, "language": self._language_code,
"media_resolution": params.media_resolution, "media_resolution": params.media_resolution,
"vad": params.vad,
"extra": params.extra if isinstance(params.extra, dict) else {}, "extra": params.extra if isinstance(params.extra, dict) else {},
} }
@@ -513,31 +526,61 @@ class GeminiMultimodalLiveLLMService(LLMService):
self._websocket = await websockets.connect(uri=uri) self._websocket = await websockets.connect(uri=uri)
self._receive_task = self.create_task(self._receive_task_handler()) self._receive_task = self.create_task(self._receive_task_handler())
self._transcribe_audio_task = self.create_task(self._transcribe_audio_handler()) self._transcribe_audio_task = self.create_task(self._transcribe_audio_handler())
config = events.Config.model_validate(
{
"setup": {
"model": self._model_name,
"generation_config": {
"frequency_penalty": self._settings["frequency_penalty"],
"max_output_tokens": self._settings["max_tokens"], # Not supported yet
"presence_penalty": self._settings["presence_penalty"],
"temperature": self._settings["temperature"],
"top_k": self._settings["top_k"],
"top_p": self._settings["top_p"],
"response_modalities": self._settings["modalities"].value,
"speech_config": {
"voice_config": {
"prebuilt_voice_config": {"voice_name": self._voice_id}
},
"language_code": self._settings["language"],
},
"media_resolution": self._settings["media_resolution"].value,
},
"output_audio_transcription": {},
},
}
)
# Create the basic configuration
config_data = {
"setup": {
"model": self._model_name,
"generation_config": {
"frequency_penalty": self._settings["frequency_penalty"],
"max_output_tokens": self._settings["max_tokens"],
"presence_penalty": self._settings["presence_penalty"],
"temperature": self._settings["temperature"],
"top_k": self._settings["top_k"],
"top_p": self._settings["top_p"],
"response_modalities": self._settings["modalities"].value,
"speech_config": {
"voice_config": {
"prebuilt_voice_config": {"voice_name": self._voice_id}
},
"language_code": self._settings["language"],
},
"media_resolution": self._settings["media_resolution"].value,
},
"output_audio_transcription": {},
}
}
# Add VAD configuration if provided
if self._settings.get("vad"):
vad_config = {}
vad_params = self._settings["vad"]
# Only add parameters that are explicitly set
if vad_params.disabled is not None:
vad_config["disabled"] = vad_params.disabled
if vad_params.start_sensitivity:
vad_config["start_of_speech_sensitivity"] = vad_params.start_sensitivity.value
if vad_params.end_sensitivity:
vad_config["end_of_speech_sensitivity"] = vad_params.end_sensitivity.value
if vad_params.prefix_padding_ms is not None:
vad_config["prefix_padding_ms"] = vad_params.prefix_padding_ms
if vad_params.silence_duration_ms is not None:
vad_config["silence_duration_ms"] = vad_params.silence_duration_ms
# Only add automatic_activity_detection if we have VAD settings
if vad_config:
realtime_config = {"automatic_activity_detection": vad_config}
config_data["setup"]["realtime_input_config"] = realtime_config
config = events.Config.model_validate(config_data)
# Add system instruction if available
system_instruction = self._system_instruction or "" system_instruction = self._system_instruction or ""
if self._context and hasattr(self._context, "extract_system_instructions"): if self._context and hasattr(self._context, "extract_system_instructions"):
system_instruction += "\n" + self._context.extract_system_instructions() system_instruction += "\n" + self._context.extract_system_instructions()
@@ -546,9 +589,13 @@ class GeminiMultimodalLiveLLMService(LLMService):
config.setup.system_instruction = events.SystemInstruction( config.setup.system_instruction = events.SystemInstruction(
parts=[events.ContentPart(text=system_instruction)] parts=[events.ContentPart(text=system_instruction)]
) )
# Add tools if available
if self._tools: if self._tools:
logger.debug(f"Gemini is configuring to use tools{self._tools}") logger.debug(f"Gemini is configuring to use tools{self._tools}")
config.setup.tools = self.get_llm_adapter().from_standard_tools(self._tools) config.setup.tools = self.get_llm_adapter().from_standard_tools(self._tools)
# Send the configuration
await self.send_client_event(config) await self.send_client_event(config)
except Exception as e: except Exception as e: