class RoutedExpertsBuffer:
"""Own incomplete and unkeyed full routed-experts blocks."""
def __init__(
self,
dtype: np.dtype[Any],
shape_per_token: tuple[int, ...],
block_size: int,
max_num_seqs: int,
max_num_batched_tokens: int,
max_concurrent_batches: int,
) -> None:
self.dtype = dtype
self.shape_per_token = shape_per_token
self.block_size = block_size
max_step_blocks = (
max_num_batched_tokens + block_size - 1
) // block_size + max_num_seqs
# In-flight full blocks plus one incomplete tail per active request.
max_blocks = max_concurrent_batches * max_step_blocks + max_num_seqs
self._rows: np.ndarray = np.empty(
(max_blocks, block_size, *shape_per_token), dtype=dtype
)
self._free_slots = list(range(len(self._rows) - 1, -1, -1))
self._owned_slots: dict[int, int] = {}
self._requests: dict[Hashable, _RequestTail] = {}
def _tail(self, request_id: Hashable, block_start: int) -> _RequestTail:
tail = self._requests.get(request_id)
if tail is not None:
assert tail.block_start == block_start, (
"auxiliary output capture skipped an incomplete block: "
f"request={request_id}, expected={tail.block_start}, "
f"actual={block_start}"
)
return tail
assert self._free_slots, "auxiliary output block pool is exhausted"
tail = _RequestTail(self._free_slots.pop(), block_start)
self._requests[request_id] = tail
return tail
def capture(
self, request_id: Hashable, token_start: int, rows: np.ndarray
) -> list[tuple[int, np.ndarray]]:
"""Stage rows and return completed blocks without retaining them."""
rows = np.asarray(rows)
assert rows.shape[1:] == self.shape_per_token and rows.dtype == self.dtype, (
"routed-experts capture profile changed"
)
if token_start < 0:
raise ValueError("auxiliary output token start must be non-negative")
completed: list[tuple[int, np.ndarray]] = []
offset = 0
while offset < len(rows):
position = token_start + offset
block_start = position // self.block_size * self.block_size
local_start = position - block_start
# Full aligned input blocks do not need tail staging.
if (
request_id not in self._requests
and local_start == 0
and len(rows) - offset >= self.block_size
):
completed.append((block_start, rows[offset : offset + self.block_size]))
offset += self.block_size
continue
tail = self._tail(request_id, block_start)
assert local_start == tail.length, (
"auxiliary output capture is not contiguous: "
f"request={request_id}, expected={block_start + tail.length}, "
f"actual={position}"
)
count = min(self.block_size - local_start, len(rows) - offset)
if count == 1:
self._rows[tail.slot, local_start] = rows[offset]
else:
self._rows[tail.slot, local_start : local_start + count] = rows[
offset : offset + count
]
tail.length += count
offset += count
if tail.length == self.block_size:
block = self._rows[tail.slot]
del self._requests[request_id]
self._owned_slots[id(block)] = tail.slot
completed.append((block_start, block))
return completed
def read(
self, request_id: Hashable, token_start: int, token_end: int
) -> np.ndarray:
tail = self._requests.get(request_id)
assert tail is not None, (
f"auxiliary output buffer is missing request {request_id}"
)
local_start = token_start - tail.block_start
local_end = token_end - tail.block_start
assert 0 <= local_start < local_end <= tail.length, (
"auxiliary output range is unavailable: "
f"request={request_id}, range=[{token_start}, {token_end}), "
f"available=[{tail.block_start}, {tail.block_start + tail.length})"
)
return self._rows[tail.slot, local_start:local_end].copy()
def retain_block(self, rows: np.ndarray) -> np.ndarray:
"""Retain one unkeyed block after the current capture call."""
if id(rows) in self._owned_slots:
return rows
assert self._free_slots, "auxiliary output block pool is exhausted"
slot = self._free_slots.pop()
retained = self._rows[slot]
retained[...] = rows
self._owned_slots[id(retained)] = slot
return retained
def release_block(self, rows: np.ndarray) -> None:
slot = self._owned_slots.pop(id(rows), None)
if slot is not None:
self._free_slots.append(slot)
def discard(self, request_id: Hashable) -> None:
tail = self._requests.pop(request_id, None)
if tail is not None:
self._free_slots.append(tail.slot)
def reset(self) -> None:
self._requests.clear()
self._owned_slots.clear()
self._free_slots[:] = range(len(self._rows) - 1, -1, -1)