From 24b85e1d06eee8f66b2360afdcda8fac11db8c53 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 1/9] streaming: scheduled left-context chunk policy for early first audio --- mstar/streaming/chunk_policy.py | 64 +++++++++++++++++++++++++++++++++ 1 file changed, 64 insertions(+) diff --git a/mstar/streaming/chunk_policy.py b/mstar/streaming/chunk_policy.py index 6c527fe46..3e945d771 100644 --- a/mstar/streaming/chunk_policy.py +++ b/mstar/streaming/chunk_policy.py @@ -1,4 +1,5 @@ from abc import ABC, abstractmethod +from collections.abc import Sequence class ChunkPolicy(ABC): @@ -150,3 +151,66 @@ def window_size(self) -> int: def continue_after_producer_done(self) -> bool: return self._continue_after_done + + +class ScheduledLeftContextChunkPolicy(ChunkPolicy): + """Left-context chunking whose chunk sizes follow a ramp. + + Streaming vocoders want the first audio out as early as possible and + larger chunks once the stream is running. Chunk ``k`` delivers + ``schedule[k]`` new items (``chunk`` once the schedule is exhausted) with + up to ``left_context`` already-delivered items in front of them, so a + causal decoder can warm up on frames it has processed before. Unlike + ``LeftContextChunkPolicy`` the first chunk may be smaller than the + context: the context is whatever has been delivered so far, capped. + + Example (Qwen3-TTS, 12 Hz frames): ``schedule=(4, 8, 16)``, ``chunk=25``, + ``left_context=25`` pops windows of 4, 4+8, 12+16, 25+25, 25+25, ... + items and the first audio leaves after four frames instead of 300. + + The consumer learns how many leading items of a window are context from + ``StreamChunk.context_items`` (the worker passes it along as + ``step_metadata["stream_chunks"][edge]["context_items"]``), so it can trim + the duplicated output without re-deriving this schedule. + """ + + def __init__(self, schedule: Sequence[int], chunk: int, left_context: int): + super().__init__() + if chunk <= 0 or any(size <= 0 for size in schedule) or left_context < 0: + raise ValueError("chunk sizes must be positive and left_context non-negative") + self._schedule = tuple(int(size) for size in schedule) + self._chunk = int(chunk) + self._left_context = int(left_context) + self._chunks_popped = 0 + self._delivered = 0 # new items handed to the consumer so far + + def _new_items(self) -> int: + if self._chunks_popped < len(self._schedule): + return self._schedule[self._chunks_popped] + return self._chunk + + def _context(self) -> int: + return min(self._left_context, self._delivered) + + def is_ready(self, buffer_len: int) -> bool: + return buffer_len >= self.window_size() + + def window_size(self) -> int: + return self._context() + self._new_items() + + def next_chunk_size(self, buffer_len: int) -> int: + # The buffer pointer sits ``context`` items before the first new item. + # After this pop it must sit ``next context`` items before the next + # chunk's first new item. + delivered_after = self._delivered + self._new_items() + next_context = min(self._left_context, delivered_after) + return (delivered_after - next_context) - (self._delivered - self._context()) + + def register_chunk(self, chunk_size: int): + super().register_chunk(chunk_size) + # A regular pop delivers exactly this chunk's new items. The only + # other caller is the terminal flush (producer done, window not + # full), after which no data-carrying chunk follows, so treating it + # the same keeps the bookkeeping trivially correct where it matters. + self._delivered += self._new_items() + self._chunks_popped += 1 From b38fe36aed0f82b95977a06931d3fdbcd627fe36 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 2/9] streaming: report each chunk's already-delivered context items --- mstar/streaming/stream_buffer.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/mstar/streaming/stream_buffer.py b/mstar/streaming/stream_buffer.py index d2ce728c0..3375f7f6b 100644 --- a/mstar/streaming/stream_buffer.py +++ b/mstar/streaming/stream_buffer.py @@ -14,6 +14,9 @@ class StreamChunk: chunk_index: int start_offset: int = 0 # global position of the first item in this chunk is_final: bool = False + # leading items of this chunk that an earlier chunk already delivered + # (sliding-window overlap / left context); the consumer trims their output + context_items: int = 0 @dataclass @@ -39,6 +42,9 @@ class StreamBuffer: _id_to_tensor: dict = field(default_factory=dict) _consumed: int = 0 _chunks_popped: int = 0 + # global position just past the last item ever handed out; everything + # before it in a later window is context, not new data + _delivered_end: int = 0 producer_done: bool = False # Set once a chunk has been popped with ``is_final=True`` (the terminal # flush). Guards the empty-buffer final flush below so it fires exactly @@ -138,7 +144,9 @@ def pop_chunk(self) -> StreamChunk: chunk_index=self._chunks_popped, start_offset=offset, is_final=is_final, + context_items=min(max(self._delivered_end - offset, 0), len(items)), ) + self._delivered_end = max(self._delivered_end, offset + len(items)) self._chunks_popped += 1 return chunk From 58fe5bc949b21186d849851385ce784c8ccf14e9 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 3/9] graph: carry stream chunk offset and context on synthetic edges --- mstar/graph/base.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/mstar/graph/base.py b/mstar/graph/base.py index 17c38de34..7b9524848 100644 --- a/mstar/graph/base.py +++ b/mstar/graph/base.py @@ -74,6 +74,11 @@ class GraphEdge: # set on a synthetic streaming-input edge carrying the final chunk, so the # consuming pass (not the earlier ingest) reports the partition done _final_stream_chunk: bool = field(default=False) + # set on a synthetic streaming-input edge: where the chunk starts in the + # stream and how many of its leading items were delivered before (context); + # None on every other edge + _stream_chunk_offset: int | None = field(default=None) + _stream_chunk_context: int | None = field(default=None) # Set for sharded configurations _total_fanin: int = 1 @@ -90,6 +95,8 @@ def clone(self): output_modality=self.output_modality, _persist_for_loop=self._persist_for_loop, _final_stream_chunk=self._final_stream_chunk, + _stream_chunk_offset=self._stream_chunk_offset, + _stream_chunk_context=self._stream_chunk_context, ) From e8954cce4df8799257a099536a4abb89add5ba83 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 4/9] worker: expose stream chunk geometry to consumers via step_metadata --- mstar/worker/worker.py | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/mstar/worker/worker.py b/mstar/worker/worker.py index 8a9ad0523..0dd09f4f4 100644 --- a/mstar/worker/worker.py +++ b/mstar/worker/worker.py @@ -789,6 +789,8 @@ def _pop_streaming_edge( name=edge_name, tensor_info=[], _final_stream_chunk=chunk.is_final, + _stream_chunk_offset=chunk.start_offset, + _stream_chunk_context=chunk.context_items, ) else: # Normal chunk — store tensor and create edge with tensor_info. @@ -807,6 +809,8 @@ def _pop_streaming_edge( name=edge_name, tensor_info=tensor_infos.get(edge_name, []), _final_stream_chunk=chunk.is_final, + _stream_chunk_offset=chunk.start_offset, + _stream_chunk_context=chunk.context_items, ) return synthetic_edge @@ -958,6 +962,7 @@ def _build_executing_batch(self, batch: ScheduledBatch) -> ExecutingBatch: for request_id, node in batch.node_objects.items(): tensors = {} + stream_chunks = {} ready_inputs = node.ready_signals.ready_inputs for input_name, edge in ready_inputs.items(): tensors[input_name] = [ @@ -967,8 +972,20 @@ def _build_executing_batch(self, batch: ScheduledBatch) -> ExecutingBatch: ] if edge._final_stream_chunk: final_stream_rids.add(request_id) + if edge._stream_chunk_context is not None: + stream_chunks[input_name] = { + "start_offset": edge._stream_chunk_offset, + "context_items": edge._stream_chunk_context, + "is_final": edge._final_stream_chunk, + } per_request_inputs[request_id] = tensors - per_request_info[request_id] = self.worker_graphs_manager.get_fwd_info(request_id, batch_partition) + fwd_info = self.worker_graphs_manager.get_fwd_info(request_id, batch_partition) + if stream_chunks: + # Where each streamed input sits in its stream and how many of + # its leading items are repeated context, for the consumer's + # ``prepare_inputs`` (a vocoder trims that context's audio). + fwd_info.step_metadata["stream_chunks"] = stream_chunks + per_request_info[request_id] = fwd_info return self._make_executing_batch( node_name=batch.node_name, From 0b393ead0cce27108f8ad1adfc57c147eb33c696 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 5/9] test: scheduled chunk policy and chunk context geometry --- test/modular/test_stream_chunk_schedule.py | 126 +++++++++++++++++++++ 1 file changed, 126 insertions(+) create mode 100644 test/modular/test_stream_chunk_schedule.py diff --git a/test/modular/test_stream_chunk_schedule.py b/test/modular/test_stream_chunk_schedule.py new file mode 100644 index 000000000..c35d6f528 --- /dev/null +++ b/test/modular/test_stream_chunk_schedule.py @@ -0,0 +1,126 @@ +"""``ScheduledLeftContextChunkPolicy`` and the chunk geometry a StreamBuffer reports. + +A streaming vocoder decodes each popped window and must drop the audio of the +leading ``context_items`` (frames an earlier chunk already delivered). These +tests pin (a) the window / context sequence of the ramped policy, (b) that every +item reaches the consumer exactly once as new data whatever the ramp, and (c) +that ``StreamChunk.context_items`` is right for every policy in the tree, so a +consumer can rely on it instead of re-deriving the policy's schedule. +""" + +import pytest +import torch + +from mstar.graph.base import GraphEdge +from mstar.streaming.chunk_policy import ( + FixedChunkPolicy, + LeftContextChunkPolicy, + ScheduledLeftContextChunkPolicy, + SlidingWindowChunkPolicy, +) +from mstar.streaming.stream_buffer import StreamBuffer + + +def _drive(policy, total_items, drain_before_done=True): + """Feed ``total_items`` one-row items; return (chunks, is_final flags).""" + buffer = StreamBuffer(request_id="r", edge_name="codec_tokens", from_partition="Talker", policy=policy) + chunks = [] + + def poll(): + for _ in range(total_items + 50): + if not buffer.has_chunk_ready(): + return + chunks.append(buffer.pop_chunk()) + raise AssertionError("has_chunk_ready never went False") + + for i in range(total_items): + buffer.pre_read_register(f"t{i}") + buffer.put(f"t{i}", torch.tensor([i])) + if drain_before_done: + poll() + buffer.signal_done() + poll() + return chunks + + +def _geometry(chunks): + """(window, context, start_offset) per data-carrying chunk.""" + out = [] + for chunk in chunks: + data = chunk.data["data"] + if data is None: + continue + items = data.reshape(-1).tolist() + out.append((len(items), chunk.context_items, chunk.start_offset, items)) + return out + + +def _new_items(chunks): + delivered = [] + for window, context, _, items in _geometry(chunks): + delivered.extend(items[context:window]) + return delivered + + +def test_scheduled_policy_ramps_windows_and_reports_context(): + policy = ScheduledLeftContextChunkPolicy(schedule=(4, 8, 16), chunk=25, left_context=25) + chunks = _drive(policy, total_items=120) + geometry = [(w, c, o) for w, c, o, _ in _geometry(chunks)] + # window = context + new: 4 | 4+8 | 12+16 | 25+25 | 25+25 | flush + assert geometry[:5] == [(4, 0, 0), (12, 4, 0), (28, 12, 0), (50, 25, 3), (50, 25, 28)] + # The terminal flush hands over the remaining 17 new items behind 25 of context. + assert geometry[-1] == (42, 25, 78) + assert _new_items(chunks) == list(range(120)) + assert sum(chunk.is_final for chunk in chunks) == 1 and chunks[-1].is_final + # Every window's leading context is exactly the items the previous chunk ended with. + geometry_all = _geometry(chunks) + for prev, curr in zip(geometry_all, geometry_all[1:], strict=False): + assert curr[3][:curr[1]] == prev[3][len(prev[3]) - curr[1]:] + + +@pytest.mark.parametrize("drain_before_done", [True, False]) +@pytest.mark.parametrize("total_items", [0, 1, 3, 4, 5, 12, 28, 53, 78, 100, 153]) +def test_scheduled_policy_delivers_everything_once_with_one_final(total_items, drain_before_done): + policy = ScheduledLeftContextChunkPolicy(schedule=(4, 8, 16), chunk=25, left_context=25) + chunks = _drive(policy, total_items, drain_before_done) + assert _new_items(chunks) == list(range(total_items)) + assert sum(chunk.is_final for chunk in chunks) == 1 and chunks[-1].is_final + + +def test_scheduled_policy_context_smaller_than_first_chunks(): + # Left context shorter than the ramp steps: context saturates at 2. + policy = ScheduledLeftContextChunkPolicy(schedule=(1, 3), chunk=5, left_context=2) + chunks = _drive(policy, total_items=14) + assert [(w, c, o) for w, c, o, _ in _geometry(chunks)] == [ + (1, 0, 0), (4, 1, 0), (7, 2, 2), (7, 2, 7), (2, 2, 12), + ] + assert _new_items(chunks) == list(range(14)) + + +def test_scheduled_policy_rejects_bad_sizes(): + with pytest.raises(ValueError): + ScheduledLeftContextChunkPolicy(schedule=(0,), chunk=5, left_context=1) + with pytest.raises(ValueError): + ScheduledLeftContextChunkPolicy(schedule=(), chunk=5, left_context=-1) + + +@pytest.mark.parametrize( + ("policy", "expected"), + [ + (FixedChunkPolicy(chunk_size=3), [(3, 0), (3, 0), (3, 0), (1, 0)]), + (SlidingWindowChunkPolicy(window=4, stride=2), [(4, 0), (4, 2), (4, 2), (4, 2), (2, 2)]), + (LeftContextChunkPolicy(chunk=4, left_context=1), [(4, 0), (5, 1), (3, 1)]), + ], +) +def test_existing_policies_report_context_items(policy, expected): + chunks = _drive(policy, total_items=10) + assert [(w, c) for w, c, _, _ in _geometry(chunks)] == expected + assert _new_items(chunks) == list(range(10)) + + +def test_graph_edge_clone_keeps_stream_chunk_geometry(): + edge = GraphEdge(next_node="Codec", name="codec_tokens", _final_stream_chunk=True, + _stream_chunk_offset=7, _stream_chunk_context=3) + clone = edge.clone() + assert (clone._stream_chunk_offset, clone._stream_chunk_context, clone._final_stream_chunk) == (7, 3, True) + assert GraphEdge(next_node="x", name="y")._stream_chunk_context is None From 2e8c1e8e76cc0326a54a001926ca93ce1fe2368b Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 6/9] streaming: report each chunk's item count --- mstar/streaming/stream_buffer.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/mstar/streaming/stream_buffer.py b/mstar/streaming/stream_buffer.py index 3375f7f6b..aa8dec71a 100644 --- a/mstar/streaming/stream_buffer.py +++ b/mstar/streaming/stream_buffer.py @@ -17,6 +17,9 @@ class StreamChunk: # leading items of this chunk that an earlier chunk already delivered # (sliding-window overlap / left context); the consumer trims their output context_items: int = 0 + # items in this chunk (context included); lets a consumer pick its + # capture bucket before the chunk tensor is unpacked + num_items: int = 0 @dataclass @@ -145,6 +148,7 @@ def pop_chunk(self) -> StreamChunk: start_offset=offset, is_final=is_final, context_items=min(max(self._delivered_end - offset, 0), len(items)), + num_items=len(items), ) self._delivered_end = max(self._delivered_end, offset + len(items)) self._chunks_popped += 1 From eab9de57de5157fec19ae48d91e01feda2e29f7a Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 7/9] graph: carry the stream chunk item count on synthetic edges --- mstar/graph/base.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/mstar/graph/base.py b/mstar/graph/base.py index 7b9524848..a4e0b8b62 100644 --- a/mstar/graph/base.py +++ b/mstar/graph/base.py @@ -79,6 +79,7 @@ class GraphEdge: # None on every other edge _stream_chunk_offset: int | None = field(default=None) _stream_chunk_context: int | None = field(default=None) + _stream_chunk_items: int | None = field(default=None) # Set for sharded configurations _total_fanin: int = 1 @@ -97,6 +98,7 @@ def clone(self): _final_stream_chunk=self._final_stream_chunk, _stream_chunk_offset=self._stream_chunk_offset, _stream_chunk_context=self._stream_chunk_context, + _stream_chunk_items=self._stream_chunk_items, ) From cc115269316a04f77df214a2350b9cf45c62d515 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 8/9] worker: include the chunk item count in stream_chunks metadata --- mstar/worker/worker.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/mstar/worker/worker.py b/mstar/worker/worker.py index 0dd09f4f4..569541e30 100644 --- a/mstar/worker/worker.py +++ b/mstar/worker/worker.py @@ -791,6 +791,7 @@ def _pop_streaming_edge( _final_stream_chunk=chunk.is_final, _stream_chunk_offset=chunk.start_offset, _stream_chunk_context=chunk.context_items, + _stream_chunk_items=chunk.num_items, ) else: # Normal chunk — store tensor and create edge with tensor_info. @@ -811,6 +812,7 @@ def _pop_streaming_edge( _final_stream_chunk=chunk.is_final, _stream_chunk_offset=chunk.start_offset, _stream_chunk_context=chunk.context_items, + _stream_chunk_items=chunk.num_items, ) return synthetic_edge @@ -976,6 +978,7 @@ def _build_executing_batch(self, batch: ScheduledBatch) -> ExecutingBatch: stream_chunks[input_name] = { "start_offset": edge._stream_chunk_offset, "context_items": edge._stream_chunk_context, + "num_items": edge._stream_chunk_items, "is_final": edge._final_stream_chunk, } per_request_inputs[request_id] = tensors From adfec01a48a175d13dbaef138e29bfd08b700c25 Mon Sep 17 00:00:00 2001 From: merceod Date: Fri, 18 Sep 2026 02:09:58 -0700 Subject: [PATCH 9/9] test: stream chunk item counts --- test/modular/test_stream_chunk_schedule.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/test/modular/test_stream_chunk_schedule.py b/test/modular/test_stream_chunk_schedule.py index c35d6f528..7abc6aa05 100644 --- a/test/modular/test_stream_chunk_schedule.py +++ b/test/modular/test_stream_chunk_schedule.py @@ -49,8 +49,10 @@ def _geometry(chunks): for chunk in chunks: data = chunk.data["data"] if data is None: + assert chunk.num_items == 0 continue items = data.reshape(-1).tolist() + assert chunk.num_items == len(items) out.append((len(items), chunk.context_items, chunk.start_offset, items)) return out @@ -120,7 +122,9 @@ def test_existing_policies_report_context_items(policy, expected): def test_graph_edge_clone_keeps_stream_chunk_geometry(): edge = GraphEdge(next_node="Codec", name="codec_tokens", _final_stream_chunk=True, - _stream_chunk_offset=7, _stream_chunk_context=3) + _stream_chunk_offset=7, _stream_chunk_context=3, _stream_chunk_items=12) clone = edge.clone() - assert (clone._stream_chunk_offset, clone._stream_chunk_context, clone._final_stream_chunk) == (7, 3, True) - assert GraphEdge(next_node="x", name="y")._stream_chunk_context is None + assert (clone._stream_chunk_offset, clone._stream_chunk_context, clone._stream_chunk_items, + clone._final_stream_chunk) == (7, 3, 12, True) + plain = GraphEdge(next_node="x", name="y") + assert plain._stream_chunk_context is None and plain._stream_chunk_items is None