File size: 8,689 Bytes
2415c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 | # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0
"""Blocking and asynchronous device-to-host output read lifecycle."""
from __future__ import annotations
import itertools
import threading
from dataclasses import dataclass, field
from typing import Any
import torch
import ttnn
@dataclass(frozen=True)
class PendingRead:
"""One retained asynchronous host destination and its completion events."""
value: Any
events: tuple[Any, ...]
sequence: int
_owner: Any = field(repr=False, compare=False)
_completed: bool = field(default=False, repr=False, compare=False)
class OutputReader:
"""Own pending output-read destinations/events until completion or drain.
`DecodeRuntime.consume` uses `read` for blocking output. vLLM's
asynchronous path uses `submit`, returns the retained host value and
events to its scheduler, and later calls `complete`. Executor cleanup
calls `drain` to retire anything the scheduler did not complete.
"""
def __init__(self, mesh_device: Any):
self.mesh_device = mesh_device
self._sequences = itertools.count()
self._pending: dict[int, PendingRead] = {}
self._pending_by_value_id: dict[int, int] = {}
self._owner = object()
self._lock = threading.Lock()
# Public API
def read(self, value: Any, *, blocking: bool = True) -> Any | PendingRead:
"""Read a nested output payload to host.
Blocking calls return the completed host payload directly. Async calls
return a ``PendingRead`` whose ``value`` and ``events`` can be adapted to
the existing vLLM tuple contract while this reader retains ownership.
"""
return self._read(value, blocking=blocking)
def submit(self, value: Any) -> PendingRead:
"""Submit an asynchronous read and retain its resources."""
result = self._read(value, blocking=False)
assert isinstance(result, PendingRead)
return result
def read_synchronized(self, value: Any) -> Any:
"""Submit nested host copies and synchronize the device once."""
retained_destinations: list[Any] = []
try:
host_value, submitted = _read_to_host(
value,
blocking=False,
retained_destinations=retained_destinations,
)
if submitted:
ttnn.synchronize_device(self.mesh_device)
except BaseException:
_synchronize_after_failed_submission(self.mesh_device)
raise
return host_value
def complete(self, pending_or_value: PendingRead | Any) -> Any:
"""Synchronize and retire a pending read, returning its host payload.
Passing the exact unwrapped ``PendingRead.value`` supports compatibility
facades that return ``(host_value, events)`` to external schedulers.
Completion is idempotent.
"""
return self._complete_pending(pending_or_value)
def drain(self) -> None:
"""Complete every pending read; repeated drains are no-ops."""
with self._lock:
pending_reads = tuple(self._pending.values())
failures = []
for pending in pending_reads:
try:
self._complete_pending(pending)
except BaseException as error:
failures.append(error)
if failures:
raise RuntimeError(f"Failed to drain {len(failures)} pending output read(s)") from failures[0]
# Private implementation
def _read(self, value: Any, *, blocking: bool) -> Any | PendingRead:
retained_destinations: list[Any] = []
try:
host_value, submitted = _read_to_host(
value,
blocking=blocking,
retained_destinations=retained_destinations,
)
except BaseException:
if not blocking:
_synchronize_after_failed_submission(self.mesh_device)
raise
if blocking:
return host_value
sequence = next(self._sequences)
if not submitted:
return PendingRead(value=host_value, events=(), sequence=sequence, _owner=self._owner, _completed=True)
record_event = getattr(ttnn, "record_event", None)
event_synchronize = getattr(ttnn, "event_synchronize", None)
if not callable(record_event) or not callable(event_synchronize):
_synchronize_after_failed_submission(self.mesh_device)
raise RuntimeError("Asynchronous output reads require ttnn.record_event and ttnn.event_synchronize")
try:
event = record_event(self.mesh_device, 0)
except BaseException:
_synchronize_after_failed_submission(self.mesh_device)
raise
pending = PendingRead(value=host_value, events=(event,), sequence=sequence, _owner=self._owner)
with self._lock:
self._pending[sequence] = pending
self._pending_by_value_id[id(host_value)] = sequence
return pending
def _complete_pending(self, pending_or_value: PendingRead | Any) -> Any:
if isinstance(pending_or_value, PendingRead):
candidate = pending_or_value
if candidate._owner is not self._owner:
raise ValueError("PendingRead is not owned by this OutputReader")
sequence = candidate.sequence
else:
with self._lock:
sequence = self._pending_by_value_id.get(id(pending_or_value))
if sequence is None:
return pending_or_value
candidate = None
with self._lock:
pending = self._pending.get(sequence)
if pending is None:
if candidate is not None:
if candidate._completed:
return candidate.value
raise ValueError("PendingRead is not active in this OutputReader")
return pending_or_value
if candidate is not None and pending is not candidate:
raise ValueError("PendingRead is not owned by this OutputReader")
synchronize = getattr(ttnn, "event_synchronize", None)
if pending.events and not callable(synchronize):
raise RuntimeError("Pending output reads require ttnn.event_synchronize")
for event in pending.events:
synchronize(event)
with self._lock:
current = self._pending.pop(sequence, None)
if current is not None:
self._pending_by_value_id.pop(id(current.value), None)
object.__setattr__(current, "_completed", True)
return pending.value
def _read_to_host(
value: Any,
*,
blocking: bool,
retained_destinations: list[Any],
) -> tuple[Any, bool]:
if value is None:
return None, False
if isinstance(value, tuple):
converted = [
_read_to_host(item, blocking=blocking, retained_destinations=retained_destinations) for item in value
]
return tuple(item for item, _ in converted), any(submitted for _, submitted in converted)
if isinstance(value, list):
converted = [
_read_to_host(item, blocking=blocking, retained_destinations=retained_destinations) for item in value
]
return [item for item, _ in converted], any(submitted for _, submitted in converted)
if isinstance(value, dict):
converted = {
key: _read_to_host(item, blocking=blocking, retained_destinations=retained_destinations)
for key, item in value.items()
}
return (
{key: item for key, (item, _) in converted.items()},
any(submitted for _, submitted in converted.values()),
)
if isinstance(value, torch.Tensor):
return value.cpu(), False
if isinstance(value, ttnn.Tensor):
if value.storage_type() == ttnn.StorageType.HOST:
return value, False
host_value = value.cpu(blocking=blocking)
if not blocking:
retained_destinations.append(host_value)
return host_value, not blocking
cpu = getattr(value, "cpu", None)
if callable(cpu):
try:
host_value = cpu(blocking=blocking)
except TypeError:
return cpu(), False
if not blocking:
retained_destinations.append(host_value)
return host_value, not blocking
return value, False
def _synchronize_after_failed_submission(mesh_device: Any) -> None:
synchronize = getattr(ttnn, "synchronize_device", None)
if callable(synchronize):
synchronize(mesh_device)
|