Add files using upload-large-folder tool
Browse files- .gitattributes +10 -0
- code/rf_detr/server/app.py +32 -12
- code/rf_detr/tt/ttnn_backbone.py +434 -47
- code/rf_detr/tt/ttnn_rf_detr.py +93 -52
- image/blobs/sha256/0aa064431bd920ad87ca8f42142afb0bbc909e68bc589afbbc6ff87800ac1153 +3 -0
- image/blobs/sha256/2cbaee648501aee7415c5bb636fcfd6efdfe01bd976f1e3b9d8fd5be48d3bc00 +3 -0
- image/blobs/sha256/4186f6a9387df261e246adfcc37ae3225f74d40791c2dced1a46fc9d12b6398c +3 -0
- image/blobs/sha256/6134c251abc1a162f6d394d0b13369424783eee9bf3c4c4ced16f162a4fad452 +3 -0
- image/blobs/sha256/77872ef52f01066f9ad1825a9a71a56d304560f756a575654f51973aa59d9ac4 +3 -0
- image/blobs/sha256/85e1ed2c56959e565b7ed00246b276c265f12dfab1060197f8c2969868ef90aa +3 -0
- image/blobs/sha256/a435df3ddd972def28ddfda336794ad5391073fd1a2c86ee50e78fc2d5a4fd46 +3 -0
- image/blobs/sha256/d0c2d85cc7e14bb778666593a7c2b8307242528023ca6f5ec53497244b5f5ea8 +3 -0
- image/blobs/sha256/f5c0a29ae70a96b2a3a149b03201b80a49e1052dd91e0bf657f3c6658e664d33 +3 -0
- image/blobs/sha256/fc2482b5110947c3e5ce45808795a1ba8be69d4c095f3482e1c83fa096562563 +3 -0
.gitattributes
CHANGED
|
@@ -64,3 +64,13 @@ image/blobs/sha256/61f1ff0c58566f161609b7eaec666117292f8464ce99234f2d245c85d0f67
|
|
| 64 |
image/blobs/sha256/9ca1750a03c00f6f105afec1d35651daf877fc2c527fe3a0aa0ec17720cd0a1a filter=lfs diff=lfs merge=lfs -text
|
| 65 |
image/blobs/sha256/c0ae3b923ed167e124b3784b3626662cd6e5a83bb8e605d432b9bf99be2dfda4 filter=lfs diff=lfs merge=lfs -text
|
| 66 |
image/blobs/sha256/8340bb434cbf4a154610cb79702042f9e17b3b240900ff8ff08eeef62f5725f9 filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
image/blobs/sha256/9ca1750a03c00f6f105afec1d35651daf877fc2c527fe3a0aa0ec17720cd0a1a filter=lfs diff=lfs merge=lfs -text
|
| 65 |
image/blobs/sha256/c0ae3b923ed167e124b3784b3626662cd6e5a83bb8e605d432b9bf99be2dfda4 filter=lfs diff=lfs merge=lfs -text
|
| 66 |
image/blobs/sha256/8340bb434cbf4a154610cb79702042f9e17b3b240900ff8ff08eeef62f5725f9 filter=lfs diff=lfs merge=lfs -text
|
| 67 |
+
image/blobs/sha256/a435df3ddd972def28ddfda336794ad5391073fd1a2c86ee50e78fc2d5a4fd46 filter=lfs diff=lfs merge=lfs -text
|
| 68 |
+
image/blobs/sha256/2cbaee648501aee7415c5bb636fcfd6efdfe01bd976f1e3b9d8fd5be48d3bc00 filter=lfs diff=lfs merge=lfs -text
|
| 69 |
+
image/blobs/sha256/6134c251abc1a162f6d394d0b13369424783eee9bf3c4c4ced16f162a4fad452 filter=lfs diff=lfs merge=lfs -text
|
| 70 |
+
image/blobs/sha256/d0c2d85cc7e14bb778666593a7c2b8307242528023ca6f5ec53497244b5f5ea8 filter=lfs diff=lfs merge=lfs -text
|
| 71 |
+
image/blobs/sha256/4186f6a9387df261e246adfcc37ae3225f74d40791c2dced1a46fc9d12b6398c filter=lfs diff=lfs merge=lfs -text
|
| 72 |
+
image/blobs/sha256/f5c0a29ae70a96b2a3a149b03201b80a49e1052dd91e0bf657f3c6658e664d33 filter=lfs diff=lfs merge=lfs -text
|
| 73 |
+
image/blobs/sha256/77872ef52f01066f9ad1825a9a71a56d304560f756a575654f51973aa59d9ac4 filter=lfs diff=lfs merge=lfs -text
|
| 74 |
+
image/blobs/sha256/0aa064431bd920ad87ca8f42142afb0bbc909e68bc589afbbc6ff87800ac1153 filter=lfs diff=lfs merge=lfs -text
|
| 75 |
+
image/blobs/sha256/85e1ed2c56959e565b7ed00246b276c265f12dfab1060197f8c2969868ef90aa filter=lfs diff=lfs merge=lfs -text
|
| 76 |
+
image/blobs/sha256/fc2482b5110947c3e5ce45808795a1ba8be69d4c095f3482e1c83fa096562563 filter=lfs diff=lfs merge=lfs -text
|
code/rf_detr/server/app.py
CHANGED
|
@@ -8,11 +8,15 @@ Served by tt-model-manager as ``kind: tt-dit-server``::
|
|
| 8 |
mesh_shape_env: TT_MESH_SHAPE
|
| 9 |
|
| 10 |
uvicorn runs this module with ``--lifespan on``. Weights are loaded, the chip is
|
| 11 |
-
opened, the TT model is built and warmed
|
| 12 |
-
|
| 13 |
-
-
|
| 14 |
-
|
| 15 |
-
``
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
|
| 17 |
Importing this module has no side effects: no device, no downloads, no env reads.
|
| 18 |
Everything device-touching happens in :func:`lifespan`; ``ttnn`` and the port's
|
|
@@ -75,7 +79,7 @@ MAX_IMAGE_SIDE = 8192
|
|
| 75 |
|
| 76 |
# Device-open parameters the port validated with (conftest.py, scripts/make_demo.py,
|
| 77 |
# rf_detr/benchmark.py): l1_small for the projector's conv2d, a 90 MB trace region for
|
| 78 |
-
# the
|
| 79 |
DEVICE_PARAMS = dict(l1_small_size=32768, trace_region_size=90_000_000, num_command_queues=1)
|
| 80 |
|
| 81 |
log = logging.getLogger("rf_detr.server")
|
|
@@ -211,15 +215,20 @@ async def lifespan(_app: FastAPI):
|
|
| 211 |
model = TtRfDetr(ref, device)
|
| 212 |
log.info("TT model built in %.1fs", time.perf_counter() - t0)
|
| 213 |
|
| 214 |
-
# 4) warm up: the first call runs the
|
| 215 |
-
#
|
| 216 |
-
#
|
|
|
|
|
|
|
|
|
|
| 217 |
log.info("Warming up: capturing trace (first call compiles kernels on a cold cache) ...")
|
| 218 |
t0 = time.perf_counter()
|
| 219 |
dummy = _warmup_input(pre)
|
| 220 |
with torch.inference_mode():
|
| 221 |
model(dummy)
|
| 222 |
t_first = time.perf_counter() - t0
|
|
|
|
|
|
|
| 223 |
t1 = time.perf_counter()
|
| 224 |
out = model(dummy)
|
| 225 |
ttnn.synchronize_device(device)
|
|
@@ -230,10 +239,16 @@ async def lifespan(_app: FastAPI):
|
|
| 230 |
)
|
| 231 |
if not (torch.isfinite(out.logits).all() and torch.isfinite(out.pred_boxes).all()):
|
| 232 |
raise RuntimeError("warm-up produced non-finite logits/boxes")
|
|
|
|
| 233 |
log.info(
|
| 234 |
-
"Warmup complete: first call %.1fs (compile + trace capture), traced call %.
|
| 235 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 236 |
)
|
|
|
|
| 237 |
except BaseException:
|
| 238 |
try:
|
| 239 |
if model is not None:
|
|
@@ -252,6 +267,9 @@ async def lifespan(_app: FastAPI):
|
|
| 252 |
"revision": None if cfg["weights_dir"] else cfg["revision"],
|
| 253 |
"files": {"model.safetensors": weights_path, "config.json": config_path}},
|
| 254 |
started_at=time.time(),
|
|
|
|
|
|
|
|
|
|
| 255 |
)
|
| 256 |
STATE["ready"] = True
|
| 257 |
try:
|
|
@@ -414,7 +432,9 @@ def info() -> dict:
|
|
| 414 |
"trace_region_size": cfg.get("trace_region_size", DEVICE_PARAMS["trace_region_size"]),
|
| 415 |
"num_command_queues": 1,
|
| 416 |
},
|
| 417 |
-
"expected_latency_ms": "~
|
|
|
|
|
|
|
| 418 |
}
|
| 419 |
|
| 420 |
|
|
|
|
| 8 |
mesh_shape_env: TT_MESH_SHAPE
|
| 9 |
|
| 10 |
uvicorn runs this module with ``--lifespan on``. Weights are loaded, the chip is
|
| 11 |
+
opened, the TT model is built and warmed inside the ASGI **lifespan**: the first call
|
| 12 |
+
runs the whole device graph eagerly (kernel JIT, program cache) and captures it as ONE
|
| 13 |
+
metal-trace (patch embedding -> backbone -> feature shaping -> projector -> transformer,
|
| 14 |
+
see ``TtRfDetr._capture_trace``), the second call replays that trace; the lifespan
|
| 15 |
+
checks the trace exists before READY. So uvicorn's ``Application startup complete``
|
| 16 |
+
-- the line ``tt-model serve`` waits for -- means the device is claimed, the trace is
|
| 17 |
+
captured and the model is warm. Every request then is: preprocess (torch) -> host im2col
|
| 18 |
+
into the persistent input buffer -> upload -> ``execute_trace`` -> two readbacks. The
|
| 19 |
+
device is closed in the lifespan shutdown (SIGTERM from ``tt-model stop``).
|
| 20 |
|
| 21 |
Importing this module has no side effects: no device, no downloads, no env reads.
|
| 22 |
Everything device-touching happens in :func:`lifespan`; ``ttnn`` and the port's
|
|
|
|
| 79 |
|
| 80 |
# Device-open parameters the port validated with (conftest.py, scripts/make_demo.py,
|
| 81 |
# rf_detr/benchmark.py): l1_small for the projector's conv2d, a 90 MB trace region for
|
| 82 |
+
# the whole-graph metal-trace, a single command queue (2 CQs regressed).
|
| 83 |
DEVICE_PARAMS = dict(l1_small_size=32768, trace_region_size=90_000_000, num_command_queues=1)
|
| 84 |
|
| 85 |
log = logging.getLogger("rf_detr.server")
|
|
|
|
| 215 |
model = TtRfDetr(ref, device)
|
| 216 |
log.info("TT model built in %.1fs", time.perf_counter() - t0)
|
| 217 |
|
| 218 |
+
# 4) warm up: the first call runs the whole device graph eagerly once (JIT compile
|
| 219 |
+
# on a cold cache, program cache, conv2d weight prep) and then captures it as ONE
|
| 220 |
+
# metal-trace with a persistent input buffer (TtRfDetr._capture_trace); the
|
| 221 |
+
# second call exercises the trace-replay path every request will take. The
|
| 222 |
+
# trace must exist before READY -- a model that silently fell back to eager
|
| 223 |
+
# would serve at 3x the latency, so that is a boot failure, not a warning.
|
| 224 |
log.info("Warming up: capturing trace (first call compiles kernels on a cold cache) ...")
|
| 225 |
t0 = time.perf_counter()
|
| 226 |
dummy = _warmup_input(pre)
|
| 227 |
with torch.inference_mode():
|
| 228 |
model(dummy)
|
| 229 |
t_first = time.perf_counter() - t0
|
| 230 |
+
if getattr(model, "use_trace", False) and getattr(model, "_trace_id", None) is None:
|
| 231 |
+
raise RuntimeError("warm-up did not capture the metal-trace (TtRfDetr._trace_id is None)")
|
| 232 |
t1 = time.perf_counter()
|
| 233 |
out = model(dummy)
|
| 234 |
ttnn.synchronize_device(device)
|
|
|
|
| 239 |
)
|
| 240 |
if not (torch.isfinite(out.logits).all() and torch.isfinite(out.pred_boxes).all()):
|
| 241 |
raise RuntimeError("warm-up produced non-finite logits/boxes")
|
| 242 |
+
pin = getattr(model, "_persistent_in", None)
|
| 243 |
log.info(
|
| 244 |
+
"Warmup complete: first call %.1fs (compile + trace capture), traced call %.1f ms; "
|
| 245 |
+
"trace_id=%s input_path=%s persistent_in=%s; boot %.1fs",
|
| 246 |
+
t_first, t_second * 1000.0, getattr(model, "_trace_id", None),
|
| 247 |
+
getattr(getattr(model, "backbone", None), "input_path", "?"),
|
| 248 |
+
None if pin is None else f"{tuple(pin.shape)} {pin.dtype} {pin.layout}",
|
| 249 |
+
time.perf_counter() - boot_t0,
|
| 250 |
)
|
| 251 |
+
warmup_ms = t_second * 1000.0
|
| 252 |
except BaseException:
|
| 253 |
try:
|
| 254 |
if model is not None:
|
|
|
|
| 267 |
"revision": None if cfg["weights_dir"] else cfg["revision"],
|
| 268 |
"files": {"model.safetensors": weights_path, "config.json": config_path}},
|
| 269 |
started_at=time.time(),
|
| 270 |
+
warmup_traced_call_ms=round(warmup_ms, 2),
|
| 271 |
+
trace={"captured": getattr(model, "_trace_id", None) is not None,
|
| 272 |
+
"input_path": getattr(model.backbone, "input_path", None)},
|
| 273 |
)
|
| 274 |
STATE["ready"] = True
|
| 275 |
try:
|
|
|
|
| 432 |
"trace_region_size": cfg.get("trace_region_size", DEVICE_PARAMS["trace_region_size"]),
|
| 433 |
"num_command_queues": 1,
|
| 434 |
},
|
| 435 |
+
"expected_latency_ms": "~12 ms per image steady state on p150a (~85 FPS, whole graph in one metal-trace)",
|
| 436 |
+
"warmup_traced_call_ms": STATE.get("warmup_traced_call_ms"),
|
| 437 |
+
"trace": STATE.get("trace"),
|
| 438 |
}
|
| 439 |
|
| 440 |
|
code/rf_detr/tt/ttnn_backbone.py
CHANGED
|
@@ -1,16 +1,78 @@
|
|
| 1 |
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
"""TTNN port of RF-DETR's windowed DINOv2-S/14 backbone.
|
| 3 |
|
| 4 |
-
Embeddings (patch conv + cls + interpolated pos-embed + window partition)
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
|
| 9 |
Windowing: 560/14 = 40 patch grid, num_windows=4 -> 16 windows of 10x10=100 patches.
|
| 10 |
Windowed layers operate on [16, 101, 384] (1 cls + 100 patches per window).
|
| 11 |
Global layers (2,5,8,11) attend over the merged [1, 1616, 384] sequence.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
"""
|
| 13 |
|
|
|
|
|
|
|
| 14 |
import torch
|
| 15 |
import ttnn
|
| 16 |
|
|
@@ -19,6 +81,11 @@ import ttnn
|
|
| 19 |
# below the 99% gate — so the deepest stage keeps bf16.)
|
| 20 |
WEIGHT_DTYPE = ttnn.bfloat16
|
| 21 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 22 |
|
| 23 |
def _lin(linear, device, dtype=WEIGHT_DTYPE):
|
| 24 |
"""torch nn.Linear -> (ttnn weight [in,out], ttnn bias [1,out] or None)."""
|
|
@@ -37,20 +104,81 @@ def _vec(t, device):
|
|
| 37 |
return ttnn.from_torch(t.detach().reshape(1, 1, -1), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device)
|
| 38 |
|
| 39 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
class TtDinoBackbone:
|
| 41 |
-
def __init__(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
self.device = device
|
| 43 |
-
#
|
| 44 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
self.mem = ttnn.L1_MEMORY_CONFIG if l1 else None
|
| 46 |
wb = ref_model.backbone[0].encoder.encoder # WindowedDinoBackbone
|
| 47 |
-
self.wb = wb # kept for host-side embeddings +
|
| 48 |
self.cfg = wb.cfg
|
|
|
|
| 49 |
self.num_heads = self.cfg.num_attention_heads
|
| 50 |
self.head_dim = self.cfg.hidden_size // self.num_heads
|
| 51 |
self.eps = self.cfg.layer_norm_eps
|
| 52 |
self.num_windows = self.cfg.num_windows
|
| 53 |
self.nw2 = self.num_windows ** 2
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 54 |
# Precision knob: compute_kernel_config controls matmul math fidelity + fp32 accumulation.
|
| 55 |
# None => ttnn default. Passed to every backbone matmul.
|
| 56 |
self.compute_config = (
|
|
@@ -61,8 +189,51 @@ class TtDinoBackbone:
|
|
| 61 |
else None
|
| 62 |
)
|
| 63 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 64 |
self.layers = []
|
| 65 |
-
for layer in wb.encoder.layer:
|
| 66 |
att = layer.attention.attention
|
| 67 |
qkv_w = torch.cat([att.query.weight, att.key.weight, att.value.weight], dim=0) # [3*384,384]
|
| 68 |
qkv_b = torch.cat([att.query.bias, att.key.bias, att.value.bias], dim=0)
|
|
@@ -75,6 +246,8 @@ class TtDinoBackbone:
|
|
| 75 |
proj_w, proj_b = _lin(layer.attention.output.dense, device, dtype=weight_dtype)
|
| 76 |
fc1_w, fc1_b = _lin(layer.mlp.fc1, device, dtype=weight_dtype)
|
| 77 |
fc2_w, fc2_b = _lin(layer.mlp.fc2, device, dtype=weight_dtype)
|
|
|
|
|
|
|
| 78 |
self.layers.append(
|
| 79 |
{
|
| 80 |
"global": layer.global_attention,
|
|
@@ -84,73 +257,287 @@ class TtDinoBackbone:
|
|
| 84 |
"qkv_b": qkv_b_tt,
|
| 85 |
"proj_w": proj_w,
|
| 86 |
"proj_b": proj_b,
|
| 87 |
-
"ls1":
|
|
|
|
| 88 |
"norm2_w": _vec(layer.norm2.weight, device),
|
| 89 |
"norm2_b": _vec(layer.norm2.bias, device),
|
| 90 |
"fc1_w": fc1_w,
|
| 91 |
"fc1_b": fc1_b,
|
| 92 |
"fc2_w": fc2_w,
|
| 93 |
"fc2_b": fc2_b,
|
| 94 |
-
"ls2":
|
|
|
|
| 95 |
}
|
| 96 |
)
|
| 97 |
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 110 |
)
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
ctx = ttnn.transformer.concatenate_heads(ctx) # [b,s,384]
|
| 116 |
ttnn.deallocate(qkv)
|
| 117 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 118 |
if p["global"]:
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 122 |
|
| 123 |
# ---- mlp (norm2 -> fc1 -> gelu -> fc2 -> layerscale2 -> residual) ----
|
| 124 |
-
residual = x
|
| 125 |
h = ttnn.layer_norm(x, weight=p["norm2_w"], bias=p["norm2_b"], epsilon=self.eps, memory_config=mc)
|
| 126 |
-
h =
|
| 127 |
-
h = ttnn.gelu(h, fast_and_approximate_mode=False, memory_config=mc)
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
return x
|
| 132 |
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 137 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 138 |
out = {}
|
| 139 |
for i, p in enumerate(self.layers):
|
| 140 |
x = self._layer(x, p)
|
| 141 |
-
if i in
|
| 142 |
out[i] = ttnn.to_torch(x).float()
|
| 143 |
return out
|
| 144 |
|
| 145 |
-
def
|
| 146 |
-
"""
|
| 147 |
-
Returns list of 4 torch tensors [1,384,40,40]."""
|
| 148 |
embed = self.wb.embeddings(pixel_values) # [16,101,384] host
|
| 149 |
hidden = self.run_layers(embed)
|
| 150 |
_, _, H, W = pixel_values.shape
|
| 151 |
feats = []
|
| 152 |
-
for i in
|
| 153 |
-
hs =
|
| 154 |
hs = self.wb.layernorm(hs)
|
| 155 |
hs = hs[:, 1:] # drop cls per window
|
| 156 |
hs = self.wb.window_unpartition(hs, H, W)
|
|
|
|
| 1 |
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
"""TTNN port of RF-DETR's windowed DINOv2-S/14 backbone.
|
| 3 |
|
| 4 |
+
Embeddings (patch conv + cls + interpolated pos-embed + window partition) run on host;
|
| 5 |
+
the transformer layers AND the feature-map shaping of the four out-index layers
|
| 6 |
+
(1/4/7/10) run on device, so the whole backbone -> projector -> transformer graph is
|
| 7 |
+
one device graph with no mid-graph host readback (see ``TtRfDetr``).
|
| 8 |
+
|
| 9 |
+
Only layers 0..10 are built and run: RF-DETR consumes the hidden states after layers
|
| 10 |
+
1/4/7/10 (reference ``out_indices`` 2/5/8/11 count the embedding output as index 0) and
|
| 11 |
+
nothing reads the output of layer 11, so the reference computes it for nothing; skipping
|
| 12 |
+
it is exact (a global layer, ~0.95 ms of the traced graph).
|
| 13 |
+
|
| 14 |
+
Per-layer op fusion (S2): attention is one ``scaled_dot_product_attention`` call (scale
|
| 15 |
+
fused, padded keys masked by the kernel), and each ``linear -> * layerscale -> + residual``
|
| 16 |
+
pair (proj, fc2) is one ``dit_minimal_matmul_addcmul_fused`` call, so a block is
|
| 17 |
+
LN, qkv matmul, split heads, SDPA, concat heads, proj+ls1+res, LN, fc1, gelu, fc2+ls2+res
|
| 18 |
+
(10 ops instead of 19). The qkv and fc1 matmuls use the same ``minimal_matmul`` kernel
|
| 19 |
+
(2-3x faster than ``ttnn.linear``'s default program on these shapes); the exact gelu stays a
|
| 20 |
+
separate op because fusing it into the matmul is slower. Global layers stay in the merged
|
| 21 |
+
layout for the whole block.
|
| 22 |
|
| 23 |
Windowing: 560/14 = 40 patch grid, num_windows=4 -> 16 windows of 10x10=100 patches.
|
| 24 |
Windowed layers operate on [16, 101, 384] (1 cls + 100 patches per window).
|
| 25 |
Global layers (2,5,8,11) attend over the merged [1, 1616, 384] sequence.
|
| 26 |
+
|
| 27 |
+
Input path (S4, knob ``RFDETR_BB_INPUT``): the per-image host work is only what the device
|
| 28 |
+
cannot do, everything else is traced device ops (``_ingest``):
|
| 29 |
+
|
| 30 |
+
* ``patch`` (default): the patch embedding runs ON DEVICE. Host: im2col of the 560x560 image
|
| 31 |
+
into a persistent fp32 ``[16, 101, 608]`` buffer (one strided copy, 0.03 ms; row 0 of every
|
| 32 |
+
window is a zero row for the cls token, cols 588..607 zero pad) -> ``from_torch`` ROW_MAJOR
|
| 33 |
+
(0.14 ms) -> ``copy_host_to_device`` (3.9 MB, 0.18 ms). Device, inside the trace:
|
| 34 |
+
``tilize_with_zero_padding`` (0.10 ms) then ONE ``dit_minimal_matmul_addcmul_fused`` call
|
| 35 |
+
``pos_full + 1.0 * (X @ W_patch) * ones`` (0.08 ms) = the conv as a matmul with the flattened
|
| 36 |
+
``[588, 384]`` kernel (bf16, HiFi4, fp32 accumulation) plus a constant ``[16, 101, 384]``
|
| 37 |
+
tensor holding ``cls + pos_cls`` on the cls rows and ``pos_patch + conv_bias`` on the patch
|
| 38 |
+
rows (built once from the reference modules, so it cannot drift). This is exact algebra
|
| 39 |
+
(verified fp32 on host: max|d| 0 vs ``wb.embeddings``) but NOT bit-exact on device: the
|
| 40 |
+
kernel is bf16 and the accumulation order differs, embed mean|d| 3e-4 vs the fp32 reference
|
| 41 |
+
(the S1-S3 bf16 embed had 1e-4); the four feature-map PCCs are unchanged to 1e-5 and the
|
| 42 |
+
outputs are deterministic run to run. Uploading the pixels as bf16 instead is NOT enough
|
| 43 |
+
(embed error 1e-3, detection-IoU 98.21 < gate); fp32 W as well needs an fp32 pos tensor +
|
| 44 |
+
a typecast and buys nothing downstream.
|
| 45 |
+
* ``rowmajor``: the S1-S3 host embeddings (torch conv + cls/pos + window partition, 0.9-1.5 ms)
|
| 46 |
+
are cast to bf16 and uploaded ROW_MAJOR ``[16, 101, 384]`` (0.11 ms of host work instead of the
|
| 47 |
+
0.9-1.5 ms host tilize); ``tilize_with_zero_padding`` on device (0.07 ms, in L1) yields the
|
| 48 |
+
padded ``[16, 128, 384]`` the layers run on. Bit-identical to ``tile`` (logits bit-equal).
|
| 49 |
+
* ``tile``: the S1-S3 path -- host tilize via ``from_torch(TILE)`` in the ``RFDETR_BB_UPLOAD``
|
| 50 |
+
layout (merged ``[1, 1616, 384]`` + one device reshape, or windowed).
|
| 51 |
+
|
| 52 |
+
Tried and rejected (S3, lever C): running ALL layers in the merged [1, 1616, 384] layout with
|
| 53 |
+
windowed attention expressed as global attention plus SDPA's on-device block-diagonal mask
|
| 54 |
+
(``cu_window_seqlens=[0, 101, ..., 1616]``). Numerically equivalent (PCC 1.0 vs the batched
|
| 55 |
+
SDPA) but the dense-mask kernel visits every K chunk for every Q chunk (16x the FLOPs:
|
| 56 |
+
0.17 vs 0.05 ms per layer), so the traced backbone got slower (4.51 vs 4.20 ms) while the
|
| 57 |
+
other ops gained only launch-floor crumbs from the 21% fewer rows; and it was
|
| 58 |
+
NON-deterministic run to run (feature maps differ by up to 1.1, detection-IoU 98.15-98.69),
|
| 59 |
+
the on-device partial-tile mask generation being the only op that changed. See
|
| 60 |
+
logs/opt-rf-detr/RESULTS.md (S3) and s3_merged_layout.patch.
|
| 61 |
+
|
| 62 |
+
Device-side feature shaping (``shape_features``): the reference does
|
| 63 |
+
``layernorm(hs)[:, 1:]`` -> ``window_unpartition`` -> raster ``[1, 384, 40, 40]``.
|
| 64 |
+
Dropping the 16 per-window cls rows and re-ordering the windows into raster order is a
|
| 65 |
+
pure row permutation, so it is expressed as ONE matmul with a constant 0/1 matrix
|
| 66 |
+
``P[1600, 1616]`` (``P @ X`` selects rows; exact in bf16 at HiFi4 -- every output row is
|
| 67 |
+
1.0 * one input row + 0 * the rest), followed by the final ``ttnn.layer_norm``. ``P`` is
|
| 68 |
+
built once on host by pushing an index tensor through the reference
|
| 69 |
+
``window_unpartition`` (``build_shaping_perm``), so it cannot drift from the reference.
|
| 70 |
+
The result is channels-last ``[1, 1600, 384]`` in raster order -- exactly what the
|
| 71 |
+
projector consumes.
|
| 72 |
"""
|
| 73 |
|
| 74 |
+
import os
|
| 75 |
+
|
| 76 |
import torch
|
| 77 |
import ttnn
|
| 78 |
|
|
|
|
| 81 |
# below the 99% gate — so the deepest stage keeps bf16.)
|
| 82 |
WEIGHT_DTYPE = ttnn.bfloat16
|
| 83 |
|
| 84 |
+
# Layers whose hidden state becomes a feature map (reference out_indices 2/5/8/11 count the
|
| 85 |
+
# embedding output as hidden_states[0], i.e. the outputs of layers 1/4/7/10). Layers after the
|
| 86 |
+
# last one feed nothing, so the device graph stops there.
|
| 87 |
+
OUT_LAYERS = (1, 4, 7, 10)
|
| 88 |
+
|
| 89 |
|
| 90 |
def _lin(linear, device, dtype=WEIGHT_DTYPE):
|
| 91 |
"""torch nn.Linear -> (ttnn weight [in,out], ttnn bias [1,out] or None)."""
|
|
|
|
| 104 |
return ttnn.from_torch(t.detach().reshape(1, 1, -1), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device)
|
| 105 |
|
| 106 |
|
| 107 |
+
def build_shaping_perm(wb, height, width):
|
| 108 |
+
"""Row-permutation matrix for the feature-map shaping of one out-index layer.
|
| 109 |
+
|
| 110 |
+
Returns ``(P, rows)``: ``P`` is float32 ``[n_patches, nw2 * seq]`` (here [1600, 1616]) with
|
| 111 |
+
exactly one 1.0 per row, and ``rows`` is the LongTensor of source rows such that, for the
|
| 112 |
+
merged windowed hidden state ``X = hs.reshape(nw2 * seq, C)``,
|
| 113 |
+
``P @ X == X[rows] == shaping(hs).flatten(2).transpose(1, 2)[0]`` where ``shaping`` is the
|
| 114 |
+
reference ``hs[:, 1:]`` -> ``window_unpartition`` -> ``reshape(1, H/p, W/p, C)`` ->
|
| 115 |
+
``permute(0, 3, 1, 2)`` chain (without the layernorm, which is row-wise and commutes).
|
| 116 |
+
Built by pushing the flat row index through the reference ``window_unpartition``.
|
| 117 |
+
"""
|
| 118 |
+
cfg = wb.cfg
|
| 119 |
+
nw2 = cfg.num_windows ** 2
|
| 120 |
+
n_h, n_w = height // cfg.patch_size, width // cfg.patch_size
|
| 121 |
+
per_window = (n_h // cfg.num_windows) * (n_w // cfg.num_windows)
|
| 122 |
+
seq = per_window + 1 # + cls
|
| 123 |
+
rows = torch.arange(nw2 * seq, dtype=torch.float32).reshape(nw2, seq, 1)[:, 1:] # drop cls
|
| 124 |
+
rows = wb.window_unpartition(rows, height, width).reshape(-1).long() # raster order
|
| 125 |
+
assert rows.numel() == n_h * n_w == nw2 * per_window
|
| 126 |
+
perm = torch.zeros(rows.numel(), nw2 * seq, dtype=torch.float32)
|
| 127 |
+
perm[torch.arange(rows.numel()), rows] = 1.0
|
| 128 |
+
return perm, rows
|
| 129 |
+
|
| 130 |
+
|
| 131 |
class TtDinoBackbone:
|
| 132 |
+
def __init__(
|
| 133 |
+
self,
|
| 134 |
+
ref_model,
|
| 135 |
+
device,
|
| 136 |
+
weight_dtype=WEIGHT_DTYPE,
|
| 137 |
+
math_fidelity=None,
|
| 138 |
+
fp32_acc=False,
|
| 139 |
+
l1=True,
|
| 140 |
+
image_size=None,
|
| 141 |
+
attn="sdpa",
|
| 142 |
+
matmul="minimal",
|
| 143 |
+
upload=None,
|
| 144 |
+
input_path=None,
|
| 145 |
+
):
|
| 146 |
self.device = device
|
| 147 |
+
# Input path (S4): what the host uploads per image and which traced device ops turn it into the
|
| 148 |
+
# windowed [16, 101, 384] TILE tensor the layers consume (module docstring). None => env
|
| 149 |
+
# RFDETR_BB_INPUT (default "patch": patch embedding on device, 11.8 ms; "rowmajor" is the bit-exact
|
| 150 |
+
# fallback, 12.6 ms; "tile" the S1-S3 host tilize, 13.9-14.2 ms).
|
| 151 |
+
self.input_path = input_path or os.environ.get("RFDETR_BB_INPUT", "patch")
|
| 152 |
+
if self.input_path not in ("patch", "rowmajor", "tile"):
|
| 153 |
+
raise ValueError(f"RFDETR_BB_INPUT must be patch|rowmajor|tile, got {self.input_path!r}")
|
| 154 |
+
# Upload layout of the "tile" input path (S3). "merged": the host embed is uploaded as [1, 1616, 384] (tile-padded to
|
| 155 |
+
# 1632 rows: 1.25 MB, host tilize 0.87 ms) and ONE on-device reshape (0.08 ms) turns it into the
|
| 156 |
+
# windowed [16, 101, 384] (padded [16, 128, 384]: 1.5 MB, host tilize 1.12 ms) the layers run on;
|
| 157 |
+
# the reshape also lands the input in L1. Exact (pure data movement). "windowed": upload as
|
| 158 |
+
# [16, 101, 384] directly (the S1/S2 path). None => env RFDETR_BB_UPLOAD (default merged).
|
| 159 |
+
self.upload = upload or os.environ.get("RFDETR_BB_UPLOAD", "merged")
|
| 160 |
+
# L1 placement: keep the layer's working set (op outputs) in on-chip L1 instead of DRAM
|
| 161 |
+
# round-trips. Default ON: on the traced graph it is 31.0 -> 26.0 ms (S1 measurement; it
|
| 162 |
+
# was a wash on the old eager pipeline, whose host syncs hid it). `l1=False` => ttnn
|
| 163 |
+
# default (DRAM interleaved). Not bit-identical (different matmul blocking/rounding),
|
| 164 |
+
# but within the gates (accuracy 98.617 vs 98.513 with DRAM).
|
| 165 |
self.mem = ttnn.L1_MEMORY_CONFIG if l1 else None
|
| 166 |
wb = ref_model.backbone[0].encoder.encoder # WindowedDinoBackbone
|
| 167 |
+
self.wb = wb # kept for host-side embeddings (+ the host shaping debug path)
|
| 168 |
self.cfg = wb.cfg
|
| 169 |
+
self.hidden = self.cfg.hidden_size
|
| 170 |
self.num_heads = self.cfg.num_attention_heads
|
| 171 |
self.head_dim = self.cfg.hidden_size // self.num_heads
|
| 172 |
self.eps = self.cfg.layer_norm_eps
|
| 173 |
self.num_windows = self.cfg.num_windows
|
| 174 |
self.nw2 = self.num_windows ** 2
|
| 175 |
+
self.out_layers = OUT_LAYERS
|
| 176 |
+
# The device graph is shape-locked to the model's input resolution (560 -> 40x40 patches).
|
| 177 |
+
self.image_size = int(image_size or getattr(ref_model.cfg, "image_resolution", 560))
|
| 178 |
+
self.grid = self.image_size // self.cfg.patch_size # 40
|
| 179 |
+
self.seq_per_window = (self.grid // self.num_windows) ** 2 + 1 # 101
|
| 180 |
+
self.merged_len = self.nw2 * self.seq_per_window # 1616
|
| 181 |
+
self.n_patches = self.grid * self.grid # 1600
|
| 182 |
# Precision knob: compute_kernel_config controls matmul math fidelity + fp32 accumulation.
|
| 183 |
# None => ttnn default. Passed to every backbone matmul.
|
| 184 |
self.compute_config = (
|
|
|
|
| 189 |
else None
|
| 190 |
)
|
| 191 |
|
| 192 |
+
# Attention kernel. "sdpa": ttnn.transformer.scaled_dot_product_attention (FlashAttention-2
|
| 193 |
+
# kernel; the 1/sqrt(d) scale is fused via scale=, and the kernel masks the tile-padded key
|
| 194 |
+
# columns 101->128 / 1616->1632 itself -- do NOT pass an attn_mask: the provided-mask path
|
| 195 |
+
# mishandles the padded columns, which was the README's old "SDPA is wrong here" finding).
|
| 196 |
+
# "matmul": the explicit q@kT -> scale -> softmax -> probs@v chain (kept for A/B debugging).
|
| 197 |
+
# Chunk sizes: windowed 101 -> one 128 chunk per window (single-pass softmax, 96 work items);
|
| 198 |
+
# global 1616 -> q 96 (17 chunks x 6 heads = 102 items on the 110-core grid), k 256.
|
| 199 |
+
# exp_approx_mode=False (no measurable cost). Default 32/32 chunks were slower (0.30 vs 0.13
|
| 200 |
+
# ms per global layer) AND dropped the detection-IoU gate (98.36): more softmax rescale steps.
|
| 201 |
+
self.attn = attn
|
| 202 |
+
grid = device.compute_with_storage_grid_size()
|
| 203 |
+
self.sdpa_pc_window = ttnn.SDPAProgramConfig(
|
| 204 |
+
compute_with_storage_grid_size=grid, q_chunk_size=128, k_chunk_size=128, exp_approx_mode=False
|
| 205 |
+
)
|
| 206 |
+
self.sdpa_pc_global = ttnn.SDPAProgramConfig(
|
| 207 |
+
compute_with_storage_grid_size=grid, q_chunk_size=96, k_chunk_size=256, exp_approx_mode=False
|
| 208 |
+
)
|
| 209 |
+
|
| 210 |
+
# Matmul kernel. "minimal": ttnn.experimental.minimal_matmul for qkv / fc1 and its fused
|
| 211 |
+
# sibling dit_minimal_matmul_addcmul_fused for proj / fc2 (out = residual + (h @ W + b) * ls,
|
| 212 |
+
# i.e. layerscale-multiply + residual-add folded into the matmul). The minimal_matmul kernel
|
| 213 |
+
# is 2-5x faster than ttnn.linear's default program on these shapes (batched fc2
|
| 214 |
+
# [16,128,1536] x [1536,384]: 0.51 -> 0.10 ms; qkv 0.19 -> 0.06). "linear": ttnn.linear +
|
| 215 |
+
# multiply + add (kept for A/B debugging). Compute config HiFi2 WITHOUT fp32 accumulation
|
| 216 |
+
# unless the fidelity knobs say otherwise: the fused op's own default (HiFi2+fp32acc) and
|
| 217 |
+
# HiFi4+fp32acc give slightly higher per-stage PCC but drop the detection-IoU gate
|
| 218 |
+
# (98.39 / 98.41 < 98.5). Explicit small tile blocks: the default 8x8x8 blocks need ~1.1 MB
|
| 219 |
+
# of L1 circular buffers per core and clash with the L1-resident working set.
|
| 220 |
+
# The fused kernel compiles ONE TensorAccessor type for both addcmul inputs, so the
|
| 221 |
+
# layerscale vector must live in the same buffer type (L1 / DRAM) as the residual: ``_ls``.
|
| 222 |
+
self.matmul = matmul
|
| 223 |
+
self.mm_compute_config = self.compute_config or ttnn.init_device_compute_kernel_config(
|
| 224 |
+
device.arch(), math_fidelity=ttnn.MathFidelity.HiFi2, fp32_dest_acc_en=False, packer_l1_acc=True
|
| 225 |
+
)
|
| 226 |
+
self.mm_config = ttnn.MinimalMatmulConfig( # qkv [.,384]x[384,1152], fc1 [.,384]x[384,1536]
|
| 227 |
+
M_block_size=8, K_block_size=4, N_block_size=4, subblock_h=2, subblock_w=2,
|
| 228 |
+
compute_with_storage_grid_size=grid,
|
| 229 |
+
)
|
| 230 |
+
self.dit_config = ttnn.MinimalMatmulConfig( # proj [.,384]x[384,384], fc2 [.,1536]x[1536,384]
|
| 231 |
+
M_block_size=4, K_block_size=4, N_block_size=4, subblock_h=2, subblock_w=2,
|
| 232 |
+
compute_with_storage_grid_size=grid,
|
| 233 |
+
)
|
| 234 |
+
|
| 235 |
self.layers = []
|
| 236 |
+
for layer in wb.encoder.layer[: max(self.out_layers) + 1]: # layer 11's output is never consumed
|
| 237 |
att = layer.attention.attention
|
| 238 |
qkv_w = torch.cat([att.query.weight, att.key.weight, att.value.weight], dim=0) # [3*384,384]
|
| 239 |
qkv_b = torch.cat([att.query.bias, att.key.bias, att.value.bias], dim=0)
|
|
|
|
| 246 |
proj_w, proj_b = _lin(layer.attention.output.dense, device, dtype=weight_dtype)
|
| 247 |
fc1_w, fc1_b = _lin(layer.mlp.fc1, device, dtype=weight_dtype)
|
| 248 |
fc2_w, fc2_b = _lin(layer.mlp.fc2, device, dtype=weight_dtype)
|
| 249 |
+
ls1 = _vec(layer.layer_scale1.lambda1, device)
|
| 250 |
+
ls2 = _vec(layer.layer_scale2.lambda1, device)
|
| 251 |
self.layers.append(
|
| 252 |
{
|
| 253 |
"global": layer.global_attention,
|
|
|
|
| 257 |
"qkv_b": qkv_b_tt,
|
| 258 |
"proj_w": proj_w,
|
| 259 |
"proj_b": proj_b,
|
| 260 |
+
"ls1": ls1,
|
| 261 |
+
"ls1_l1": ttnn.to_memory_config(ls1, ttnn.L1_MEMORY_CONFIG) if l1 else ls1,
|
| 262 |
"norm2_w": _vec(layer.norm2.weight, device),
|
| 263 |
"norm2_b": _vec(layer.norm2.bias, device),
|
| 264 |
"fc1_w": fc1_w,
|
| 265 |
"fc1_b": fc1_b,
|
| 266 |
"fc2_w": fc2_w,
|
| 267 |
"fc2_b": fc2_b,
|
| 268 |
+
"ls2": ls2,
|
| 269 |
+
"ls2_l1": ttnn.to_memory_config(ls2, ttnn.L1_MEMORY_CONFIG) if l1 else ls2,
|
| 270 |
}
|
| 271 |
)
|
| 272 |
|
| 273 |
+
# ---- device-side feature shaping: final layernorm + row-permutation matmul ----
|
| 274 |
+
self.final_norm_w = _vec(wb.layernorm.weight, device)
|
| 275 |
+
self.final_norm_b = _vec(wb.layernorm.bias, device)
|
| 276 |
+
self.final_eps = float(wb.layernorm.eps)
|
| 277 |
+
perm, self.perm_rows = build_shaping_perm(wb, self.image_size, self.image_size)
|
| 278 |
+
# 0/1 matrix, exact in bf16. HiFi4 keeps the full bf16 mantissa of X, so P @ X is a
|
| 279 |
+
# bit-exact row gather (one nonzero term per output; verified bit-exact on device, and
|
| 280 |
+
# fp32 accumulation adds nothing but +50% op time, so it stays off). LoFi is NOT exact.
|
| 281 |
+
self.perm = ttnn.from_torch(perm, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device)
|
| 282 |
+
self.perm_compute_config = ttnn.init_device_compute_kernel_config(
|
| 283 |
+
device.arch(), math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=False, packer_l1_acc=False
|
| 284 |
+
)
|
| 285 |
+
|
| 286 |
+
# ---- on-device patch embedding (input path "patch") ----
|
| 287 |
+
# X[16, 101, 608] (fp32 im2col, zero cls row, zero pad cols) @ W_patch[608, 384] (bf16, rows 588.. zero)
|
| 288 |
+
# + pos_full[16, 101, 384] (bf16): cls rows = cls_token + pos[cls], patch rows = pos[patch] + conv bias,
|
| 289 |
+
# window-partitioned by the reference module. The fused op needs its two addcmul inputs (pos_full,
|
| 290 |
+
# ones) in the same buffer type; both live in DRAM. fp32 pixels + HiFi4 + fp32 accumulation: bf16
|
| 291 |
+
# pixels lose the detection-IoU gate (98.21), see the module docstring.
|
| 292 |
+
self.patch_dim = self.cfg.num_channels * self.cfg.patch_size ** 2 # 588
|
| 293 |
+
self.patch_dim_padded = -(-self.patch_dim // 32) * 32 # 608
|
| 294 |
+
self.patch_w = self.patch_pos = self.patch_ones = None
|
| 295 |
+
self._im2col_buf = None
|
| 296 |
+
if self.input_path == "patch":
|
| 297 |
+
self._build_patch_embed()
|
| 298 |
+
|
| 299 |
+
# ------------------------------------------------------------------- input
|
| 300 |
+
@property
|
| 301 |
+
def input_shape(self):
|
| 302 |
+
""""tile" path: shape the host embed [16, 101, 384] is reshaped to (a free torch view) before
|
| 303 |
+
``from_torch``; the other paths upload the windowed shape."""
|
| 304 |
+
if self.input_path == "tile" and self.upload == "merged":
|
| 305 |
+
return (1, self.merged_len, self.hidden)
|
| 306 |
+
return (self.nw2, self.seq_per_window, self.hidden)
|
| 307 |
+
|
| 308 |
+
def _build_patch_embed(self):
|
| 309 |
+
emb = self.wb.embeddings
|
| 310 |
+
conv = emb.patch_embeddings.projection
|
| 311 |
+
w = torch.zeros(self.patch_dim_padded, self.hidden)
|
| 312 |
+
w[: self.patch_dim] = conv.weight.detach().reshape(self.hidden, self.patch_dim).t() # (c, kh, kw) order
|
| 313 |
+
self.patch_w = ttnn.from_torch(w, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=self.device)
|
| 314 |
+
pos = emb.interpolate_pos_encoding(torch.zeros(1, self.n_patches + 1, self.hidden), self.image_size, self.image_size)
|
| 315 |
+
cls_row = emb.cls_token.detach() + pos[:, :1]
|
| 316 |
+
patch_rows = pos[:, 1:] + conv.bias.detach().reshape(1, 1, -1)
|
| 317 |
+
pos_full = emb.window_partition(torch.cat([cls_row, patch_rows], dim=1), self.image_size, self.image_size)
|
| 318 |
+
assert tuple(pos_full.shape) == (self.nw2, self.seq_per_window, self.hidden)
|
| 319 |
+
self.patch_pos = ttnn.from_torch(
|
| 320 |
+
pos_full, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=self.device, memory_config=ttnn.DRAM_MEMORY_CONFIG
|
| 321 |
+
)
|
| 322 |
+
self.patch_ones = ttnn.from_torch(
|
| 323 |
+
torch.ones(1, 1, self.hidden), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=self.device,
|
| 324 |
+
memory_config=ttnn.DRAM_MEMORY_CONFIG,
|
| 325 |
+
)
|
| 326 |
+
self.patch_compute_config = ttnn.init_device_compute_kernel_config(
|
| 327 |
+
self.device.arch(), math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True, packer_l1_acc=True
|
| 328 |
+
)
|
| 329 |
+
self._im2col_buf = torch.zeros(self.nw2, self.seq_per_window, self.patch_dim_padded, dtype=torch.float32)
|
| 330 |
+
|
| 331 |
+
def _im2col_host(self, pixel_values):
|
| 332 |
+
"""[1, 3, 560, 560] -> the persistent fp32 im2col buffer [16, 101, 608] in window order: window
|
| 333 |
+
(wh, ww) raster, token (h_pw, w_pw) raster after the zero cls row, features (c, kh, kw) like the
|
| 334 |
+
flattened conv kernel. One strided copy (casts if pixel_values is not fp32)."""
|
| 335 |
+
if pixel_values.shape[0] != 1:
|
| 336 |
+
raise ValueError("the device graph is built for batch 1")
|
| 337 |
+
c, p, nw = self.cfg.num_channels, self.cfg.patch_size, self.num_windows
|
| 338 |
+
w = self.grid // nw # patches per window side
|
| 339 |
+
src = pixel_values.reshape(c, nw, w, p, nw, w, p).permute(1, 4, 2, 5, 0, 3, 6) # wh ww hpw wpw c kh kw
|
| 340 |
+
self._im2col_buf[:, 1:, : self.patch_dim].view(nw, nw, w, w, c, p, p).copy_(src)
|
| 341 |
+
return self._im2col_buf
|
| 342 |
+
|
| 343 |
+
def host_input(self, pixel_values):
|
| 344 |
+
"""Per-image host work -> host ttnn tensor to upload (the trace's persistent input has this spec)."""
|
| 345 |
+
if self.input_path == "patch":
|
| 346 |
+
return ttnn.from_torch(self._im2col_host(pixel_values), dtype=ttnn.float32, layout=ttnn.ROW_MAJOR_LAYOUT)
|
| 347 |
+
return self.host_input_from_embed(self.wb.embeddings(pixel_values))
|
| 348 |
+
|
| 349 |
+
def host_input_from_embed(self, embed):
|
| 350 |
+
"""Host embed [16, 101, 384] (torch) -> host ttnn tensor, "rowmajor" or "tile" style ("patch" has no
|
| 351 |
+
embed input; it uses the exact rowmajor upload, which is what ``run_layers`` wants)."""
|
| 352 |
+
if self.input_path == "tile":
|
| 353 |
+
return ttnn.from_torch(embed.reshape(self.input_shape), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT)
|
| 354 |
+
return ttnn.from_torch(
|
| 355 |
+
embed.reshape(self.nw2, self.seq_per_window, self.hidden).to(torch.bfloat16),
|
| 356 |
+
dtype=ttnn.bfloat16, layout=ttnn.ROW_MAJOR_LAYOUT,
|
| 357 |
)
|
| 358 |
+
|
| 359 |
+
def _ingest(self, x):
|
| 360 |
+
"""Device input tensor (any input path, recognised by layout/shape) -> windowed [16, 101, 384] TILE
|
| 361 |
+
(padded [16, 128, 384]) in ``self.mem``; traced with the rest of the graph."""
|
| 362 |
+
if x.layout == ttnn.ROW_MAJOR_LAYOUT:
|
| 363 |
+
x = ttnn.tilize_with_zero_padding(x, memory_config=self.mem, use_multicore=True)
|
| 364 |
+
if x.shape[2] == self.patch_dim_padded: # im2col -> patch embedding
|
| 365 |
+
return ttnn.experimental.dit_minimal_matmul_addcmul_fused(
|
| 366 |
+
x, self.patch_w, 1.0, self.patch_pos, self.patch_ones,
|
| 367 |
+
bias_tensor=None, config=self.dit_config, memory_config=self.mem, dtype=ttnn.bfloat16,
|
| 368 |
+
compute_kernel_config=self.patch_compute_config,
|
| 369 |
+
)
|
| 370 |
+
return x
|
| 371 |
+
if x.shape[0] == 1: # merged TILE upload
|
| 372 |
+
return ttnn.reshape(x, (self.nw2, self.seq_per_window, self.hidden), memory_config=self.mem)
|
| 373 |
+
return x
|
| 374 |
+
|
| 375 |
+
# ------------------------------------------------------------------ layers
|
| 376 |
+
def _merge_windows(self, x):
|
| 377 |
+
"""[16, 101, 384] -> [1, 1616, 384] (the global-attention / shaping layout)."""
|
| 378 |
+
b, s, c = x.shape
|
| 379 |
+
return ttnn.reshape(x, (b // self.nw2, self.nw2 * s, c))
|
| 380 |
+
|
| 381 |
+
def _attention(self, normed, p, is_global):
|
| 382 |
+
"""norm1 output [b, s, 384] -> multi-head self-attention context [b, s, 384] (heads concatenated)."""
|
| 383 |
+
cc = self.compute_config
|
| 384 |
+
mc = self.mem
|
| 385 |
+
qkv = self._matmul(normed, p["qkv_w"], p["qkv_b"])
|
| 386 |
+
if self.attn == "sdpa":
|
| 387 |
+
q, k, v = ttnn.transformer.split_query_key_value_and_split_heads(
|
| 388 |
+
qkv, num_heads=self.num_heads, transpose_key=False
|
| 389 |
+
) # 3 x [b,h,s,d]
|
| 390 |
+
ctx = ttnn.transformer.scaled_dot_product_attention(
|
| 391 |
+
q, k, v,
|
| 392 |
+
is_causal=False,
|
| 393 |
+
scale=self.head_dim ** -0.5,
|
| 394 |
+
program_config=self.sdpa_pc_global if is_global else self.sdpa_pc_window,
|
| 395 |
+
compute_kernel_config=cc,
|
| 396 |
+
memory_config=mc,
|
| 397 |
+
)
|
| 398 |
+
else:
|
| 399 |
+
q, k, v = ttnn.transformer.split_query_key_value_and_split_heads(
|
| 400 |
+
qkv, num_heads=self.num_heads, transpose_key=True
|
| 401 |
+
)
|
| 402 |
+
scores = ttnn.matmul(q, k, compute_kernel_config=cc, memory_config=mc)
|
| 403 |
+
scores = ttnn.multiply(scores, self.head_dim ** -0.5, memory_config=mc)
|
| 404 |
+
probs = ttnn.softmax(scores, dim=-1, memory_config=mc)
|
| 405 |
+
ctx = ttnn.matmul(probs, v, compute_kernel_config=cc, memory_config=mc) # [b,h,s,d]
|
| 406 |
ctx = ttnn.transformer.concatenate_heads(ctx) # [b,s,384]
|
| 407 |
ttnn.deallocate(qkv)
|
| 408 |
+
return ctx
|
| 409 |
+
|
| 410 |
+
@staticmethod
|
| 411 |
+
def _ls(p, key, residual):
|
| 412 |
+
"""Layerscale vector in the same buffer type as ``residual`` (layer 0's residual is the DRAM
|
| 413 |
+
input buffer, later residuals live in L1 when l1=True)."""
|
| 414 |
+
if residual.memory_config().buffer_type == ttnn.BufferType.L1:
|
| 415 |
+
return p[key + "_l1"]
|
| 416 |
+
return p[key]
|
| 417 |
+
|
| 418 |
+
def _matmul(self, h, w, b):
|
| 419 |
+
"""h @ w + b (qkv, fc1)."""
|
| 420 |
+
if self.matmul == "minimal":
|
| 421 |
+
return ttnn.experimental.minimal_matmul(
|
| 422 |
+
h, w, bias_tensor=b, config=self.mm_config, memory_config=self.mem,
|
| 423 |
+
compute_kernel_config=self.mm_compute_config,
|
| 424 |
+
)
|
| 425 |
+
return ttnn.linear(h, w, bias=b, compute_kernel_config=self.compute_config, memory_config=self.mem)
|
| 426 |
+
|
| 427 |
+
def _matmul_ls_residual(self, h, w, b, ls_key, p, residual):
|
| 428 |
+
"""residual + layerscale * (h @ w + b) (proj, fc2): one fused op, or linear -> multiply -> add."""
|
| 429 |
+
mc = self.mem
|
| 430 |
+
if self.matmul == "minimal":
|
| 431 |
+
return ttnn.experimental.dit_minimal_matmul_addcmul_fused(
|
| 432 |
+
h, w, 1.0, residual, self._ls(p, ls_key, residual),
|
| 433 |
+
bias_tensor=b,
|
| 434 |
+
config=self.dit_config,
|
| 435 |
+
memory_config=mc,
|
| 436 |
+
compute_kernel_config=self.mm_compute_config,
|
| 437 |
+
)
|
| 438 |
+
y = ttnn.linear(h, w, bias=b, compute_kernel_config=self.compute_config, memory_config=mc)
|
| 439 |
+
y = ttnn.multiply(y, p[ls_key], memory_config=mc)
|
| 440 |
+
return ttnn.add(y, residual, memory_config=mc)
|
| 441 |
+
|
| 442 |
+
def _layer(self, x, p, x_merged=None):
|
| 443 |
+
"""One DINO block. ``x``: [16, 101, 384]. ``x_merged`` may pass the already-merged view of ``x``
|
| 444 |
+
for a global layer.
|
| 445 |
+
|
| 446 |
+
Global layers run the WHOLE block (attention + MLP + both residuals) in the merged
|
| 447 |
+
[1, 1616, 384] layout and reshape back to the windowed layout once at the end: the MLP is
|
| 448 |
+
row-wise, so this is the same math, but a [1632, 1536] x [1536, 384] matmul is ~3x faster
|
| 449 |
+
than the batched [16, 128, 1536] x [1536, 384] one and the 21% window padding is not
|
| 450 |
+
computed (windowed layers cannot avoid it: their attention needs the [16, ...] layout).
|
| 451 |
+
"""
|
| 452 |
+
windowed_shape = (x.shape[0], x.shape[1], x.shape[2])
|
| 453 |
if p["global"]:
|
| 454 |
+
x = x_merged if x_merged is not None else self._merge_windows(x)
|
| 455 |
+
mc = self.mem # None => DRAM default; L1_MEMORY_CONFIG keeps the working set on-chip
|
| 456 |
+
|
| 457 |
+
# ---- attention (norm1 -> MHA -> proj -> layerscale1 -> residual) ----
|
| 458 |
+
normed = ttnn.layer_norm(x, weight=p["norm1_w"], bias=p["norm1_b"], epsilon=self.eps, memory_config=mc)
|
| 459 |
+
ctx = self._attention(normed, p, p["global"]) # [b,s,384]
|
| 460 |
+
x = self._matmul_ls_residual(ctx, p["proj_w"], p["proj_b"], "ls1", p, x)
|
| 461 |
|
| 462 |
# ---- mlp (norm2 -> fc1 -> gelu -> fc2 -> layerscale2 -> residual) ----
|
|
|
|
| 463 |
h = ttnn.layer_norm(x, weight=p["norm2_w"], bias=p["norm2_b"], epsilon=self.eps, memory_config=mc)
|
| 464 |
+
h = self._matmul(h, p["fc1_w"], p["fc1_b"])
|
| 465 |
+
h = ttnn.gelu(h, fast_and_approximate_mode=False, memory_config=mc) # exact; fusing it is slower
|
| 466 |
+
x = self._matmul_ls_residual(h, p["fc2_w"], p["fc2_b"], "ls2", p, x)
|
| 467 |
+
if p["global"]:
|
| 468 |
+
x = ttnn.reshape(x, windowed_shape) # back to [16, 101, 384] for the next windowed layer
|
| 469 |
return x
|
| 470 |
|
| 471 |
+
# --------------------------------------------------------- device shaping
|
| 472 |
+
def shape_features(self, x_merged, apply_norm=True):
|
| 473 |
+
"""[1, 1616, 384] merged hidden state -> feature map [1, 1600, 384], channels-last, raster order.
|
| 474 |
+
|
| 475 |
+
``P @ X`` drops the per-window cls rows and re-orders windows (exact row gather), then the
|
| 476 |
+
backbone's final layernorm. ``apply_norm=False`` returns the bare gather (for exactness tests).
|
| 477 |
+
"""
|
| 478 |
+
y = ttnn.matmul(
|
| 479 |
+
self.perm, x_merged, compute_kernel_config=self.perm_compute_config, memory_config=self.mem
|
| 480 |
+
) # [1, 1600, 384]
|
| 481 |
+
if not apply_norm:
|
| 482 |
+
return y
|
| 483 |
+
return ttnn.layer_norm(
|
| 484 |
+
y, weight=self.final_norm_w, bias=self.final_norm_b, epsilon=self.final_eps, memory_config=self.mem
|
| 485 |
)
|
| 486 |
+
|
| 487 |
+
def forward_device(self, x):
|
| 488 |
+
"""Whole device backbone. ``x``: the device input tensor of the configured input path (the fp32
|
| 489 |
+
ROW_MAJOR im2col [16, 101, 608], the bf16 ROW_MAJOR embed [16, 101, 384], or the TILE embed in
|
| 490 |
+
``input_shape``), see ``_ingest``.
|
| 491 |
+
|
| 492 |
+
Returns the 4 shaped feature maps as device tensors [1, 1600, 384] (channels-last, raster
|
| 493 |
+
order, final LN applied) -- no host round trip, so it is metal-trace-able end to end.
|
| 494 |
+
The merged view built for the shaping is reused by the global layer that follows; the
|
| 495 |
+
loop ends with the last out layer (10), see the module docstring.
|
| 496 |
+
"""
|
| 497 |
+
feats = []
|
| 498 |
+
merged = None
|
| 499 |
+
x = self._ingest(x)
|
| 500 |
+
for i, p in enumerate(self.layers):
|
| 501 |
+
x = self._layer(x, p, x_merged=merged)
|
| 502 |
+
merged = None
|
| 503 |
+
if i in self.out_layers:
|
| 504 |
+
merged = self._merge_windows(x)
|
| 505 |
+
feats.append(self.shape_features(merged))
|
| 506 |
+
return feats
|
| 507 |
+
|
| 508 |
+
def feature_maps(self, pixel_values):
|
| 509 |
+
"""Full backbone through the configured input path: host_input -> device ingest + layers + shaping
|
| 510 |
+
-> host. Returns list of 4 torch tensors [1,384,40,40] (reference layout)."""
|
| 511 |
+
x = ttnn.to_device(self.host_input(pixel_values), self.device)
|
| 512 |
+
feats = self.forward_device(x)
|
| 513 |
+
n = self.grid
|
| 514 |
+
return [
|
| 515 |
+
ttnn.to_torch(f).float().reshape(pixel_values.shape[0], n, n, -1).permute(0, 3, 1, 2).contiguous()
|
| 516 |
+
for f in feats
|
| 517 |
+
]
|
| 518 |
+
|
| 519 |
+
# ------------------------------------------------ eager / host debug path
|
| 520 |
+
def run_layers(self, embed_windowed_torch):
|
| 521 |
+
"""Debug path: embed [16, 101, 384] -> dict idx->torch hidden after layers 1,4,7,10 (readbacks).
|
| 522 |
+
Starts from the given (exact) embed, so it isolates the layers from the input path."""
|
| 523 |
+
x = ttnn.to_device(self.host_input_from_embed(embed_windowed_torch), self.device)
|
| 524 |
+
x = self._ingest(x)
|
| 525 |
out = {}
|
| 526 |
for i, p in enumerate(self.layers):
|
| 527 |
x = self._layer(x, p)
|
| 528 |
+
if i in self.out_layers:
|
| 529 |
out[i] = ttnn.to_torch(x).float()
|
| 530 |
return out
|
| 531 |
|
| 532 |
+
def feature_maps_host(self, pixel_values):
|
| 533 |
+
"""Debug path (the pre-S1 pipeline): device layers with 4 readbacks + host shaping
|
| 534 |
+
(torch LN + drop cls + window_unpartition). Returns list of 4 torch tensors [1,384,40,40]."""
|
| 535 |
embed = self.wb.embeddings(pixel_values) # [16,101,384] host
|
| 536 |
hidden = self.run_layers(embed)
|
| 537 |
_, _, H, W = pixel_values.shape
|
| 538 |
feats = []
|
| 539 |
+
for i in self.out_layers:
|
| 540 |
+
hs = hidden[i]
|
| 541 |
hs = self.wb.layernorm(hs)
|
| 542 |
hs = hs[:, 1:] # drop cls per window
|
| 543 |
hs = self.wb.window_unpartition(hs, H, W)
|
code/rf_detr/tt/ttnn_rf_detr.py
CHANGED
|
@@ -1,23 +1,35 @@
|
|
| 1 |
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
"""End-to-end RF-DETR on Tenstorrent.
|
| 3 |
|
| 4 |
-
On-device chain:
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
"""
|
| 22 |
|
| 23 |
import os
|
|
@@ -33,66 +45,95 @@ NUM_CLASSES = 91
|
|
| 33 |
|
| 34 |
|
| 35 |
class TtRfDetr:
|
| 36 |
-
def __init__(self, ref_model, device):
|
| 37 |
self.ref = ref_model.eval()
|
| 38 |
self.device = device
|
| 39 |
-
# Backbone
|
| 40 |
-
# RFDETR_BB_FIDELITY = LoFi|HiFi2|HiFi4 ; RFDETR_BB_FP32ACC = 1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
_fid = os.environ.get("RFDETR_BB_FIDELITY")
|
| 42 |
_mf = getattr(ttnn.MathFidelity, _fid) if _fid else None
|
| 43 |
_fp32 = os.environ.get("RFDETR_BB_FP32ACC", "0") == "1"
|
| 44 |
-
_l1 = os.environ.get("RFDETR_BB_L1", "
|
| 45 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
self.projector = TtProjector(ref_model, device)
|
| 47 |
self.transformer = TtTransformer(ref_model, device)
|
| 48 |
|
| 49 |
device.enable_program_cache()
|
| 50 |
|
|
|
|
|
|
|
| 51 |
self._trace_id = None
|
| 52 |
-
self._persistent_in = None #
|
| 53 |
self._logits_out = None # device tensor (trace output)
|
| 54 |
self._boxes_out = None # device tensor (trace output)
|
| 55 |
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 60 |
|
| 61 |
-
|
|
|
|
| 62 |
device = self.device
|
| 63 |
-
# Persistent device input
|
| 64 |
-
self._persistent_in =
|
| 65 |
-
ttnn.from_torch(f, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device)
|
| 66 |
-
for f in feats_cl_host
|
| 67 |
-
]
|
| 68 |
|
| 69 |
# Warm run (eager) so conv2d prepared-weights are cached and the program cache
|
| 70 |
# is populated; mutating ops (p["w"] = prepared) must run OUTSIDE the capture.
|
| 71 |
-
|
| 72 |
-
logits, boxes = self.transformer.forward_device(source)
|
| 73 |
ttnn.synchronize_device(device)
|
|
|
|
|
|
|
| 74 |
|
| 75 |
-
# Capture the projector + transformer device graph.
|
| 76 |
self._trace_id = ttnn.begin_trace_capture(device, cq_id=0)
|
| 77 |
-
|
| 78 |
-
self._logits_out, self._boxes_out = self.transformer.forward_device(source)
|
| 79 |
ttnn.end_trace_capture(device, self._trace_id, cq_id=0)
|
| 80 |
ttnn.synchronize_device(device)
|
| 81 |
|
| 82 |
-
def _run_trace(self,
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
ttnn.
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
return logits_t, boxes_t
|
| 92 |
|
|
|
|
| 93 |
def __call__(self, pixel_values):
|
| 94 |
-
|
| 95 |
-
if self.
|
| 96 |
-
self.
|
| 97 |
-
|
|
|
|
|
|
|
|
|
|
| 98 |
return RfDetrOutput(logits=logits, pred_boxes=pred_boxes)
|
|
|
|
| 1 |
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
"""End-to-end RF-DETR on Tenstorrent.
|
| 3 |
|
| 4 |
+
On-device chain: patch embedding (S4) -> windowed DINOv2 backbone (12 layers + device
|
| 5 |
+
feature shaping) -> C2f projector -> two-stage deformable transformer + heads. The only
|
| 6 |
+
host glue left is an im2col of the 560x560 image into a persistent fp32 ``[16, 101, 608]``
|
| 7 |
+
buffer (one strided copy) that is uploaded ROW_MAJOR; tilize + the patch-embed matmul
|
| 8 |
+
(+ cls/pos) run inside the trace (see ``TtDinoBackbone``, knob ``RFDETR_BB_INPUT``;
|
| 9 |
+
``rowmajor`` keeps the S1-S3 host embeddings but skips the host tilize, bit-identical to
|
| 10 |
+
``tile``).
|
| 11 |
+
|
| 12 |
+
The WHOLE device graph is captured into ONE metal-trace and replayed per inference:
|
| 13 |
+
``__call__`` builds the host input (``backbone.host_input``), ``copy_host_to_device_tensor``
|
| 14 |
+
into the persistent device input buffer the trace reads from, ``execute_trace``, then two
|
| 15 |
+
``to_torch`` readbacks (logits, boxes). There is no
|
| 16 |
+
mid-graph host work any more: the pre-S1 pipeline read the hidden state back after
|
| 17 |
+
layers 1/4/7/10, did LN + drop-cls + window_unpartition on host and re-uploaded 4 maps
|
| 18 |
+
(~8 ms of syncs, host shaping and host tilize per image); that shaping now runs on
|
| 19 |
+
device (``TtDinoBackbone.shape_features``), so the backbone's ops no longer pay host
|
| 20 |
+
dispatch either. The backbone layer itself is 10 fused ops (SDPA attention, matmul +
|
| 21 |
+
layerscale + residual fused, minimal_matmul kernels; layer 11 skipped as unused), see
|
| 22 |
+
``ttnn_backbone.py``: 39.6 ms (eager, host shaping) -> 25.5 (S1, one trace) -> 14.4 ms
|
| 23 |
+
(S2) -> 13.9 ms (S3, merged embed upload) -> 12.6 ms (S4a, no host tilize) -> 11.8 ms
|
| 24 |
+
(S4b, patch embedding on device) per image on p150a at detection-IoU 98.71 / 98.74.
|
| 25 |
+
|
| 26 |
+
Debug paths: ``TtRfDetr(..., use_trace=False)`` or env ``RFDETR_EAGER=1`` runs the same
|
| 27 |
+
device graph eagerly (no trace); ``TtDinoBackbone.feature_maps_host`` keeps the old
|
| 28 |
+
host-shaping pipeline for A/B comparisons.
|
| 29 |
+
|
| 30 |
+
2-CQ overlap (upload on cq_id=1 + event so compute waits) was implemented and measured
|
| 31 |
+
on the old pipeline: it *regressed* FPS because the benchmark runs each inference fully
|
| 32 |
+
synchronized (no cross-inference pipelining to overlap). The device stays on a single CQ.
|
| 33 |
"""
|
| 34 |
|
| 35 |
import os
|
|
|
|
| 45 |
|
| 46 |
|
| 47 |
class TtRfDetr:
|
| 48 |
+
def __init__(self, ref_model, device, use_trace=None):
|
| 49 |
self.ref = ref_model.eval()
|
| 50 |
self.device = device
|
| 51 |
+
# Backbone knobs, env-tunable:
|
| 52 |
+
# RFDETR_BB_FIDELITY = LoFi|HiFi2|HiFi4 (default: ttnn default) ; RFDETR_BB_FP32ACC = 1 (default 0)
|
| 53 |
+
# RFDETR_BB_L1 = 0|1 (default 1: backbone working set in L1; 31.0 -> 26.0 ms on the traced graph)
|
| 54 |
+
# RFDETR_BB_ATTN = sdpa|matmul (default sdpa: fused scaled_dot_product_attention kernel)
|
| 55 |
+
# RFDETR_BB_MATMUL = minimal|linear (default minimal: minimal_matmul kernel for qkv/fc1 and the
|
| 56 |
+
# fused matmul+layerscale+residual op for proj/fc2; linear = ttnn.linear chain)
|
| 57 |
+
# RFDETR_BB_INPUT = patch|rowmajor|tile (default patch: fp32 im2col upload + patch-embed matmul on
|
| 58 |
+
# device; rowmajor = host embed uploaded bf16 ROW_MAJOR + device tilize, bit-exact
|
| 59 |
+
# vs tile; tile = the S1-S3 host tilize). Read by TtDinoBackbone (tests follow it).
|
| 60 |
+
# RFDETR_BB_UPLOAD = merged|windowed (tile path only, default merged: embed uploaded as [1,1616,384]
|
| 61 |
+
# + one device reshape)
|
| 62 |
_fid = os.environ.get("RFDETR_BB_FIDELITY")
|
| 63 |
_mf = getattr(ttnn.MathFidelity, _fid) if _fid else None
|
| 64 |
_fp32 = os.environ.get("RFDETR_BB_FP32ACC", "0") == "1"
|
| 65 |
+
_l1 = os.environ.get("RFDETR_BB_L1", "1") == "1"
|
| 66 |
+
_attn = os.environ.get("RFDETR_BB_ATTN", "sdpa")
|
| 67 |
+
_mm = os.environ.get("RFDETR_BB_MATMUL", "minimal")
|
| 68 |
+
self.backbone = TtDinoBackbone(
|
| 69 |
+
ref_model, device, math_fidelity=_mf, fp32_acc=_fp32, l1=_l1, attn=_attn, matmul=_mm
|
| 70 |
+
)
|
| 71 |
self.projector = TtProjector(ref_model, device)
|
| 72 |
self.transformer = TtTransformer(ref_model, device)
|
| 73 |
|
| 74 |
device.enable_program_cache()
|
| 75 |
|
| 76 |
+
# Trace is the production path; RFDETR_EAGER=1 (or use_trace=False) runs the graph eagerly.
|
| 77 |
+
self.use_trace = (os.environ.get("RFDETR_EAGER", "0") != "1") if use_trace is None else bool(use_trace)
|
| 78 |
self._trace_id = None
|
| 79 |
+
self._persistent_in = None # persistent device input (spec = backbone.host_input's tensor)
|
| 80 |
self._logits_out = None # device tensor (trace output)
|
| 81 |
self._boxes_out = None # device tensor (trace output)
|
| 82 |
|
| 83 |
+
# ------------------------------------------------------------------ pieces
|
| 84 |
+
def _embed_host(self, pixel_values):
|
| 85 |
+
"""Per-image host work -> host ttnn tensor (not yet uploaded): the fp32 ROW_MAJOR im2col
|
| 86 |
+
[16, 101, 608] (input path "patch", default), or the bf16 embed (ROW_MAJOR [16, 101, 384] / TILE)."""
|
| 87 |
+
return self.backbone.host_input(pixel_values)
|
| 88 |
+
|
| 89 |
+
def _device_graph(self, embed_dev):
|
| 90 |
+
"""The whole device graph: backbone layers + shaping -> projector -> transformer. Device in/out."""
|
| 91 |
+
feats = self.backbone.forward_device(embed_dev) # 4 x [1,1600,384]
|
| 92 |
+
source = self.projector(feats) # [1,1600,256]
|
| 93 |
+
return self.transformer.forward_device(source) # device (logits, pred_boxes)
|
| 94 |
+
|
| 95 |
+
@staticmethod
|
| 96 |
+
def _read_outputs(logits_dev, boxes_dev):
|
| 97 |
+
logits_t = ttnn.to_torch(logits_dev).float().reshape(1, N_QUERIES, NUM_CLASSES)
|
| 98 |
+
boxes_t = ttnn.to_torch(boxes_dev).float().reshape(1, N_QUERIES, 4)
|
| 99 |
+
return logits_t, boxes_t
|
| 100 |
|
| 101 |
+
# ------------------------------------------------------------------- trace
|
| 102 |
+
def _capture_trace(self, embed_host):
|
| 103 |
device = self.device
|
| 104 |
+
# Persistent device input buffer the trace reads from.
|
| 105 |
+
self._persistent_in = ttnn.to_device(embed_host, device)
|
|
|
|
|
|
|
|
|
|
| 106 |
|
| 107 |
# Warm run (eager) so conv2d prepared-weights are cached and the program cache
|
| 108 |
# is populated; mutating ops (p["w"] = prepared) must run OUTSIDE the capture.
|
| 109 |
+
logits, boxes = self._device_graph(self._persistent_in)
|
|
|
|
| 110 |
ttnn.synchronize_device(device)
|
| 111 |
+
ttnn.deallocate(logits)
|
| 112 |
+
ttnn.deallocate(boxes)
|
| 113 |
|
| 114 |
+
# Capture the whole backbone + projector + transformer device graph.
|
| 115 |
self._trace_id = ttnn.begin_trace_capture(device, cq_id=0)
|
| 116 |
+
self._logits_out, self._boxes_out = self._device_graph(self._persistent_in)
|
|
|
|
| 117 |
ttnn.end_trace_capture(device, self._trace_id, cq_id=0)
|
| 118 |
ttnn.synchronize_device(device)
|
| 119 |
|
| 120 |
+
def _run_trace(self, embed_host):
|
| 121 |
+
ttnn.copy_host_to_device_tensor(embed_host, self._persistent_in, cq_id=0)
|
| 122 |
+
ttnn.execute_trace(self.device, self._trace_id, cq_id=0, blocking=False)
|
| 123 |
+
return self._read_outputs(self._logits_out, self._boxes_out)
|
| 124 |
+
|
| 125 |
+
def _run_eager(self, embed_host):
|
| 126 |
+
embed_dev = ttnn.to_device(embed_host, self.device)
|
| 127 |
+
logits, boxes = self._device_graph(embed_dev)
|
| 128 |
+
return self._read_outputs(logits, boxes)
|
|
|
|
| 129 |
|
| 130 |
+
# -------------------------------------------------------------------- call
|
| 131 |
def __call__(self, pixel_values):
|
| 132 |
+
embed_host = self._embed_host(pixel_values)
|
| 133 |
+
if not self.use_trace:
|
| 134 |
+
logits, pred_boxes = self._run_eager(embed_host)
|
| 135 |
+
else:
|
| 136 |
+
if self._trace_id is None:
|
| 137 |
+
self._capture_trace(embed_host)
|
| 138 |
+
logits, pred_boxes = self._run_trace(embed_host)
|
| 139 |
return RfDetrOutput(logits=logits, pred_boxes=pred_boxes)
|
image/blobs/sha256/0aa064431bd920ad87ca8f42142afb0bbc909e68bc589afbbc6ff87800ac1153
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0aa064431bd920ad87ca8f42142afb0bbc909e68bc589afbbc6ff87800ac1153
|
| 3 |
+
size 187652608
|
image/blobs/sha256/2cbaee648501aee7415c5bb636fcfd6efdfe01bd976f1e3b9d8fd5be48d3bc00
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2cbaee648501aee7415c5bb636fcfd6efdfe01bd976f1e3b9d8fd5be48d3bc00
|
| 3 |
+
size 454599168
|
image/blobs/sha256/4186f6a9387df261e246adfcc37ae3225f74d40791c2dced1a46fc9d12b6398c
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:4186f6a9387df261e246adfcc37ae3225f74d40791c2dced1a46fc9d12b6398c
|
| 3 |
+
size 130853888
|
image/blobs/sha256/6134c251abc1a162f6d394d0b13369424783eee9bf3c4c4ced16f162a4fad452
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6134c251abc1a162f6d394d0b13369424783eee9bf3c4c4ced16f162a4fad452
|
| 3 |
+
size 198656
|
image/blobs/sha256/77872ef52f01066f9ad1825a9a71a56d304560f756a575654f51973aa59d9ac4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:77872ef52f01066f9ad1825a9a71a56d304560f756a575654f51973aa59d9ac4
|
| 3 |
+
size 130853888
|
image/blobs/sha256/85e1ed2c56959e565b7ed00246b276c265f12dfab1060197f8c2969868ef90aa
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:85e1ed2c56959e565b7ed00246b276c265f12dfab1060197f8c2969868ef90aa
|
| 3 |
+
size 191707136
|
image/blobs/sha256/a435df3ddd972def28ddfda336794ad5391073fd1a2c86ee50e78fc2d5a4fd46
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a435df3ddd972def28ddfda336794ad5391073fd1a2c86ee50e78fc2d5a4fd46
|
| 3 |
+
size 4678144
|
image/blobs/sha256/d0c2d85cc7e14bb778666593a7c2b8307242528023ca6f5ec53497244b5f5ea8
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d0c2d85cc7e14bb778666593a7c2b8307242528023ca6f5ec53497244b5f5ea8
|
| 3 |
+
size 65318400
|
image/blobs/sha256/f5c0a29ae70a96b2a3a149b03201b80a49e1052dd91e0bf657f3c6658e664d33
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:f5c0a29ae70a96b2a3a149b03201b80a49e1052dd91e0bf657f3c6658e664d33
|
| 3 |
+
size 218112
|
image/blobs/sha256/fc2482b5110947c3e5ce45808795a1ba8be69d4c095f3482e1c83fa096562563
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fc2482b5110947c3e5ce45808795a1ba8be69d4c095f3482e1c83fa096562563
|
| 3 |
+
size 1036517888
|