Merge pull request #885 from pipecat-ai/aleix/parallelpipeline-wait-for-slowest-endframe

pipeline(parallel): wait for slowest endframe
This commit is contained in:
Aleix Conchillo Flaqué
2024-12-19 15:18:22 -08:00
committed by GitHub
2 changed files with 84 additions and 31 deletions

View File

@@ -57,6 +57,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
### Fixed ### Fixed
- Fixed an issue that would cause `ParallelPipeline` to handle `EndFrame`
incorrectly causing the main pipeline to not terminate or terminate too early.
- Fixed an audio stuttering issue in `FastPitchTTSService`. - Fixed an audio stuttering issue in `FastPitchTTSService`.
- Fixed a `BaseOutputTransport` issue that was causing non-audio frames being - Fixed a `BaseOutputTransport` issue that was causing non-audio frames being

View File

@@ -6,11 +6,18 @@
import asyncio import asyncio
from itertools import chain from itertools import chain
from typing import Awaitable, Callable, List from typing import Awaitable, Callable, Dict, List
from loguru import logger from loguru import logger
from pipecat.frames.frames import CancelFrame, EndFrame, Frame, StartFrame, SystemFrame from pipecat.frames.frames import (
CancelFrame,
EndFrame,
Frame,
StartFrame,
StartInterruptionFrame,
SystemFrame,
)
from pipecat.pipeline.base_pipeline import BasePipeline from pipecat.pipeline.base_pipeline import BasePipeline
from pipecat.pipeline.pipeline import Pipeline from pipecat.pipeline.pipeline import Pipeline
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
@@ -72,11 +79,10 @@ class ParallelPipeline(BasePipeline):
self._sources = [] self._sources = []
self._sinks = [] self._sinks = []
self._seen_ids = set() self._seen_ids = set()
self._endframe_counter: Dict[int, int] = {}
self._up_queue = asyncio.Queue() self._up_queue = asyncio.Queue()
self._down_queue = asyncio.Queue() self._down_queue = asyncio.Queue()
self._up_task: asyncio.Task | None = None
self._down_task: asyncio.Task | None = None
self._pipelines = [] self._pipelines = []
@@ -111,18 +117,19 @@ class ParallelPipeline(BasePipeline):
# #
async def cleanup(self): async def cleanup(self):
await asyncio.gather(*[s.cleanup() for s in self._sources])
await asyncio.gather(*[p.cleanup() for p in self._pipelines]) await asyncio.gather(*[p.cleanup() for p in self._pipelines])
await asyncio.gather(*[s.cleanup() for s in self._sinks])
async def _start_tasks(self):
loop = self.get_event_loop()
self._up_task = loop.create_task(self._process_up_queue())
self._down_task = loop.create_task(self._process_down_queue())
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
await super().process_frame(frame, direction) await super().process_frame(frame, direction)
if isinstance(frame, StartFrame): if isinstance(frame, StartFrame):
await self._start_tasks() await self._start()
elif isinstance(frame, EndFrame):
self._endframe_counter[frame.id] = len(self._pipelines)
elif isinstance(frame, CancelFrame):
await self._cancel()
if direction == FrameDirection.UPSTREAM: if direction == FrameDirection.UPSTREAM:
# If we get an upstream frame we process it in each sink. # If we get an upstream frame we process it in each sink.
@@ -131,16 +138,45 @@ class ParallelPipeline(BasePipeline):
# If we get a downstream frame we process it in each source. # If we get a downstream frame we process it in each source.
await asyncio.gather(*[s.queue_frame(frame, direction) for s in self._sources]) await asyncio.gather(*[s.queue_frame(frame, direction) for s in self._sources])
# If we get an EndFrame we stop our queue processing tasks and wait on # Handle interruptions after everything has been cancelled.
# all the pipelines to finish. if isinstance(frame, StartInterruptionFrame):
if isinstance(frame, (CancelFrame, EndFrame)): await self._handle_interruption()
# Use None to indicate when queues should be done processing. # Wait for tasks to finish.
await self._up_queue.put(None) elif isinstance(frame, EndFrame):
await self._down_queue.put(None) await self._stop()
if self._up_task:
await self._up_task async def _start(self):
if self._down_task: await self._create_tasks()
await self._down_task
async def _stop(self):
# The up task doesn't receive an EndFrame, so we just cancel it.
self._up_task.cancel()
await self._up_task
# The down tasks waits for the last EndFrame send by the internal
# pipelines.
await self._down_task
async def _cancel(self):
self._up_task.cancel()
await self._up_task
self._down_task.cancel()
await self._down_task
async def _create_tasks(self):
loop = self.get_event_loop()
self._up_task = loop.create_task(self._process_up_queue())
self._down_task = loop.create_task(self._process_down_queue())
async def _drain_queues(self):
while not self._up_queue.empty:
await self._up_queue.get()
while not self._down_queue.empty:
await self._down_queue.get()
async def _handle_interruption(self):
await self._cancel()
await self._drain_queues()
await self._create_tasks()
async def _parallel_push_frame(self, frame: Frame, direction: FrameDirection): async def _parallel_push_frame(self, frame: Frame, direction: FrameDirection):
if frame.id not in self._seen_ids: if frame.id not in self._seen_ids:
@@ -148,19 +184,33 @@ class ParallelPipeline(BasePipeline):
await self.push_frame(frame, direction) await self.push_frame(frame, direction)
async def _process_up_queue(self): async def _process_up_queue(self):
running = True while True:
while running: try:
frame = await self._up_queue.get() frame = await self._up_queue.get()
if frame:
await self._parallel_push_frame(frame, FrameDirection.UPSTREAM) await self._parallel_push_frame(frame, FrameDirection.UPSTREAM)
running = frame is not None self._up_queue.task_done()
self._up_queue.task_done() except asyncio.CancelledError:
break
async def _process_down_queue(self): async def _process_down_queue(self):
running = True running = True
while running: while running:
frame = await self._down_queue.get() try:
if frame: frame = await self._down_queue.get()
await self._parallel_push_frame(frame, FrameDirection.DOWNSTREAM)
running = frame is not None endframe_counter = self._endframe_counter.get(frame.id, 0)
self._down_queue.task_done()
# If we have a counter, decrement it.
if endframe_counter > 0:
self._endframe_counter[frame.id] -= 1
endframe_counter = self._endframe_counter[frame.id]
# If we don't have a counter or we reached 0, push the frame.
if endframe_counter == 0:
await self._parallel_push_frame(frame, FrameDirection.DOWNSTREAM)
running = not (endframe_counter == 0 and isinstance(frame, EndFrame))
self._down_queue.task_done()
except asyncio.CancelledError:
break