File size: 70,104 Bytes
4d9b003 | 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 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 1067 1068 1069 1070 1071 1072 1073 1074 1075 1076 1077 1078 1079 1080 1081 1082 1083 1084 1085 1086 1087 1088 1089 1090 1091 1092 1093 1094 1095 1096 1097 1098 1099 1100 1101 1102 1103 1104 1105 1106 1107 1108 1109 1110 1111 1112 1113 1114 1115 1116 1117 1118 1119 1120 1121 1122 1123 1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137 1138 1139 1140 1141 1142 1143 1144 1145 1146 1147 1148 1149 1150 1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178 1179 1180 1181 1182 1183 1184 1185 1186 1187 1188 1189 1190 1191 1192 1193 1194 1195 1196 1197 1198 1199 1200 1201 1202 1203 1204 1205 1206 1207 1208 1209 1210 1211 1212 1213 1214 1215 1216 1217 1218 1219 1220 1221 1222 1223 1224 1225 1226 1227 1228 1229 1230 1231 1232 1233 1234 1235 1236 1237 1238 1239 1240 1241 1242 1243 1244 1245 1246 1247 1248 1249 1250 1251 1252 1253 1254 1255 1256 1257 1258 1259 1260 1261 1262 1263 1264 1265 1266 1267 1268 1269 1270 1271 1272 1273 1274 1275 1276 1277 1278 1279 1280 1281 1282 1283 1284 1285 1286 1287 1288 1289 1290 1291 1292 1293 1294 1295 1296 1297 1298 1299 1300 1301 1302 1303 1304 1305 1306 1307 1308 1309 1310 1311 1312 1313 1314 1315 1316 1317 1318 | # SPDX-License-Identifier: Apache-2.0
"""C02: ``TraceRunner`` -- every device stage of a port runs inside metal traces.
What it enforces (TT_PLATFORM.md section 3, REFERENCE_PATTERNS.md section 1.4):
1. **Persistent I/O before any capture.** Inputs, RT-dev parameters, state buffers and (optionally) output buffers
are allocated (and given defined initial values: never uninitialised index buffers) when they are added, i.e.
before the first capture. A variant returns either the tensors its last ops produce (allocated during capture,
valid while the trace lives) or persistent outputs written with ``ctx.write_output`` (stable addresses shared by
several variants, never overwritten by another variant's replay).
2. **Warm-up, then capture.** Every variant runs eagerly ``warmup_runs`` times (kernel JIT, program cache, prepared
conv weights) before *any* trace is captured. Adding a variant after a capture releases all traces, warms the new
one and recaptures everything, because a warm-up after a capture could allocate into a trace's freed
intermediates.
3. **Strict capture.** ``device.set_program_cache_misses_allowed(False)`` during capture (a miss raises with the op
name instead of aborting with "Writes are not supported during trace capture"); ``end_trace_capture`` runs in a
``finally`` block and a failed capture is released (an open capture left the process spinning in close_device:
RP section 1.4); the program-cache entry count must not change during capture.
4. **Variants.** Several traces keyed by name (shape buckets, modes, segments), chosen by the caller per run.
5. **RT-dev parameters.** Per-frame values that must not change the program (thresholds, timesteps, poses) live
in persistent device tensors refreshed before ``execute_trace``; unchanged values are not re-uploaded.
6. **1CQ / 2CQ.** With 2 CQs inputs are uploaded on CQ1 and ordered with events, exactly like
``common/tools/check_dispatch.py`` (CQ1 waits for the last trace, uploads, records; CQ0 waits, replays, records).
``stage_inputs=True`` uses the ``tt_cnn`` executor pattern instead: CQ1 writes a DRAM staging copy while the
previous trace still runs, and an eager copy on CQ0 moves it into the trace input before the replay. Every CQ0
write into a trace input outside the protocol (:meth:`TraceRunner.write_input`) re-records the event CQ1 waits for.
7. **State inside the trace.** ``ctx.write_state(name, value)`` is ``ttnn.copy(value, buffer)`` into a persistent
buffer (FLOAT32 / UINT32 / BFLOAT16 / ... all supported), or no op at all when ``value`` was computed straight
into ``ctx.write_target(name)`` (``output_tensor=``). ``pingpong=True`` states use two buffers and two traces
per variant (phase 0 reads A writes B, phase 1 reads B writes A); the phase flips after every run of a variant
that writes the ping-pong states. Per-stream state (PLAN.md D16): ``add_state(..., banks=n)`` +
:meth:`TraceRunner.save_state` / :meth:`TraceRunner.load_state`, with the policy in :class:`StreamBanks`.
8. **Readback.** One packed output (:func:`pack_outputs`) gives one D2H; reads go into preallocated host tensors,
on CQ0 or (segmented pipelines, 2CQ) on CQ1 after a host-side event wait.
9. **Alloc tracking.** With ``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``, ``ttnn.execute_trace``
refuses to replay over live unsafe buffers; the runner acknowledges the outputs of traces captured after the
first one (they may be overwritten by an earlier trace's replay: read outputs before running another variant).
Example::
runner = TraceRunner(device, num_command_queues=2)
runner.add_input("x", shape=(1, 1, 64, 64), dtype="bfloat16", layout=ttnn.TILE_LAYOUT)
runner.add_param("scale", 1.0) # fp32 [1,1,1,1], TILE
runner.add_variant("default", lambda ctx: ttnn.relu(ttnn.multiply(ctx["x"], ctx["scale"])))
runner.capture() # warm-up, then capture
y = runner("default", inputs={"x": x_np}, params={"scale": 0.5}) # upload, replay, read -> numpy
"""
from __future__ import annotations
import contextlib
import math
import time
from dataclasses import dataclass, field
from typing import Any, Callable, Dict, Iterator, List, Mapping, Optional, Sequence, Tuple
import numpy as np
from .io import InputError
from .tensors import TILE, dtype_name, round_up, to_host_tensor, to_numpy, ttnn_dtype
__all__ = [
"CQ_COMPUTE",
"CQ_INPUT",
"PackEntry",
"PackLayout",
"Packed",
"pack_outputs",
"SINGLE_ROW_MAX_ELEMS",
"PACK_ROW_ELEMS",
"TraceContext",
"TraceRunner",
"StreamBanks",
"alloc_tracking_enabled",
]
CQ_COMPUTE = 0 # programs, traces and (by default) readback
CQ_INPUT = 1 # host -> device uploads when the device has 2 command queues
def alloc_tracking_enabled() -> bool:
"""True when ``TT_METAL_TRACE_ALLOC_TRACKING=1`` was set before ttnn was imported (read from tt-metal)."""
try:
from ttnn.tools import trace_allocation_tracker as tracker
except ImportError:
return False
return bool(getattr(tracker, "TRACE_ALLOC_TRACKING", False))
# ----------------------------------------------------------------------------------------------- packing
# Packed layouts (PORT_LOG Q9 of the YOLOX port). A ROW_MAJOR tensor is stored page by page, one page per row, and
# the RM reshape / concat programs stage whole pages in L1 (reshape_rm_program_factory.cpp: 2 x the destination page
# per kernel copy when the pages are not 16-byte aligned), so a single-row pack of more than ~0.6 MB fails with
# "RM reshape dest staging does not fit in L1". Above SINGLE_ROW_MAX_ELEMS the pack is a [1, 1, rows, R] tensor of
# R = PACK_ROW_ELEMS elements per row: every RM page the packing programs touch stays below ROW_PAGE_MAX_BYTES,
# whatever the size of the outputs (tens of MB), and the readback is still ONE device-to-host copy.
SINGLE_ROW_MAX_ELEMS = 131072 # one [1, 1, 1, total] row up to 512 KiB of float32 (the 0.1.0 - 0.14.0 layout)
PACK_ROW_ELEMS = 1024 # elements per row of the multi-row layout (4 KiB float32 pages)
FLAT_MAX_BYTES = 64 << 10 # multi-row: tensors up to this size are flattened to one row, then cut into rows
ROW_PAGE_MAX_BYTES = 128 << 10 # multi-row: largest RM page (last dim x element size) a packed tensor may have
_ELEM_BYTES = {"float32": 4, "uint32": 4, "int32": 4, "bfloat16": 2, "uint16": 2, "uint8": 1}
@dataclass(frozen=True)
class PackEntry:
"""One tensor inside a packed readback: elements ``[offset, offset + numel)`` reshaped to ``shape``.
``pitch > 0`` (multi-row layout only): the tensor's rows (its last dim, ``shape[-1]`` elements)
are stored ``pitch`` elements apart, zero-padded, i.e. elements ``[offset, offset + numel // shape[-1] * pitch)``
viewed as ``(rows, pitch)`` hold it in their first ``shape[-1]`` columns. ``0``: contiguous."""
name: str
offset: int
numel: int
shape: Tuple[int, ...]
pitch: int = 0
def view(self, flat: np.ndarray) -> np.ndarray:
"""This tensor inside the flat readback ``flat`` (a view, strided when ``pitch`` is set)."""
if not self.pitch:
return flat[self.offset:self.offset + self.numel].reshape(self.shape)
cols = int(self.shape[-1]) if self.shape else 1
rows = self.numel // cols
return flat[self.offset:self.offset + rows * self.pitch].reshape(rows, self.pitch)[:, :cols].reshape(
self.shape)
@dataclass(frozen=True)
class PackLayout:
"""Host-side description of a packed output (built at capture time, applied at every read).
``entries`` index the packed tensor's elements in row-major order, so :meth:`unpack` does not depend on the
device shape ``(1, 1, rows, row_elems)``: one row of ``total`` elements (``rows == 1``) or rows of
``row_elems`` elements, every tensor starting on a row boundary (multi-row layout; an entry with a ``pitch``
stores its rows zero-padded to that many elements, see :class:`PackEntry`)."""
entries: Tuple[PackEntry, ...]
total: int
rows: int = 1
row_elems: int = 0 # 0: one row of ``total`` elements
@property
def shape(self) -> Tuple[int, int, int, int]:
"""Shape of the packed device tensor."""
return (1, 1, self.rows, self.row_elems or self.total)
def unpack(self, flat: Any) -> Dict[str, np.ndarray]:
"""Flat readback (any shape with ``total`` elements) -> ``{name: array}`` (views, no copies; the view of an
entry with a ``pitch`` is strided, not C-contiguous)."""
a = np.asarray(flat).reshape(-1)
if a.size != self.total:
raise ValueError(f"packed readback has {a.size} elements, layout expects {self.total}")
return {e.name: e.view(a) for e in self.entries}
@dataclass
class Packed:
"""A packed device tensor plus its layout; return it from a variant function to get ``{name: array}`` back."""
tensor: Any
layout: PackLayout
def _elem_bytes(dtype: Any) -> int:
return _ELEM_BYTES.get(dtype_name(dtype), 4)
def _row_major(ttnn: Any, t: Any, dtype: Any) -> Any:
if t.dtype != dtype:
t = ttnn.typecast(t, dtype)
if t.layout != ttnn.ROW_MAJOR_LAYOUT:
t = ttnn.to_layout(t, ttnn.ROW_MAJOR_LAYOUT)
return t
def _flat_rows(ttnn: Any, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int) -> Tuple[Any, int]:
"""A small RM tensor -> one row -> zero-padded to whole rows -> ``[1, 1, k, row_elems]``."""
padded = round_up(numel, row_elems)
if shape != (1, 1, 1, numel):
t = ttnn.reshape(t, (1, 1, 1, numel))
if padded != numel:
t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0)
if padded != row_elems:
t = ttnn.reshape(t, (1, 1, padded // row_elems, row_elems))
return t, padded
def _fill_rows(rows: int, cols: int, row_elems: int) -> int:
"""Rows of ``cols`` elements a ``[rows, cols]`` tensor is zero-padded to so that it fills whole packed rows."""
return round_up(rows, row_elems // math.gcd(cols, row_elems))
def _row_pitches(cols: int, row_elems: int, elem: int) -> List[int]:
"""Row pitches > ``cols`` worth trying: ``cols`` rounded up to every power of two dividing ``row_elems`` (rows of
that pitch fill a packed row every ``row_elems / gcd`` rows) and to ``row_elems`` (every row fills whole packed
rows), within the page budget."""
steps, m = {row_elems}, 2
while row_elems % m == 0:
steps.add(m)
m *= 2
return sorted({p for p in (round_up(cols, m) for m in steps) if p != cols and p * elem <= ROW_PAGE_MAX_BYTES})
def _pack_rows_segment(ttnn: Any, name: str, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int,
elem: int) -> Tuple[Any, int, int]:
"""One ROW_MAJOR tensor -> ``[1, 1, k, row_elems]`` holding its elements in order, zero-padded to whole rows.
Returns ``(segment, k * row_elems, pitch)`` (``pitch``: see :class:`PackEntry`). No RM page above
``ROW_PAGE_MAX_BYTES`` is created or reshaped."""
cols = int(shape[-1]) if shape else 1
rows = numel // cols
if cols * elem <= ROW_PAGE_MAX_BYTES:
# [rows, cols] is a view of the RM tensor; rows * cols fills whole packed rows when rows % unit == 0
rows_p = _fill_rows(rows, cols, row_elems)
if rows_p != rows and numel * elem <= FLAT_MAX_BYTES:
return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0) # small: pad < row_elems elements, not rows
if rows_p != rows:
# Appending zero rows costs up to row_elems / gcd(cols, row_elems) - 1 rows: 80 MB for a [1, 1, 2, 20001]
# fp32 tensor. Take one flat row (when it fits the page budget) or rows zero-padded to a pitch instead
# when that is smaller by more than 1/8 of the tensor (typical shapes keep the contiguous layout).
options = [(p * _fill_rows(rows, p, row_elems), p) for p in _row_pitches(cols, row_elems, elem)]
if round_up(numel, row_elems) * elem <= ROW_PAGE_MAX_BYTES:
options.append((round_up(numel, row_elems), 0))
if options and rows_p * cols - min(options)[0] > numel // 8:
size, pitch = min(options) # ties: the flat row, then the narrower pitch
if not pitch:
return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0)
prows = size // pitch
if shape != (1, 1, rows, cols):
t = ttnn.reshape(t, (1, 1, rows, cols))
t = ttnn.pad(t, [(0, 0), (0, 0), (0, prows - rows), (0, pitch - cols)], 0.0)
if pitch != row_elems:
t = ttnn.reshape(t, (1, 1, size // row_elems, row_elems))
return t, size, pitch
# zero rows appended, then one RM reshape into rows of row_elems (source pages of cols elements, destination
# pages of row_elems elements: both small)
if shape != (1, 1, rows, cols):
t = ttnn.reshape(t, (1, 1, rows, cols))
if rows_p != rows:
t = ttnn.pad(t, [(0, 0), (0, 0), (0, rows_p - rows), (0, 0)], 0.0)
if cols != row_elems:
t = ttnn.reshape(t, (1, 1, rows_p * cols // row_elems, row_elems))
return t, rows_p * cols, 0
if rows != 1:
raise ValueError(f"pack_outputs: {name!r} {shape} has rows of {cols * elem} B; the multi-row layout reads "
f"RM rows of at most {ROW_PAGE_MAX_BYTES} B: give it a narrower last dim (e.g. "
f"[1, 1, -1, {row_elems}]) before packing")
# one wide row (a flat vector): cut it into chunks of whole packed rows, each well inside the L1 budget
if shape != (1, 1, 1, numel):
t = ttnn.reshape(t, (1, 1, 1, numel))
width = max(row_elems, (ROW_PAGE_MAX_BYTES // elem) // row_elems * row_elems)
parts, total = [], 0
for start in range(0, numel, width):
stop = min(start + width, numel)
chunk = ttnn.slice(t, [0, 0, 0, start], [1, 1, 1, stop])
chunk, padded = _flat_rows(ttnn, chunk, (1, 1, 1, stop - start), stop - start, row_elems)
parts.append(chunk)
total += padded
return (parts[0] if len(parts) == 1 else ttnn.concat(parts, dim=2)), total, 0
def pack_outputs(tensors: Mapping[str, Any], *, dtype: Any = "float32", align: int = 32,
row_elems: Optional[int] = None) -> Packed:
"""Pack several device tensors into ONE ROW_MAJOR tensor so the host reads them with one device-to-host copy.
Each tensor is typecast to ``dtype`` (float32 default: exact for bf16 and for integers below 2**24) and converted
to ROW_MAJOR; call this inside the variant function (the ops become part of the trace) and return the result.
:meth:`PackLayout.unpack` (done by ``TraceRunner.read``) gives ``{name: array}`` back in the original shapes.
Device layout: while the packed total is at most :data:`SINGLE_ROW_MAX_ELEMS`, one ``[1, 1, 1, total]`` row (each
tensor flattened and zero-padded to ``align`` elements, then concatenated): a ``[1, 1, N, 1]`` RM readback is N
pages and costs milliseconds, one row is one page (RP section 1.4). Above it (or with ``row_elems=R``) a
``[1, 1, rows, R]`` tensor (default R = :data:`PACK_ROW_ELEMS`): each tensor is zero-padded to whole rows of R
elements and the segments are concatenated on the row dim, so no RM page exceeds :data:`ROW_PAGE_MAX_BYTES`
and outputs of tens of MB pack (a single row of more than ~0.6 MB does not fit the RM reshape's L1 staging:
YOLOX PORT_LOG Q9). A tensor is padded with zero rows (contiguous), or, when that wastes more than 1/8 of it
(an awkward last dim with few rows: e.g. ``[1, 1, 2, 20001]`` would take 1024 rows), flattened (up to the page
budget) or stored with its rows zero-padded to a wider pitch (:class:`PackEntry` ``pitch``). A last dim wider
than ``ROW_PAGE_MAX_BYTES`` is accepted for a flat vector (``[1, 1, 1, N]``, cut into chunks with
``ttnn.slice``); any other tensor needs a narrower last dim (``ValueError``)."""
import ttnn
if not tensors:
raise ValueError("pack_outputs needs at least one tensor")
if align <= 0:
raise ValueError("align must be positive")
if row_elems is not None and (int(row_elems) <= 0 or int(row_elems) % TILE):
raise ValueError(f"row_elems={row_elems}: expected a positive multiple of {TILE}")
dtype = ttnn_dtype(dtype)
elem = _elem_bytes(dtype)
items = [(str(name), t, tuple(int(s) for s in t.shape)) for name, t in tensors.items()]
empty = [name for name, _, shape in items if math.prod(shape) == 0]
if empty:
raise ValueError(f"pack_outputs: {empty} have no elements")
single_total = sum(round_up(math.prod(shape), align) for _, _, shape in items)
if row_elems is None and single_total <= SINGLE_ROW_MAX_ELEMS:
segments, entries, offset = [], [], 0
for name, t, shape in items: # the single-row layout of ttaw 0.1.0 - 0.14.0
numel = math.prod(shape)
t = _row_major(ttnn, t, dtype)
if shape != (1, 1, 1, numel):
t = ttnn.reshape(t, (1, 1, 1, numel))
padded = round_up(numel, align)
if padded != numel:
t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0)
segments.append(t)
entries.append(PackEntry(name, offset, numel, shape))
offset += padded
packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=-1)
return Packed(packed, PackLayout(tuple(entries), offset))
r = int(row_elems or PACK_ROW_ELEMS)
segments, entries, offset = [], [], 0
for name, t, shape in items:
numel = math.prod(shape)
t, padded, pitch = _pack_rows_segment(ttnn, name, _row_major(ttnn, t, dtype), shape, numel, r, elem)
segments.append(t)
entries.append(PackEntry(name, offset, numel, shape, pitch))
offset += padded
packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=2)
return Packed(packed, PackLayout(tuple(entries), offset, rows=offset // r, row_elems=r))
# --------------------------------------------------------------------------------------------- internals
@dataclass
class _Slot:
"""A persistent device tensor (input, parameter or state) allocated before any capture."""
name: str
kind: str # "input" | "param" | "state" | "output"
shape: Tuple[int, ...]
dtype: Any
layout: Any
memory_config: Any
buffers: List[Any] # 1 buffer, or 2 for a ping-pong state
init: Any # host tensor holding the initial value (states, params) / warm-up value (inputs)
staging: Any = None # DRAM staging copy (2CQ stage_inputs mode)
stage_fn: Optional[Callable[[Any, Any], Any]] = None
last_value: Optional[bytes] = None # params: bytes of the last uploaded value (skip unchanged uploads)
banks: List[Any] = field(default_factory=list) # states: save / load buffers (D16 per-stream state banks)
@property
def pingpong(self) -> bool:
return len(self.buffers) == 2
def tensors(self) -> List[Any]:
"""Every device tensor the slot owns."""
return self.buffers + ([self.staging] if self.staging is not None else []) + self.banks
@dataclass
class _Variant:
name: str
fn: Callable[["TraceContext"], Any]
warmup_runs: int
@dataclass
class _Trace:
variant: str
phase: int
trace_id: Any
outputs: Any
leaves: List[Tuple[Tuple, Any, Optional[PackLayout]]] # (path, device tensor, pack layout or None)
steps: bool # writes the ping-pong states (flips the phase)
capture_ms: float
host_buffers: Optional[List[Any]] = None
done_event: Any = None
def _flatten(obj: Any, path: Tuple = ()) -> Iterator[Tuple[Tuple, Any]]:
"""Leaves of a tensor / Packed / list / tuple / dict structure with their paths."""
if isinstance(obj, Mapping):
for k, v in obj.items():
yield from _flatten(v, path + (k,))
elif isinstance(obj, (list, tuple)):
for i, v in enumerate(obj):
yield from _flatten(v, path + (i,))
else:
yield path, obj
def _rebuild(obj: Any, values: Dict[Tuple, Any], path: Tuple = ()) -> Any:
if isinstance(obj, Mapping):
return {k: _rebuild(v, values, path + (k,)) for k, v in obj.items()}
if isinstance(obj, (list, tuple)):
return type(obj)(_rebuild(v, values, path + (i,)) for i, v in enumerate(obj))
return values[path]
def _buffer_address(t: Any) -> Optional[int]:
try:
return int(t.buffer_address())
except Exception: # noqa: BLE001 -- host tensor, deallocated tensor or a fake
return None
def _same_buffer(a: Any, b: Any) -> bool:
"""``a`` is ``b`` or a tensor over the same device buffer."""
if a is b:
return True
addr = _buffer_address(a)
return addr is not None and addr == _buffer_address(b)
def _check_copy(what: str, value: Any, slot: "_Slot") -> None:
"""The preconditions of ``ttnn.copy(value, <slot buffer>)``, checked in Python so a mistake is a clear error
instead of a TT_FATAL inside an open capture: same logical shape and layout; a dtype change only in TILE."""
import ttnn
shape = tuple(int(s) for s in value.shape)
if shape != slot.shape:
raise ValueError(f"{what}: value shape {shape} != {slot.shape}")
if value.layout != slot.layout:
raise ValueError(f"{what}: value layout {value.layout} != {slot.layout} (ttnn.copy keeps the layout)")
if value.dtype != slot.dtype and slot.layout != ttnn.TILE_LAYOUT:
raise ValueError(f"{what}: dtype {value.dtype} -> {slot.dtype} needs TILE layout (ttnn.copy)")
class TraceContext(Mapping):
"""What a variant function receives: persistent tensors by name (``ctx["x"]``), and state writes.
``ctx[name]`` returns an input, a parameter, a persistent output buffer, or the buffer a state is *read* from in
this phase. ``ctx.write_state(name, value)`` copies ``value`` into the buffer the state is *written* to (in-place
states: the same buffer, so read everything you need from it before writing); ``ctx.write_output(name, value)``
copies into a persistent output and returns it."""
def __init__(self, runner: "TraceRunner", variant: str, phase: int, capturing: bool):
self._runner = runner
self.variant = variant
self.phase = phase
self.capturing = capturing
self.writes: set = set()
@property
def device(self):
return self._runner.device
def __getitem__(self, name: str):
slot = self._runner._slot(name)
if slot.kind == "state":
return self.state(name)
return slot.buffers[0]
def __iter__(self) -> Iterator[str]:
return iter(self._runner._slots)
def __len__(self) -> int:
return len(self._runner._slots)
def state(self, name: str):
"""The buffer state ``name`` is read from in this phase."""
slot = self._runner._slot(name, "state")
return slot.buffers[self.phase] if slot.pingpong else slot.buffers[0]
def write_target(self, name: str):
"""The buffer state ``name`` is written to in this phase (ping-pong: the other buffer; in place: the same
one). Pass it as ``output_tensor=`` of the op that produces the new state and then call
``write_state(name, it)``: no copy program is traced (probe P13: saves one program per state per frame)."""
slot = self._runner._slot(name, "state")
return slot.buffers[1 - self.phase] if slot.pingpong else slot.buffers[0]
def write_state(self, name: str, value) -> None:
"""``ttnn.copy(value, <write buffer>)``: same logical shape and layout as the state (dtype may differ in
TILE layout). The copy is part of the trace; it is skipped when ``value`` already is the write buffer
(:meth:`write_target`)."""
import ttnn
slot = self._runner._slot(name, "state")
target = self.write_target(name)
_check_copy(f"state {name!r}", value, slot)
if not _same_buffer(value, target):
ttnn.copy(value, target)
self.writes.add(name)
def write_output(self, name: str, value):
"""``ttnn.copy(value, <persistent output>)`` (part of the trace); returns the output buffer, which the
variant can return as (part of) its outputs. Skipped when ``value`` already is that buffer."""
import ttnn
slot = self._runner._slot(name, "output")
_check_copy(f"output {name!r}", value, slot)
if not _same_buffer(value, slot.buffers[0]):
ttnn.copy(value, slot.buffers[0])
return slot.buffers[0]
class TraceRunner:
"""Persistent device I/O + warm-up + capture + replay of one model's traced stages (see the module docstring).
Args:
device: an open ttnn device.
num_command_queues: 1 or 2. ``None`` uses what :func:`.device.open_device` recorded (1 if unknown).
warmup_runs: eager runs of each variant (and phase) before any capture.
stage_inputs: 2CQ only: upload into DRAM staging buffers and copy them into the trace inputs with an eager
op on CQ0 (the upload of frame k+1 then overlaps the replay of frame k).
forbid_cache_misses: call ``device.set_program_cache_misses_allowed(False)`` during capture.
alloc_tracking: ``True`` raises unless the process tracks trace allocations
(``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``). Whenever tracking is active, the outputs
of traces captured after the first are acknowledged as corruptible (see the module docstring).
name: label used in messages and ``describe()``.
"""
def __init__(self, device, *, num_command_queues: Optional[int] = None, warmup_runs: int = 1,
stage_inputs: bool = False, forbid_cache_misses: bool = True, alloc_tracking: Optional[bool] = None,
name: str = "model"):
from .device import open_info
opened = open_info(device).get("num_command_queues")
if num_command_queues is None:
num_command_queues = opened or 1
if num_command_queues not in (1, 2):
raise ValueError(f"num_command_queues={num_command_queues}: expected 1 or 2")
if opened is not None and num_command_queues > opened:
raise ValueError(f"the device was opened with {opened} command queue(s); cannot run {num_command_queues}")
if stage_inputs and num_command_queues != 2:
raise ValueError("stage_inputs=True needs num_command_queues=2")
if warmup_runs < 1:
raise ValueError("warmup_runs must be >= 1 (capture needs a warm program cache)")
tracking = alloc_tracking_enabled()
if alloc_tracking and not tracking:
raise RuntimeError("alloc_tracking=True but trace allocation tracking is off: export "
"TT_METAL_TRACE_ALLOC_TRACKING=1 (optionally TT_METAL_TRACE_ALLOC_TRACEBACKS=1) "
"before Python imports ttnn")
self.device = device
self.name = name
self.num_command_queues = int(num_command_queues)
self.warmup_runs = int(warmup_runs)
self.stage_inputs = bool(stage_inputs)
self.forbid_cache_misses = bool(forbid_cache_misses)
self.alloc_tracking = bool(tracking)
self._slots: Dict[str, _Slot] = {}
self._variants: Dict[str, _Variant] = {}
self._traces: Dict[Tuple[str, int], _Trace] = {}
self._phase = 0
self._last_phase: Dict[str, int] = {}
self._last_variant: Optional[str] = None
self._pending_params: Dict[str, Any] = {}
self._op_event = None
self._stage_free_event = None
self._captures = 0
self._capturing = False
self._eager_warmups: List[Callable[[], Any]] = []
self._eager_warmed = False
self._closed = False
self.timings_ms: Dict[str, Dict[str, float]] = {}
if hasattr(device, "enable_program_cache"):
device.enable_program_cache()
# ------------------------------------------------------------------------------------- registration
def _check_registration(self, what: str) -> None:
self._check_open()
if self._capturing:
raise RuntimeError(f"cannot add {what} while a variant is being warmed up or captured: register "
"inputs, params, states, outputs and variants before capture()")
def _new_slot(self, name: str, kind: str, init: Any, shape: Optional[Sequence[int]], dtype: Any, layout: Any,
memory_config: Any, n_buffers: int, stage_fn: Optional[Callable], n_banks: int = 0) -> _Slot:
import ttnn
self._check_registration(f"{kind} {name!r}")
if name in self._slots:
raise ValueError(f"{name!r} is already registered (as {self._slots[name].kind})")
if self._traces:
raise RuntimeError(f"cannot add {kind} {name!r} after capture: persistent tensors must exist before the "
"first capture (release() and rebuild)")
layout = ttnn.ROW_MAJOR_LAYOUT if layout is None else layout
memory_config = ttnn.DRAM_MEMORY_CONFIG if memory_config is None else memory_config
dtype = ttnn_dtype(dtype)
if init is None:
if shape is None:
raise ValueError(f"{kind} {name!r}: give init= or shape=")
init = np.zeros(tuple(shape), np.float32 if dtype_name(dtype) in ("float32", "bfloat16", "bfloat8_b",
"bfloat4_b") else np.int64)
if isinstance(init, ttnn.Tensor):
host = to_host_tensor(init, dtype, layout, shape=shape)
else:
arr = to_numpy(init) if hasattr(init, "detach") else np.asarray(init)
if shape is not None:
arr = np.broadcast_to(arr, tuple(shape))
host = to_host_tensor(arr, dtype, layout)
shape_t = tuple(int(s) for s in host.shape)
def allocate(config: Any):
buf = ttnn.allocate_tensor_on_device(ttnn.Shape(list(shape_t)), dtype, layout, self.device, config)
ttnn.copy_host_to_device_tensor(host, buf, cq_id=CQ_COMPUTE) # defined contents, never garbage
return buf
buffers = [allocate(memory_config) for _ in range(n_buffers)]
staging = allocate(ttnn.DRAM_MEMORY_CONFIG) if self.stage_inputs and kind in ("input", "param") else None
banks = [allocate(ttnn.DRAM_MEMORY_CONFIG) for _ in range(n_banks)]
slot = _Slot(name, kind, shape_t, dtype, layout, memory_config, buffers, host, staging,
stage_fn or (lambda src, dst: ttnn.copy(src, dst)), banks=banks)
self._slots[name] = slot
return slot
def add_input(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None,
dtype: Any = "bfloat16", layout: Any = None, memory_config: Any = None,
stage_fn: Optional[Callable[[Any, Any], Any]] = None):
"""A persistent trace input (default ROW_MAJOR in DRAM). ``init`` (array / torch / ttnn host tensor) is the
initial content and the warm-up input; zeros of ``shape`` otherwise -- give real data when the graph
gathers with these values (garbage indices can hang the chip). ``stage_fn(staging, persistent)`` replaces
the eager ``ttnn.copy`` in ``stage_inputs`` mode (e.g. a reshard into a sharded L1 input). Returns the
device tensor."""
return self._new_slot(name, "input", init, shape, dtype, layout, memory_config, 1, stage_fn).buffers[0]
def add_param(self, name: str, value: Any = 0.0, *, shape: Sequence[int] = (1, 1, 1, 1), dtype: Any = "float32",
layout: Any = None, memory_config: Any = None):
"""An RT-dev parameter: a persistent device tensor (default fp32 ``[1, 1, 1, 1]`` TILE, which broadcasts in
ttnn binary ops) refreshed before a replay when ``run(params=...)`` / :meth:`set_params` changes it."""
import ttnn
layout = ttnn.TILE_LAYOUT if layout is None else layout
slot = self._new_slot(name, "param", np.broadcast_to(np.asarray(value), tuple(shape)), tuple(shape), dtype,
layout, memory_config, 1, None)
slot.last_value = self._param_bytes(slot, value)
return slot.buffers[0]
def add_state(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None, dtype: Any = "float32",
layout: Any = None, memory_config: Any = None, pingpong: bool = False, banks: int = 0):
"""Temporal state kept on the device across replays (memory queues, previous BEV, ring buffers).
In-place (default): one buffer, read via ``ctx[name]`` and written via ``ctx.write_state``. ``pingpong``:
two buffers and two traces per variant. ``init`` is restored by :meth:`reset_state`. ``banks``: extra DRAM
buffers of the same spec for :meth:`save_state` / :meth:`load_state` (one per stream id of a
:class:`StreamBanks`, PLAN.md D16), allocated now because nothing may be allocated after a capture.
Returns the buffer(s)."""
import ttnn
if banks < 0:
raise ValueError("banks must be >= 0")
layout = ttnn.TILE_LAYOUT if layout is None else layout
slot = self._new_slot(name, "state", init, shape, dtype, layout, memory_config, 2 if pingpong else 1, None,
n_banks=int(banks))
return tuple(slot.buffers) if pingpong else slot.buffers[0]
def add_output(self, name: str, *, shape: Sequence[int], dtype: Any = "float32", layout: Any = None,
memory_config: Any = None):
"""A persistent output buffer (default TILE in DRAM, zeros) allocated before any capture. Variants write it
with ``ctx.write_output(name, value)`` (one traced ``ttnn.copy``) and return it; its address survives
recaptures and is shared by every variant (e.g. shape buckets with one readback). Returns the buffer."""
import ttnn
layout = ttnn.TILE_LAYOUT if layout is None else layout
return self._new_slot(name, "output", None, shape, dtype, layout, memory_config, 1, None).buffers[0]
def add_variant(self, name: str, fn: Callable[[TraceContext], Any], *, warmup_runs: Optional[int] = None) -> None:
"""Register a traced function. ``fn(ctx)`` runs ttnn ops on ``ctx[...]`` tensors and returns a device tensor,
a :class:`Packed`, or a list / tuple / dict of them (the persistent trace outputs). No host I/O, no
``synchronize``, no torch ops inside ``fn``: it is called for warm-up and then recorded."""
self._check_registration(f"variant {name!r}")
if name in self._variants:
raise ValueError(f"variant {name!r} already exists")
self._variants[name] = _Variant(name, fn, self.warmup_runs if warmup_runs is None else int(warmup_runs))
def add_eager_warmup(self, fn: Callable[[], Any]) -> None:
"""Register eager device work that the model runs *between* replays (host-fallback glue, an eager layout
change of a read-back tensor, ...): ``fn()`` runs once in the first :meth:`capture`, before any trace is
captured, so its programs are compiled -- and their kernel binaries allocated in DRAM -- before the first
capture. A program compiled after a capture shares the address space of the traces' freed intermediates and
a replay can overwrite its binaries (tt-metal ``tech_reports/.../TraceCorrectness.md``: corruption or a
hang); ``TT_METAL_TRACE_ALLOC_TRACKING=1`` reports it."""
self._check_registration("an eager warm-up")
if self._traces or self._eager_warmed:
raise RuntimeError("add eager warm-ups before the first capture (release() and rebuild)")
self._eager_warmups.append(fn)
# ------------------------------------------------------------------------------------------ capture
@property
def phases(self) -> int:
"""2 when a ping-pong state exists (two traces per variant), else 1."""
return 2 if any(s.pingpong for s in self._slots.values()) else 1
@property
def captured(self) -> bool:
return bool(self._variants) and all((v, p) in self._traces for v in self._variants for p in range(self.phases))
def capture(self) -> None:
"""Warm up and capture every registered variant that has no trace yet (idempotent). If traces already exist
and a new variant was added, all traces are released, the new variants warmed up, and all recaptured."""
import ttnn
self._check_open()
if not self._variants:
raise RuntimeError("no variant registered (add_variant)")
pending = [v for v in self._variants if any((v, p) not in self._traces for p in range(self.phases))]
if not pending:
return
if self._traces:
ttnn.synchronize_device(self.device) # no replay of a trace being released is still in flight
self._release_traces()
warm, pending = pending, list(self._variants)
else:
warm = pending
self._capturing = True
try:
if not self._eager_warmed:
self._warm_eager_programs()
for vname in warm:
t0 = time.perf_counter()
variant = self._variants[vname]
for phase in range(self.phases):
for _ in range(variant.warmup_runs):
ctx = TraceContext(self, vname, phase, capturing=False)
outputs = variant.fn(ctx)
ttnn.synchronize_device(self.device)
self._free_transient(outputs)
self.timings_ms.setdefault(vname, {})["warmup"] = (time.perf_counter() - t0) * 1e3
self._phase = 0
self.reset_state()
for vname in pending:
self.timings_ms.setdefault(vname, {})["capture"] = 0.0
for phase in range(self.phases):
self._capture_one(vname, phase)
finally:
self._capturing = False
ttnn.synchronize_device(self.device)
if self.num_command_queues == 2:
self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
self._stage_free_event = self._op_event
def _warm_eager_programs(self) -> None:
"""Compile the runner's own eager programs before the first capture: the staging copies (``stage_inputs``),
the state <-> bank copies (``banks``) and the registered eager warm-ups. Contents are kept: a staging buffer
mirrors its input (:meth:`write_input` writes both), and a state and its first bank hold the same value
before the states are reset ahead of the captures."""
import ttnn
for slot in self._slots.values():
if slot.staging is not None:
slot.stage_fn(slot.staging, slot.buffers[0])
if slot.banks:
ttnn.copy(slot.buffers[0], slot.banks[0])
ttnn.copy(slot.banks[0], slot.buffers[0])
for fn in self._eager_warmups:
fn()
ttnn.synchronize_device(self.device)
self._eager_warmed = True
@contextlib.contextmanager
def _eager(self, what: str) -> Iterator[None]:
"""Eager runner work after a capture must not compile a program (see :meth:`add_eager_warmup`)."""
dev = self.device
count = getattr(dev, "num_program_cache_entries", None)
before = count() if (self._traces and count is not None) else None
yield
if before is not None and count() != before:
raise RuntimeError(f"{self.name}: {what} compiled a new program after capture; its kernel binaries share "
"DRAM with the traces' freed intermediates, so a replay can overwrite them. Run it "
"before the first capture (add_eager_warmup) or keep the tensor specs of the warmed "
"program")
def _capture_one(self, vname: str, phase: int) -> None:
import ttnn
dev = self.device
variant = self._variants[vname]
ctx = TraceContext(self, vname, phase, capturing=True)
entries_before = dev.num_program_cache_entries() if hasattr(dev, "num_program_cache_entries") else None
forbid = self.forbid_cache_misses and hasattr(dev, "set_program_cache_misses_allowed")
t0 = time.perf_counter()
if forbid:
dev.set_program_cache_misses_allowed(False)
trace_id = ttnn.begin_trace_capture(dev, cq_id=CQ_COMPUTE)
ok = ended = False
try:
outputs = variant.fn(ctx)
ok = True
finally:
try:
ttnn.end_trace_capture(dev, trace_id, cq_id=CQ_COMPUTE)
ended = True
finally:
if forbid:
dev.set_program_cache_misses_allowed(True)
if not (ok and ended):
self._safe_release(trace_id)
try:
entries_after = dev.num_program_cache_entries() if entries_before is not None else None
if entries_before is not None and entries_after != entries_before:
raise RuntimeError(f"{self.name}/{vname}: the program cache grew during capture ({entries_before} -> "
f"{entries_after}); warm-up does not cover the traced graph")
leaves = []
for path, leaf in _flatten(outputs):
if isinstance(leaf, Packed):
leaves.append((path, leaf.tensor, leaf.layout))
elif isinstance(leaf, ttnn.Tensor):
leaves.append((path, leaf, None))
else:
raise TypeError(f"{self.name}/{vname}: output {path} is {type(leaf).__name__}, expected a device "
"tensor or Packed")
if not leaves:
raise ValueError(f"{self.name}/{vname}: the variant returned no output tensor")
pingpong = {s.name for s in self._slots.values() if s.pingpong}
written = ctx.writes & pingpong
if written and written != pingpong:
raise RuntimeError(f"{self.name}/{vname}: writes ping-pong states {sorted(written)} but not "
f"{sorted(pingpong - written)}; a stepping variant must write all of them")
except BaseException:
self._safe_release(trace_id)
raise
if self.alloc_tracking and self._captures > 0:
from ttnn.tools import trace_allocation_tracker as tracker
keep = self._persistent_addresses()
for _, tensor, _ in leaves:
if _buffer_address(tensor) not in keep:
tracker.acknowledge_corruptible(tensor)
self._captures += 1
ms = (time.perf_counter() - t0) * 1e3
self._traces[(vname, phase)] = _Trace(vname, phase, trace_id, outputs, leaves, bool(written), ms)
self.timings_ms[vname]["capture"] += ms
def _persistent_addresses(self) -> set:
addrs = set()
for slot in self._slots.values():
for t in slot.tensors():
a = _buffer_address(t)
if a is not None:
addrs.add(a)
return addrs
def _deallocate_outputs(self, tensors: Iterator[Any]) -> None:
"""Deallocate op-produced output tensors, never a persistent buffer or a view of one. ``force=False``: a
tensor sharing its device memory with another owner (a view of a model weight returned as an output) is
left to its owners -- ``ttnn.deallocate`` forces by default and would free the weight."""
import ttnn
keep = self._persistent_addresses()
seen = set()
for tensor in tensors:
if not isinstance(tensor, ttnn.Tensor) or id(tensor) in seen:
continue
seen.add(id(tensor))
if tensor.is_allocated() and _buffer_address(tensor) not in keep:
ttnn.deallocate(tensor, False)
def _free_transient(self, outputs: Any) -> None:
"""Deallocate warm-up / eager outputs (see :meth:`_deallocate_outputs`)."""
self._deallocate_outputs(leaf.tensor if isinstance(leaf, Packed) else leaf for _, leaf in _flatten(outputs))
def _safe_release(self, trace_id) -> None:
import ttnn
try:
ttnn.release_trace(self.device, trace_id)
except Exception: # noqa: BLE001 -- best effort on an error path; the original error is re-raised
pass
# -------------------------------------------------------------------------------------------- inputs
def _slot(self, name: str, kind: Optional[str] = None) -> _Slot:
slot = self._slots.get(name)
if slot is None:
raise KeyError(f"{self.name}: no input / param / state named {name!r}")
if kind is not None and slot.kind != kind:
raise KeyError(f"{self.name}: {name!r} is a {slot.kind}, not a {kind}")
return slot
@staticmethod
def _param_array(slot: _Slot, value: Any) -> np.ndarray:
dt = np.float32 if dtype_name(slot.dtype) in ("float32", "bfloat16", "bfloat8_b", "bfloat4_b") else np.int64
return np.ascontiguousarray(np.broadcast_to(np.asarray(value, dtype=dt), slot.shape))
def _param_bytes(self, slot: _Slot, value: Any) -> bytes:
return self._param_array(slot, value).tobytes()
def set_params(self, **values: Any) -> None:
"""Queue RT-dev parameter values for the next run (uploaded only if they changed)."""
for name in values:
self._slot(name, "param")
self._pending_params.update(values)
def _collect_uploads(self, inputs: Optional[Mapping[str, Any]],
params: Optional[Mapping[str, Any]]) -> List[Tuple[_Slot, Any, Optional[bytes]]]:
uploads = []
for name, value in (inputs or {}).items():
slot = self._slot(name, "input")
uploads.append((slot, to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape), None))
merged = dict(self._pending_params)
merged.update(params or {})
for name, value in merged.items():
slot = self._slot(name, "param")
arr = self._param_array(slot, value)
key = arr.tobytes()
if key != slot.last_value:
uploads.append((slot, to_host_tensor(arr, slot.dtype, slot.layout), key))
return uploads
def _ensure_events(self) -> None:
"""2CQ: the CQ0 events CQ1 waits for exist (``capture()`` records them; a partially failed capture, which
leaves earlier variants runnable, does not get that far)."""
import ttnn
if self._op_event is None:
self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
if self._stage_free_event is None:
self._stage_free_event = self._op_event
def _enqueue_uploads(self, uploads: List[Tuple[_Slot, Any, Optional[bytes]]]) -> None:
import ttnn
if uploads:
if self.num_command_queues == 1:
for slot, host, _ in uploads:
ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
elif self.stage_inputs:
self._ensure_events()
ttnn.wait_for_event(CQ_INPUT, self._stage_free_event)
for slot, host, _ in uploads:
ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_INPUT)
written = ttnn.record_event(self.device, CQ_INPUT)
ttnn.wait_for_event(CQ_COMPUTE, written)
with self._eager("a stage_fn copy"):
for slot, _, _ in uploads:
slot.stage_fn(slot.staging, slot.buffers[0])
self._stage_free_event = ttnn.record_event(self.device, CQ_COMPUTE)
else:
self._ensure_events()
ttnn.wait_for_event(CQ_INPUT, self._op_event)
for slot, host, _ in uploads:
ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_INPUT)
written = ttnn.record_event(self.device, CQ_INPUT)
ttnn.wait_for_event(CQ_COMPUTE, written)
for slot, _, key in uploads:
if key is not None:
slot.last_value = key
self._pending_params.clear()
def write_input(self, name: str, value: Any) -> None:
"""Upload ``value`` into input ``name`` now (CQ0, outside any trace): e.g. a realistic warm-up sample.
With 2 CQs the event that later CQ1 uploads wait for is re-recorded after this write, so an upload of the
next frame can never land before it (both write the same buffer from different queues)."""
import ttnn
self._check_open()
slot = self._slot(name, "input")
host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape)
ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
if slot.staging is not None: # the staging buffer mirrors the input (the stage-copy warm-up keeps it)
ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_COMPUTE)
if self.num_command_queues == 2 and self._op_event is not None:
self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
if self.stage_inputs:
self._stage_free_event = self._op_event
# --------------------------------------------------------------------------------------------- run
def _trace_for(self, variant: Optional[str], phase: Optional[int] = None) -> _Trace:
self._check_open()
name = variant or self._last_variant or (next(iter(self._variants)) if len(self._variants) == 1 else None)
if name is None:
raise ValueError("variant name required")
if name not in self._variants:
raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}")
p = self._phase if phase is None else phase
trace = self._traces.get((name, p))
if trace is None:
raise RuntimeError(f"{self.name}: variant {name!r} is not captured; call capture() first")
return trace
def _execute(self, trace: _Trace) -> None:
import ttnn
ttnn.execute_trace(self.device, trace.trace_id, cq_id=CQ_COMPUTE, blocking=False)
if self.num_command_queues == 2:
self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
trace.done_event = self._op_event
self._last_variant = trace.variant
self._last_phase[trace.variant] = trace.phase
if trace.steps:
self._phase ^= 1
def upload(self, inputs: Optional[Mapping[str, Any]] = None, params: Optional[Mapping[str, Any]] = None) -> int:
"""Enqueue the uploads of ``inputs`` and changed ``params`` (CQ0, or CQ1 + events with 2 CQs) without
replaying anything; the next :meth:`run` / :meth:`replay` consumes them. Returns the number of tensors
uploaded."""
self._check_open()
if not self.captured:
raise RuntimeError(f"{self.name}: capture() before uploading")
uploads = self._collect_uploads(inputs, params)
self._enqueue_uploads(uploads)
return len(uploads)
def run(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
params: Optional[Mapping[str, Any]] = None):
"""Upload ``inputs`` (name -> numpy / torch / ttnn host tensor) and changed ``params``, then replay the
variant's trace without blocking. Returns its device outputs (valid once the replay finished)."""
trace = self._trace_for(variant)
self._enqueue_uploads(self._collect_uploads(inputs, params))
self._execute(trace)
return trace.outputs
def replay(self, variant: Optional[str] = None, n: int = 1) -> None:
"""Replay ``n`` times with no uploads (back-to-back device timing); honours ping-pong phases."""
for _ in range(n):
self._execute(self._trace_for(variant))
def outputs(self, variant: Optional[str] = None, phase: Optional[int] = None):
"""Device outputs of a variant's trace (default: the phase that ran last for it, else phase 0)."""
name = variant or self._last_variant
if phase is None and name is not None:
phase = self._last_phase.get(name, 0)
return self._trace_for(name, phase if phase is not None else 0).outputs
def read(self, variant: Optional[str] = None, *, cq_id: int = CQ_COMPUTE, as_torch: bool = False):
"""Read the outputs of the last run of ``variant`` into preallocated host tensors (blocking).
Returns the output structure with numpy arrays (``as_torch=True``: torch tensors); a :class:`Packed` leaf
becomes ``{name: array}``. ``cq_id=1`` (2CQ) waits on the host for the trace's completion event and reads on
CQ1, so CQ0 can already replay the next segment (segmented D2H)."""
import ttnn
name = variant or self._last_variant
if name is None or name not in self._last_phase:
raise RuntimeError(f"{self.name}: variant {name!r} has not run yet")
trace = self._trace_for(name, self._last_phase[name])
if cq_id not in (CQ_COMPUTE, CQ_INPUT):
raise ValueError(f"cq_id={cq_id}: expected 0 or 1")
if cq_id == CQ_INPUT:
if self.num_command_queues != 2:
raise ValueError("cq_id=1 needs a device opened with 2 command queues")
ttnn.event_synchronize(trace.done_event)
if trace.host_buffers is None:
trace.host_buffers = [ttnn.allocate_tensor_on_host(t.spec, self.device) for _, t, _ in trace.leaves]
values: Dict[Tuple, Any] = {}
for (path, tensor, layout), host in zip(trace.leaves, trace.host_buffers):
ttnn.copy_device_to_host_tensor(tensor, host, blocking=True, cq_id=cq_id)
arr = to_numpy(host)
value: Any = layout.unpack(arr) if layout is not None else arr
if as_torch:
import torch
value = ({k: torch.from_numpy(np.ascontiguousarray(v)) for k, v in value.items()}
if isinstance(value, dict) else torch.from_numpy(np.ascontiguousarray(value)))
values[path] = value
structure = trace.outputs
if isinstance(structure, Packed) or not isinstance(structure, (Mapping, list, tuple)):
return values[()]
return _rebuild(structure, values)
def __call__(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
params: Optional[Mapping[str, Any]] = None, *, as_torch: bool = False):
"""``run`` + ``read`` (the common synchronous path)."""
self.run(variant, inputs, params)
return self.read(variant, as_torch=as_torch)
def run_eager(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
params: Optional[Mapping[str, Any]] = None):
"""Upload, run the variant's function *eagerly* (no trace) and return its outputs read to numpy, freeing
every eager buffer before returning (safe next to captured traces). State writes happen as in a replay.
Use it for replay-vs-eager bit checks."""
import ttnn
self._check_open()
name = variant or self._last_variant or (next(iter(self._variants)) if self._variants else None)
if name not in self._variants:
raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}")
variant_obj = self._variants[name]
uploads = self._collect_uploads(inputs, params)
for slot, host, key in uploads:
ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
if key is not None:
slot.last_value = key
self._pending_params.clear()
ctx = TraceContext(self, name, self._phase, capturing=False)
with self._eager(f"run_eager({name!r})"):
outputs = variant_obj.fn(ctx)
values = {}
for path, leaf in _flatten(outputs):
if isinstance(leaf, Packed):
values[path] = leaf.layout.unpack(to_numpy(leaf.tensor))
else:
values[path] = to_numpy(leaf)
self._free_transient(outputs)
pingpong = {s.name for s in self._slots.values() if s.pingpong}
if pingpong and (ctx.writes & pingpong) == pingpong:
self._phase ^= 1
if isinstance(outputs, Packed) or not isinstance(outputs, (Mapping, list, tuple)):
return values[()]
return _rebuild(outputs, values)
# ------------------------------------------------------------------------------------------- state
@property
def phase(self) -> int:
"""Which buffer of each ping-pong state the next run reads (0 = the first buffer)."""
return self._phase
def state_buffer(self, name: str):
"""The device buffer the next run reads for state ``name`` (ping-pong: the buffer of the current
:attr:`phase`)."""
slot = self._slot(name, "state")
return slot.buffers[self._phase] if slot.pingpong else slot.buffers[0]
def reset_state(self, name: Optional[str] = None, value: Any = None) -> None:
"""Write the initial value (or ``value``) of state ``name`` (all states when ``None``) into the buffer the
next run reads. ``value``: numpy / scalar / torch / ttnn host tensor (uploaded) or a ttnn *device* tensor of
the same shape (``ttnn.copy`` on the device). Enqueued on CQ0, so it lands after any replay already
enqueued."""
import ttnn
self._check_open()
names = [name] if name is not None else [s.name for s in self._slots.values() if s.kind == "state"]
for n in names:
slot = self._slot(n, "state")
target = self.state_buffer(n)
if isinstance(value, ttnn.Tensor) and value.storage_type() == ttnn.StorageType.DEVICE:
_check_copy(f"reset_state({n!r})", value, slot)
with self._eager(f"reset_state({n!r}) from a device tensor"):
ttnn.copy(value, target)
continue
if value is None:
host = slot.init
elif isinstance(value, ttnn.Tensor):
host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape)
else:
arr = to_numpy(value) if hasattr(value, "detach") else np.asarray(value)
host = to_host_tensor(np.broadcast_to(arr, slot.shape), slot.dtype, slot.layout)
ttnn.copy_host_to_device_tensor(host, target, cq_id=CQ_COMPUTE)
def read_state(self, name: str) -> np.ndarray:
"""The value the next run will read for state ``name`` (blocking read on CQ0)."""
return to_numpy(self.state_buffer(name))
def _bank(self, name: str, bank: int):
slot = self._slot(name, "state")
if not 0 <= bank < len(slot.banks):
raise IndexError(f"{self.name}: state {name!r} has {len(slot.banks)} bank(s) (add_state(banks=...)), "
f"not bank {bank}")
return slot.banks[bank]
def save_state(self, name: str, bank: int) -> None:
"""Copy the current value of state ``name`` into its ``bank`` (``ttnn.copy`` on CQ0, eager, after any
replay already enqueued): the first half of a stream switch (PLAN.md D16; :class:`StreamBanks`)."""
import ttnn
self._check_open()
with self._eager(f"save_state({name!r})"):
ttnn.copy(self.state_buffer(name), self._bank(name, bank))
def load_state(self, name: str, bank: int) -> None:
"""Copy ``bank`` back into the buffer the next run reads for state ``name`` (``ttnn.copy`` on CQ0)."""
import ttnn
self._check_open()
with self._eager(f"load_state({name!r})"):
ttnn.copy(self._bank(name, bank), self.state_buffer(name))
# ------------------------------------------------------------------------------------------- misc
def trace_ids(self) -> Dict[Tuple[str, int], Any]:
"""``{(variant, phase): trace_id}`` for profiling scripts."""
return {k: t.trace_id for k, t in self._traces.items()}
def describe(self) -> Dict[str, Any]:
"""A JSON-able summary for ``model.info`` / OPT_BASELINE."""
def spec(s: _Slot) -> Dict[str, Any]:
return {"shape": list(s.shape), "dtype": dtype_name(s.dtype), "layout": str(s.layout).rsplit(".", 1)[-1],
**({"pingpong": s.pingpong, "banks": len(s.banks)} if s.kind == "state" else {})}
entries = None
if hasattr(self.device, "num_program_cache_entries"):
entries = int(self.device.num_program_cache_entries())
return {
"name": self.name, "num_command_queues": self.num_command_queues, "stage_inputs": self.stage_inputs,
"warmup_runs": self.warmup_runs, "phases": self.phases, "alloc_tracking": self.alloc_tracking,
"variants": sorted(self._variants), "traces": len(self._traces),
"inputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "input"},
"params": {s.name: spec(s) for s in self._slots.values() if s.kind == "param"},
"states": {s.name: spec(s) for s in self._slots.values() if s.kind == "state"},
"outputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "output"},
"timings_ms": {k: {kk: round(vv, 3) for kk, vv in v.items()} for k, v in self.timings_ms.items()},
"program_cache_entries": entries,
}
def _release_traces(self) -> None:
traces, self._traces = self._traces, {}
for trace in traces.values():
self._safe_release(trace.trace_id)
self._last_phase.clear()
self._last_variant = None
self._captures = 0
self._deallocate_outputs(tensor for trace in traces.values() for _, tensor, _ in trace.leaves)
def release(self) -> None:
"""Release every trace and deallocate the persistent tensors. Idempotent; the persistent tensors are freed
even when the device sync or a trace release raises (the error propagates afterwards)."""
import ttnn
if self._closed:
return
self._closed = True
try:
try:
ttnn.synchronize_device(self.device)
finally:
self._release_traces()
finally:
slots, self._slots = list(self._slots.values()), {}
for slot in slots:
for t in slot.tensors():
if t.is_allocated():
ttnn.deallocate(t)
def _check_open(self) -> None:
if self._closed:
raise RuntimeError(f"{self.name}: the TraceRunner was released")
def __enter__(self) -> "TraceRunner":
return self
def __exit__(self, *exc) -> None:
self.release()
class StreamBanks:
"""Per-stream device state of a temporal model (PLAN.md D16) on top of the states of a :class:`TraceRunner`.
The traces read and write one set of state buffers: the *active* stream's. With ``max_streams > 1`` every known
stream id owns one bank of each state (``add_state(..., banks=max_streams)``) and a switch saves the active
stream into its bank and loads the selected one (``ttnn.copy`` on CQ0, after the replays already enqueued;
probe P13 measured ~80 us eager for a StreamPETR-size state). A new id beyond ``max_streams`` is refused
(``on_full="reject"``: :class:`~.io.InputError`, HTTP 400) or takes over the least recently used stream
(``"evict"``; with ``max_streams=1`` that simply restarts the one state). A stream starts fresh -- its states
reset to their ``init`` values -- when it is new, when ``reset=True``, or when its timestamp goes backwards or
jumps by more than ``max_gap_s``; :meth:`select` returns True then, so the model can also set its first-frame
RT-dev params (BEVFormer ``use_prev_bev=0``, BEVDet ``flag``, ...). Create it after the last ``capture()`` (a
capture resets every state) and call :meth:`select` under the model lock, before the run of each frame::
S_MAX = 1 # first publish (D16)
runner.add_state("prev_bev", shape=(1, 1, 22500, 256), dtype="bfloat16", pingpong=True,
banks=S_MAX if S_MAX > 1 else 0)
streams = StreamBanks(runner, ["prev_bev"], max_streams=S_MAX, on_full="evict", max_gap_s=2.0)
...
fresh = streams.select(stream.get("id", "default"), reset=stream.get("reset", False),
timestamp_s=stream.get("timestamp_s"))
out = runner("frame", inputs=..., params={"use_prev_bev": 0.0 if fresh else 1.0})
"""
def __init__(self, runner: TraceRunner, states: Sequence[str], *, max_streams: int = 1, on_full: str = "reject",
max_gap_s: Optional[float] = None):
if int(max_streams) < 1:
raise ValueError("max_streams must be >= 1")
if on_full not in ("reject", "evict"):
raise ValueError(f"on_full={on_full!r}: expected 'reject' or 'evict'")
self.runner = runner
self.states = tuple(states)
if not self.states:
raise ValueError("StreamBanks needs at least one state")
self.max_streams = int(max_streams)
for name in self.states:
banks = len(runner._slot(name, "state").banks)
if self.max_streams > 1 and banks < self.max_streams:
raise ValueError(f"state {name!r} has {banks} bank(s); StreamBanks(max_streams={self.max_streams}) "
f"needs add_state(..., banks={self.max_streams})")
self.on_full = on_full
self.max_gap_s = None if max_gap_s is None else float(max_gap_s)
self.active: Optional[str] = None
self._bank: Dict[str, int] = {} # stream id -> bank index (max_streams > 1)
self._last_t: Dict[str, float] = {} # stream id -> timestamp of its last frame
self._used: Dict[str, int] = {} # stream id -> last use (LRU clock)
self._clock = 0
@property
def streams(self) -> List[str]:
"""Known stream ids, most recently used first."""
return sorted(self._used, key=lambda s: -self._used[s])
def _save_active(self) -> None:
if self.active is not None and self.active in self._bank:
for name in self.states:
self.runner.save_state(name, self._bank[self.active])
def forget(self, stream_id: str) -> None:
"""Drop a stream (its bank becomes free; if it was active, the next :meth:`select` starts fresh)."""
sid = str(stream_id)
self._used.pop(sid, None)
self._bank.pop(sid, None)
self._last_t.pop(sid, None)
if self.active == sid:
self.active = None
def select(self, stream_id: Any = "default", *, reset: bool = False, timestamp_s: Optional[float] = None) -> bool:
"""Make ``stream_id`` the active stream for the next run; returns True when its state starts fresh."""
sid = str(stream_id)
fresh = bool(reset)
if sid != self.active:
if sid in self._used: # a known stream parked in its bank
self._save_active()
if not fresh:
for name in self.states:
self.runner.load_state(name, self._bank[sid])
else: # a new stream
if len(self._used) >= self.max_streams:
if self.on_full == "reject":
raise InputError(f"stream {sid!r}: this model keeps device state for {self.max_streams} "
f"stream(s), in use by {self.streams}; reuse an id")
self.forget(min(self._used, key=self._used.__getitem__))
self._save_active()
if self.max_streams > 1:
self._bank[sid] = min(set(range(self.max_streams)) - set(self._bank.values()))
fresh = True
self.active = sid
if not fresh and timestamp_s is not None and self.max_gap_s is not None and sid in self._last_t:
dt = float(timestamp_s) - self._last_t[sid]
fresh = dt < 0 or dt > self.max_gap_s
if fresh:
for name in self.states:
self.runner.reset_state(name)
if timestamp_s is not None:
self._last_t[sid] = float(timestamp_s)
elif fresh:
self._last_t.pop(sid, None)
self._clock += 1
self._used[sid] = self._clock
return fresh
def describe(self) -> Dict[str, Any]:
"""JSON-able summary for ``model.info``."""
return {"max_streams": self.max_streams, "on_full": self.on_full, "max_gap_s": self.max_gap_s,
"states": list(self.states), "active": self.active, "streams": self.streams}
|