examples: fix whisper examples

This commit is contained in:
Aleix Conchillo Flaqué
2024-04-05 13:34:34 -07:00
parent efdfb74dc3
commit 9590cc2fbc
2 changed files with 23 additions and 11 deletions

View File

@@ -1,12 +1,16 @@
import asyncio import asyncio
import logging import logging
from dailyai.pipeline.frames import EndFrame, TranscriptionFrame
from dailyai.transports.daily_transport import DailyTransport from dailyai.transports.daily_transport import DailyTransport
from dailyai.services.whisper_ai_services import WhisperSTTService from dailyai.services.whisper_ai_services import WhisperSTTService
from dailyai.pipeline.pipeline import Pipeline from dailyai.pipeline.pipeline import Pipeline
from runner import configure from runner import configure
from dotenv import load_dotenv
load_dotenv(override=True)
logging.basicConfig(format=f"%(levelno)s %(asctime)s %(message)s") logging.basicConfig(format=f"%(levelno)s %(asctime)s %(message)s")
logger = logging.getLogger("dailyai") logger = logging.getLogger("dailyai")
logger.setLevel(logging.DEBUG) logger.setLevel(logging.DEBUG)
@@ -26,17 +30,27 @@ async def main(room_url: str):
stt = WhisperSTTService() stt = WhisperSTTService()
transcription_output_queue = asyncio.Queue() transcription_output_queue = asyncio.Queue()
transport_done = asyncio.Event()
pipeline = Pipeline([stt]) pipeline = Pipeline([stt], source=transport.receive_queue, sink=transcription_output_queue)
pipeline.set_sink(transcription_output_queue)
async def handle_transcription(): async def handle_transcription():
print("`````````TRANSCRIPTION`````````") print("`````````TRANSCRIPTION`````````")
while True: while not transport_done.is_set():
item = await transcription_output_queue.get() item = await transcription_output_queue.get()
print(item.text) print("got item from queue", item)
if isinstance(item, TranscriptionFrame):
print(item.text)
elif isinstance(item, EndFrame):
break
print("handle_transcription done")
await asyncio.gather(transport.run(pipeline), handle_transcription()) async def run_until_done():
await transport.run()
transport_done.set()
print("run_until_done done")
await asyncio.gather(run_until_done(), pipeline.run_pipeline(), handle_transcription())
if __name__ == "__main__": if __name__ == "__main__":

View File

@@ -15,11 +15,10 @@ async def main():
meeting_duration_minutes = 1 meeting_duration_minutes = 1
transport = LocalTransport( transport = LocalTransport(
mic_enabled=False, mic_enabled=True,
camera_enabled=False, camera_enabled=False,
speaker_enabled=True, speaker_enabled=True,
duration_minutes=meeting_duration_minutes, duration_minutes=meeting_duration_minutes,
start_transcription=False,
) )
stt = WhisperSTTService() stt = WhisperSTTService()
@@ -27,8 +26,7 @@ async def main():
transcription_output_queue = asyncio.Queue() transcription_output_queue = asyncio.Queue()
transport_done = asyncio.Event() transport_done = asyncio.Event()
pipeline = Pipeline([stt]) pipeline = Pipeline([stt], source=transport.receive_queue, sink=transcription_output_queue)
pipeline.set_sink(transcription_output_queue)
async def handle_transcription(): async def handle_transcription():
print("`````````TRANSCRIPTION`````````") print("`````````TRANSCRIPTION`````````")
@@ -42,11 +40,11 @@ async def main():
print("handle_transcription done") print("handle_transcription done")
async def run_until_done(): async def run_until_done():
await transport.run(pipeline) await transport.run()
transport_done.set() transport_done.set()
print("run_until_done done") print("run_until_done done")
await asyncio.gather(run_until_done(), handle_transcription()) await asyncio.gather(run_until_done(), pipeline.run_pipeline(), handle_transcription())
if __name__ == "__main__": if __name__ == "__main__":