scripts: allow storing audio for release evals
This commit is contained in:
@@ -6,13 +6,17 @@
|
|||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import io
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
|
import wave
|
||||||
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
|
import aiofiles
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from utils import (
|
from utils import (
|
||||||
EvalResult,
|
EvalResult,
|
||||||
@@ -31,6 +35,7 @@ from pipecat.pipeline.pipeline import Pipeline
|
|||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
||||||
|
from pipecat.processors.audio.audio_buffer_processor import AudioBufferProcessor
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.cartesia.tts import CartesiaTTSService
|
from pipecat.services.cartesia.tts import CartesiaTTSService
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
@@ -50,12 +55,17 @@ PIPELINE_IDLE_TIMEOUT_SECS = 30
|
|||||||
|
|
||||||
|
|
||||||
class EvalRunner:
|
class EvalRunner:
|
||||||
def __init__(self, pattern: str = ""):
|
def __init__(self, *, pattern: str = "", record_audio: bool = False):
|
||||||
self._pattern = f".*{pattern}.*" if pattern else ""
|
self._pattern = f".*{pattern}.*" if pattern else ""
|
||||||
|
self._record_audio = record_audio
|
||||||
self._total_success = 0
|
self._total_success = 0
|
||||||
self._tests: List[EvalResult] = []
|
self._tests: List[EvalResult] = []
|
||||||
self._queue = asyncio.Queue()
|
self._queue = asyncio.Queue()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def record_audio(self):
|
||||||
|
return self._record_audio
|
||||||
|
|
||||||
async def assert_eval(self, params: FunctionCallParams):
|
async def assert_eval(self, params: FunctionCallParams):
|
||||||
reasoning = params.arguments["reasoning"]
|
reasoning = params.arguments["reasoning"]
|
||||||
logger.debug(f"🧠 EVAL REASONING: {reasoning}")
|
logger.debug(f"🧠 EVAL REASONING: {reasoning}")
|
||||||
@@ -78,7 +88,7 @@ class EvalRunner:
|
|||||||
await asyncio.wait(
|
await asyncio.wait(
|
||||||
[
|
[
|
||||||
asyncio.create_task(run_example_pipeline(example_file)),
|
asyncio.create_task(run_example_pipeline(example_file)),
|
||||||
asyncio.create_task(run_eval_pipeline(self, prompt, eval)),
|
asyncio.create_task(run_eval_pipeline(self, example_file, prompt, eval)),
|
||||||
],
|
],
|
||||||
timeout=90,
|
timeout=90,
|
||||||
)
|
)
|
||||||
@@ -126,7 +136,31 @@ async def run_example_pipeline(example_file: str):
|
|||||||
await module.run_example(transport, argparse.Namespace(), True)
|
await module.run_example(transport, argparse.Namespace(), True)
|
||||||
|
|
||||||
|
|
||||||
async def run_eval_pipeline(eval_runner: EvalRunner, prompt: str, eval: Optional[str]):
|
async def save_audio(audio: bytes, sample_rate: int, num_channels: int, name: str):
|
||||||
|
if len(audio) > 0:
|
||||||
|
recordings_dir = os.path.join(SCRIPT_DIR, "recordings")
|
||||||
|
base_name = os.path.splitext(name)[0]
|
||||||
|
os.makedirs(recordings_dir, exist_ok=True)
|
||||||
|
filename = os.path.join(
|
||||||
|
recordings_dir,
|
||||||
|
f"{base_name}_{datetime.now().strftime('%Y%m%d_%H%M%S')}.wav",
|
||||||
|
)
|
||||||
|
with io.BytesIO() as buffer:
|
||||||
|
with wave.open(buffer, "wb") as wf:
|
||||||
|
wf.setsampwidth(2)
|
||||||
|
wf.setnchannels(num_channels)
|
||||||
|
wf.setframerate(sample_rate)
|
||||||
|
wf.writeframes(audio)
|
||||||
|
async with aiofiles.open(filename, "wb") as file:
|
||||||
|
await file.write(buffer.getvalue())
|
||||||
|
logger.debug(f"Saving {name} audio to {filename}")
|
||||||
|
else:
|
||||||
|
logger.warning(f"There's no audio to save for {name}")
|
||||||
|
|
||||||
|
|
||||||
|
async def run_eval_pipeline(
|
||||||
|
eval_runner: EvalRunner, example_file: str, prompt: str, eval: Optional[str]
|
||||||
|
):
|
||||||
logger.info(f"Starting eval bot")
|
logger.info(f"Starting eval bot")
|
||||||
|
|
||||||
room_url = os.getenv("DAILY_SAMPLE_ROOM_URL")
|
room_url = os.getenv("DAILY_SAMPLE_ROOM_URL")
|
||||||
@@ -185,6 +219,8 @@ async def run_eval_pipeline(eval_runner: EvalRunner, prompt: str, eval: Optional
|
|||||||
context = OpenAILLMContext(messages, tools)
|
context = OpenAILLMContext(messages, tools)
|
||||||
context_aggregator = llm.create_context_aggregator(context)
|
context_aggregator = llm.create_context_aggregator(context)
|
||||||
|
|
||||||
|
audio_buffer = AudioBufferProcessor()
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(), # Transport user input
|
transport.input(), # Transport user input
|
||||||
@@ -193,19 +229,30 @@ async def run_eval_pipeline(eval_runner: EvalRunner, prompt: str, eval: Optional
|
|||||||
llm, # LLM
|
llm, # LLM
|
||||||
tts, # TTS
|
tts, # TTS
|
||||||
transport.output(), # Transport bot output
|
transport.output(), # Transport bot output
|
||||||
|
audio_buffer,
|
||||||
context_aggregator.assistant(), # Assistant spoken responses
|
context_aggregator.assistant(), # Assistant spoken responses
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
task = PipelineTask(
|
task = PipelineTask(
|
||||||
pipeline,
|
pipeline,
|
||||||
params=PipelineParams(allow_interruptions=True),
|
params=PipelineParams(
|
||||||
|
allow_interruptions=True,
|
||||||
|
audio_in_sample_rate=16000,
|
||||||
|
audio_out_sample_rate=16000,
|
||||||
|
),
|
||||||
idle_timeout_secs=PIPELINE_IDLE_TIMEOUT_SECS,
|
idle_timeout_secs=PIPELINE_IDLE_TIMEOUT_SECS,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@audio_buffer.event_handler("on_audio_data")
|
||||||
|
async def on_audio_data(buffer, audio, sample_rate, num_channels):
|
||||||
|
if eval_runner.record_audio:
|
||||||
|
await save_audio(audio, sample_rate, num_channels, example_file)
|
||||||
|
|
||||||
@transport.event_handler("on_client_connected")
|
@transport.event_handler("on_client_connected")
|
||||||
async def on_client_connected(transport, client):
|
async def on_client_connected(transport, client):
|
||||||
logger.info(f"Client connected")
|
logger.info(f"Client connected")
|
||||||
|
await audio_buffer.start_recording()
|
||||||
|
|
||||||
@transport.event_handler("on_client_disconnected")
|
@transport.event_handler("on_client_disconnected")
|
||||||
async def on_client_disconnected(transport, client):
|
async def on_client_disconnected(transport, client):
|
||||||
|
|||||||
@@ -117,7 +117,7 @@ async def main(args: argparse.Namespace):
|
|||||||
if not check_env_variables():
|
if not check_env_variables():
|
||||||
return
|
return
|
||||||
|
|
||||||
runner = EvalRunner(args.pattern)
|
runner = EvalRunner(pattern=args.pattern, record_audio=args.audio)
|
||||||
for test, prompt, eval in TESTS:
|
for test, prompt, eval in TESTS:
|
||||||
await runner.run_eval(test, prompt, eval)
|
await runner.run_eval(test, prompt, eval)
|
||||||
|
|
||||||
@@ -126,6 +126,7 @@ async def main(args: argparse.Namespace):
|
|||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser(description="Pipecat Eval Runner")
|
parser = argparse.ArgumentParser(description="Pipecat Eval Runner")
|
||||||
|
parser.add_argument("--audio", "-a", action="store_true", help="Record audio for each test")
|
||||||
parser.add_argument("--pattern", "-p", help="Only run tests that match the pattern")
|
parser.add_argument("--pattern", "-p", help="Only run tests that match the pattern")
|
||||||
parser.add_argument("--verbose", "-v", action="count", default=0)
|
parser.add_argument("--verbose", "-v", action="count", default=0)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|||||||
Reference in New Issue
Block a user