Merge pull request #157 from pipecat-ai/khk-improved-wake-word
Improved wake word filter
This commit is contained in:
10
CHANGELOG.md
10
CHANGELOG.md
@@ -7,6 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- Added `WakeCheckFilter` which allows you to pass information downstream only
|
||||||
|
if you say a certain phrase/word.
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- `Filter` has been renamed to `FrameFilter` and it's now under
|
||||||
|
`processors/filters`.
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
- Re-add exponential smoothing after volume calculation. This makes sure the
|
- Re-add exponential smoothing after volume calculation. This makes sure the
|
||||||
|
|||||||
@@ -12,14 +12,7 @@ import sys
|
|||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import Frame, ImageRawFrame, SpriteFrame
|
||||||
Frame,
|
|
||||||
SystemFrame,
|
|
||||||
TextFrame,
|
|
||||||
ImageRawFrame,
|
|
||||||
SpriteFrame,
|
|
||||||
TranscriptionFrame,
|
|
||||||
)
|
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
@@ -27,6 +20,7 @@ from pipecat.processors.aggregators.llm_context import (
|
|||||||
LLMUserContextAggregator,
|
LLMUserContextAggregator,
|
||||||
LLMAssistantContextAggregator,
|
LLMAssistantContextAggregator,
|
||||||
)
|
)
|
||||||
|
from pipecat.processors.filters.wake_check_filter import WakeCheckFilter
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
||||||
@@ -84,33 +78,6 @@ thinking_list = [
|
|||||||
thinking_frame = SpriteFrame(thinking_list)
|
thinking_frame = SpriteFrame(thinking_list)
|
||||||
|
|
||||||
|
|
||||||
class NameCheckFilter(FrameProcessor):
|
|
||||||
def __init__(self, names: list[str]):
|
|
||||||
super().__init__()
|
|
||||||
self._names = names
|
|
||||||
self._sentence = ""
|
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
||||||
if isinstance(frame, SystemFrame):
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
return
|
|
||||||
|
|
||||||
content: str = ""
|
|
||||||
|
|
||||||
# TODO: split up transcription by participant
|
|
||||||
if isinstance(frame, TranscriptionFrame):
|
|
||||||
content = frame.text
|
|
||||||
self._sentence += content
|
|
||||||
if self._sentence.endswith((".", "?", "!")):
|
|
||||||
if any(name in self._sentence for name in self._names):
|
|
||||||
await self.push_frame(TextFrame(self._sentence))
|
|
||||||
self._sentence = ""
|
|
||||||
else:
|
|
||||||
self._sentence = ""
|
|
||||||
else:
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
|
|
||||||
|
|
||||||
class ImageSyncAggregator(FrameProcessor):
|
class ImageSyncAggregator(FrameProcessor):
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
@@ -155,17 +122,17 @@ async def main(room_url: str, token):
|
|||||||
|
|
||||||
tma_in = LLMUserContextAggregator(messages)
|
tma_in = LLMUserContextAggregator(messages)
|
||||||
tma_out = LLMAssistantContextAggregator(messages)
|
tma_out = LLMAssistantContextAggregator(messages)
|
||||||
ncf = NameCheckFilter(["Santa Cat", "Santa"])
|
wcf = WakeCheckFilter(["Santa Cat", "Santa"])
|
||||||
|
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline([
|
||||||
transport.input(),
|
transport.input(), # Transport user input
|
||||||
isa,
|
isa, # Cat talking/quiet images
|
||||||
ncf,
|
wcf, # Filter out speech not directed at Santa Cat
|
||||||
tma_in,
|
tma_in, # User responses
|
||||||
llm,
|
llm, # LLM
|
||||||
tts,
|
tts, # TTS
|
||||||
transport.output(),
|
transport.output(), # Transport bot output
|
||||||
tma_out
|
tma_out # Santa Cat spoken responses
|
||||||
])
|
])
|
||||||
|
|
||||||
@transport.event_handler("on_first_participant_joined")
|
@transport.event_handler("on_first_participant_joined")
|
||||||
|
|||||||
99
examples/foundational/14-wake-phrase.py
Normal file
99
examples/foundational/14-wake-phrase.py
Normal file
@@ -0,0 +1,99 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import aiohttp
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from pipecat.processors.filters.wake_check_filter import WakeCheckFilter
|
||||||
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
|
from pipecat.pipeline.task import PipelineTask
|
||||||
|
from pipecat.processors.aggregators.llm_response import (
|
||||||
|
LLMAssistantResponseAggregator, LLMUserResponseAggregator)
|
||||||
|
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
||||||
|
from pipecat.services.openai import OpenAILLMService
|
||||||
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
from pipecat.vad.silero import SileroVADAnalyzer
|
||||||
|
|
||||||
|
from runner import configure
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
load_dotenv(override=True)
|
||||||
|
|
||||||
|
logger.remove(0)
|
||||||
|
logger.add(sys.stderr, level="DEBUG")
|
||||||
|
|
||||||
|
|
||||||
|
async def main(room_url: str, token):
|
||||||
|
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
transport = DailyTransport(
|
||||||
|
room_url,
|
||||||
|
token,
|
||||||
|
"Robot",
|
||||||
|
DailyParams(
|
||||||
|
audio_out_enabled=True,
|
||||||
|
transcription_enabled=True,
|
||||||
|
vad_enabled=True,
|
||||||
|
vad_analyzer=SileroVADAnalyzer()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
tts = ElevenLabsTTSService(
|
||||||
|
aiohttp_session=session,
|
||||||
|
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
||||||
|
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
|
||||||
|
)
|
||||||
|
|
||||||
|
llm = OpenAILLMService(
|
||||||
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
|
model="gpt-4o")
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{
|
||||||
|
"role": "system",
|
||||||
|
"content": "You are a helpful assistant. Respond to what the user said in a creative and helpful way. Keep your responses brief.",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
hey_robot_filter = WakeCheckFilter(["hey robot", "hey, robot"])
|
||||||
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
|
|
||||||
|
pipeline = Pipeline([
|
||||||
|
transport.input(), # Transport user input
|
||||||
|
hey_robot_filter, # Filter out speech not directed at the robot
|
||||||
|
tma_in, # User responses
|
||||||
|
llm, # LLM
|
||||||
|
tts, # TTS
|
||||||
|
transport.output(), # Transport bot output
|
||||||
|
tma_out # Assistant spoken responses
|
||||||
|
])
|
||||||
|
|
||||||
|
task = PipelineTask(pipeline, allow_interruptions=True)
|
||||||
|
|
||||||
|
@transport.event_handler("on_first_participant_joined")
|
||||||
|
async def on_first_participant_joined(transport, participant):
|
||||||
|
transport.capture_participant_transcription(participant["id"])
|
||||||
|
await tts.say("Hi! If you want to talk to me, just say 'Hey Robot'.")
|
||||||
|
|
||||||
|
# Kick off the conversation.
|
||||||
|
# messages.append(
|
||||||
|
# {"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
|
# await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
|
runner = PipelineRunner()
|
||||||
|
|
||||||
|
await runner.run(task)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
(url, token) = configure()
|
||||||
|
asyncio.run(main(url, token))
|
||||||
@@ -132,7 +132,7 @@ class TranscriptionFrame(TextFrame):
|
|||||||
timestamp: str
|
timestamp: str
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return f"{self.name}(user: {self.user_id}, text: {self.text}, timestamp: {self.timestamp})"
|
return f"{self.name}(user_id: {self.user_id}, text: {self.text}, timestamp: {self.timestamp})"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from pipecat.frames.frames import AppFrame, ControlFrame, Frame, SystemFrame
|
|||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
|
||||||
|
|
||||||
class Filter(FrameProcessor):
|
class FrameFilter(FrameProcessor):
|
||||||
|
|
||||||
def __init__(self, types: List[type]):
|
def __init__(self, types: List[type]):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
84
src/pipecat/processors/filters/wake_check_filter.py
Normal file
84
src/pipecat/processors/filters/wake_check_filter.py
Normal file
@@ -0,0 +1,84 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
|
||||||
|
class WakeCheckFilter(FrameProcessor):
|
||||||
|
"""
|
||||||
|
This filter looks for wake phrases in the transcription frames and only passes through frames
|
||||||
|
after a wake phrase has been detected. It also has a keepalive timeout to allow for a brief
|
||||||
|
period of continued conversation after a wake phrase has been detected.
|
||||||
|
"""
|
||||||
|
class WakeState(Enum):
|
||||||
|
IDLE = 1
|
||||||
|
AWAKE = 2
|
||||||
|
|
||||||
|
class ParticipantState:
|
||||||
|
def __init__(self, participant_id: str):
|
||||||
|
self.participant_id = participant_id
|
||||||
|
self.state = WakeCheckFilter.WakeState.IDLE
|
||||||
|
self.wake_timer = 0.0
|
||||||
|
self.accumulator = ""
|
||||||
|
|
||||||
|
def __init__(self, wake_phrases: list[str], keepalive_timeout: float = 2):
|
||||||
|
super().__init__()
|
||||||
|
self._participant_states = {}
|
||||||
|
self._keepalive_timeout = keepalive_timeout
|
||||||
|
self._wake_patterns = []
|
||||||
|
for name in wake_phrases:
|
||||||
|
pattern = re.compile(r'\b' + r'\s*'.join(re.escape(word)
|
||||||
|
for word in name.split()) + r'\b', re.IGNORECASE)
|
||||||
|
self._wake_patterns.append(pattern)
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
try:
|
||||||
|
if isinstance(frame, TranscriptionFrame):
|
||||||
|
p = self._participant_states.get(frame.user_id)
|
||||||
|
if p is None:
|
||||||
|
p = WakeCheckFilter.ParticipantState(frame.user_id)
|
||||||
|
self._participant_states[frame.user_id] = p
|
||||||
|
|
||||||
|
# If we have been AWAKE within the last keepalive_timeout seconds, pass
|
||||||
|
# the frame through
|
||||||
|
if p.state == WakeCheckFilter.WakeState.AWAKE:
|
||||||
|
if time.time() - p.wake_timer < self._keepalive_timeout:
|
||||||
|
logger.debug(
|
||||||
|
"Wake phrase keepalive timeout has not expired. Passing frame through.")
|
||||||
|
p.wake_timer = time.time()
|
||||||
|
await self.push_frame(frame)
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
p.state = WakeCheckFilter.WakeState.IDLE
|
||||||
|
|
||||||
|
p.accumulator += frame.text
|
||||||
|
for pattern in self._wake_patterns:
|
||||||
|
match = pattern.search(p.accumulator)
|
||||||
|
if match:
|
||||||
|
logger.debug(f"Wake phrase triggered: {match.group()}")
|
||||||
|
# Found the wake word. Discard from the accumulator up to the start of the match
|
||||||
|
# and modify the frame in place.
|
||||||
|
p.state = WakeCheckFilter.WakeState.AWAKE
|
||||||
|
p.wake_timer = time.time()
|
||||||
|
frame.text = p.accumulator[match.start():]
|
||||||
|
p.accumulator = ""
|
||||||
|
await self.push_frame(frame)
|
||||||
|
else:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = f"Error in wake word filter: {e}"
|
||||||
|
logger.error(error_msg)
|
||||||
|
await self.push_error(ErrorFrame(error_msg))
|
||||||
@@ -8,7 +8,7 @@ import asyncio
|
|||||||
from asyncio import AbstractEventLoop
|
from asyncio import AbstractEventLoop
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame
|
from pipecat.frames.frames import ErrorFrame, Frame
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|||||||
Reference in New Issue
Block a user