Skip to content
9 changes: 9 additions & 0 deletions mstar/graph/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,12 @@ 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)
_stream_chunk_items: int | None = field(default=None)

# Set for sharded configurations
_total_fanin: int = 1
Expand All @@ -90,6 +96,9 @@ 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,
_stream_chunk_items=self._stream_chunk_items,
)


Expand Down
64 changes: 64 additions & 0 deletions mstar/streaming/chunk_policy.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from abc import ABC, abstractmethod
from collections.abc import Sequence


class ChunkPolicy(ABC):
Expand Down Expand Up @@ -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
12 changes: 12 additions & 0 deletions mstar/streaming/stream_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,12 @@ 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
# items in this chunk (context included); lets a consumer pick its
# capture bucket before the chunk tensor is unpacked
num_items: int = 0


@dataclass
Expand All @@ -39,6 +45,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
Expand Down Expand Up @@ -138,7 +147,10 @@ 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)),
num_items=len(items),
)
self._delivered_end = max(self._delivered_end, offset + len(items))
self._chunks_popped += 1
return chunk

Expand Down
22 changes: 21 additions & 1 deletion mstar/worker/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -789,6 +789,9 @@ 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,
_stream_chunk_items=chunk.num_items,
)
else:
# Normal chunk — store tensor and create edge with tensor_info.
Expand All @@ -807,6 +810,9 @@ 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,
_stream_chunk_items=chunk.num_items,
)
return synthetic_edge

Expand Down Expand Up @@ -958,6 +964,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] = [
Expand All @@ -967,8 +974,21 @@ 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,
"num_items": edge._stream_chunk_items,
"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,
Expand Down
130 changes: 130 additions & 0 deletions test/modular/test_stream_chunk_schedule.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
"""``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:
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


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, _stream_chunk_items=12)
clone = edge.clone()
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
Loading