changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
14.6 kB
# SPDX-License-Identifier: Apache-2.0
"""C24: SECOND backbone + SECONDFPN neck on the device (PLAN.md section 1.2 row C24; CenterPoint, PointPainting,
TransFusion and BEVFusion use them), built from BN-folded host weights on the C17 builders (``ttaw.ops.conv``).
- :class:`Second`: blocks of 3x3 conv (zero padding 1) + ReLU; the first conv of a block carries the block stride
(mmdet3d ``SECOND``: ``layer_strides`` [1, 2, 2], CenterPoint-tiny [2, 2, 2]). Returns every block output.
- :class:`SecondFPN`: one deblock per block output, each up-sampling to the same grid, + ReLU (mmdet3d
``SECONDFPN``): a ``"conv"`` deblock is a 1x1 stride-1 conv (a :class:`RowLinear`), a ``"conv_transpose"`` deblock
a ConvTranspose with kernel == stride (k = 1 a :class:`RowLinear`, k = 2 ``ttnn.conv_transpose2d``, k >= 3
linear + depth-to-space: C17 ``ConvTranspose2d``, probe P5). The concat of the deblock outputs is NOT built: the
consumer takes the list (C25's K-split shared conv reads the parts, so CenterPoint's 177 MB concat never exists).
Placement: every module output is a DRAM interleaved TILE ``[1, 1, N*H*W, C]`` map (the robust first-port layout:
3x3 convs on DRAM inputs run ttnn's automatic DRAM slicing, probe P14; L1 residency is optimization work).
Precision: :data:`DEFAULT_PRECISION` (C17's ``CONV_PRECISION``, HiFi4 + fp32 accumulation + packer L1
accumulation: without it a conv keeps its partial sums in bf16 between inner blocks) unless a
``ttaw.precision.PrecisionPolicy`` says otherwise for ``<prefix>.block<i>.conv<j>`` / ``<neck>.deblock<i>``; a
precision's ``activations`` dtype (``:a=fp32``) is the dtype of that layer's OUTPUT (bf16 by default), the knob of a
deep plain conv stack whose bf16 activation rounding compounds (S:centerpoint:396-408).
:class:`ConvSpec` records are plain numpy, so a port maps its own weight loader onto them; :func:`second_numpy` /
:func:`second_fpn_numpy` are the fp32 torch oracles of the device modules.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Sequence
import numpy as np
from ..ops.conv import CONV_PRECISION, Conv2d, ConvTranspose2d, FeatureMap, is_dram
from ..precision import Precision, PrecisionPolicy
__all__ = ["ConvSpec", "RowLinear", "Second", "SecondFPN", "DEFAULT_PRECISION", "default_policy", "second_numpy",
"second_fpn_numpy", "free_maps"]
DEFAULT_PRECISION = CONV_PRECISION
@dataclass(frozen=True)
class ConvSpec:
"""One conv layer with BN folded: ``y = relu?(conv(x, weight) + bias)``.
``kind="conv"``: weight ``[Cout, Cin, k, k]``; ``kind="conv_transpose"``: weight ``[Cin, Cout, s, s]`` (kernel ==
stride, no padding). ``padding`` None = ``k // 2`` for convs."""
weight: np.ndarray
bias: Optional[np.ndarray]
stride: int = 1
padding: Optional[int] = None
relu: bool = True
kind: str = "conv"
name: str = ""
@property
def kernel(self) -> int:
return int(np.shape(self.weight)[-1])
@property
def in_channels(self) -> int:
return int(np.shape(self.weight)[1 if self.kind == "conv" else 0])
@property
def out_channels(self) -> int:
return int(np.shape(self.weight)[0 if self.kind == "conv" else 1])
def default_policy() -> PrecisionPolicy:
""":data:`DEFAULT_PRECISION` everywhere (bf16 weights)."""
return PrecisionPolicy(default=DEFAULT_PRECISION)
def free_maps(*maps: Any) -> None:
"""Deallocate device feature maps / tensors that still hold their buffer (``force=False``: memory shared with
a view or another owner is left alone; ``ttnn.deallocate`` forces by default, probe P2)."""
import ttnn
for m in maps:
t = m.tensor if isinstance(m, FeatureMap) else m
if t is not None and t.is_allocated():
ttnn.deallocate(t, False)
def _to_dram(fm: FeatureMap) -> FeatureMap:
"""The map in DRAM (moved to DRAM interleaved when an op left it in L1, e.g. ``conv_transpose2d``'s sharded
output; a no-op for DRAM maps)."""
import ttnn
if is_dram(fm.tensor):
return fm
out = fm.with_tensor(ttnn.to_memory_config(fm.tensor, ttnn.DRAM_MEMORY_CONFIG))
free_maps(fm)
return out
class RowLinear:
"""``y = act(x @ W + b)`` on the rows of a feature map (a 1x1 stride-1 conv): ``[1, 1, R, K]`` TILE ->
``[1, 1, R, N]`` TILE in DRAM, through ``ttnn.linear`` with an explicit compute config and the device grid as
``core_grid`` (which keeps the activation fused in the matmul). ``weight``: ``[K, N]`` (in, out) or a 1x1 conv
weight ``[N, K, 1, 1]``. Device weights are uploaded by the first call (outside any capture)."""
def __init__(self, weight: Any, bias: Any = None, *, activation: Optional[str] = None,
precision: Any = DEFAULT_PRECISION, output_dtype: str = "bfloat16", name: str = "linear"):
w = np.asarray(weight, dtype=np.float32)
if w.ndim == 4:
if w.shape[2:] != (1, 1):
raise ValueError(f"{name}: a conv weight must be 1x1, got {w.shape}")
w = w[:, :, 0, 0].T
if w.ndim != 2:
raise ValueError(f"{name}: weight must be [K, N], got {w.shape}")
if activation not in (None, "relu"):
raise ValueError(f"{name}: activation {activation!r}: None or 'relu'")
self.name = name
self.weight_host = np.ascontiguousarray(w)
self.in_channels, self.out_channels = (int(s) for s in w.shape)
b = np.zeros(self.out_channels, np.float32) if bias is None else np.asarray(bias, np.float32)
if b.shape != (self.out_channels,):
raise ValueError(f"{name}: bias must be ({self.out_channels},), got {b.shape}")
if not (np.all(np.isfinite(w)) and np.all(np.isfinite(b))):
raise ValueError(f"{name}: non-finite weights")
self.bias_host = b
self.activation = activation
self.precision = Precision.parse(precision)
self.output_dtype = output_dtype
self._device: Optional[tuple] = None
def _upload(self, device: Any):
import torch
import ttnn
from ..tensors import ttnn_dtype
wdt = ttnn_dtype(self.precision.weights)
bdt = ttnn.bfloat16 if wdt == ttnn.bfloat8_b else wdt
g = device.compute_with_storage_grid_size()
self._device = (
ttnn.from_torch(torch.from_numpy(self.weight_host), dtype=wdt, layout=ttnn.TILE_LAYOUT, device=device,
memory_config=ttnn.DRAM_MEMORY_CONFIG),
ttnn.from_torch(torch.from_numpy(self.bias_host.reshape(1, -1)), dtype=bdt, layout=ttnn.TILE_LAYOUT,
device=device, memory_config=ttnn.DRAM_MEMORY_CONFIG),
ttnn.CoreGrid(y=int(g.y), x=int(g.x)),
self.precision.compute_kernel_config())
return self._device
def __call__(self, x: Any) -> Any:
"""``x``: a TILE device tensor ``[..., R, K]`` or a :class:`FeatureMap` (returned as one)."""
import ttnn
from ..tensors import ttnn_dtype
t = x.tensor if isinstance(x, FeatureMap) else x
if int(t.shape[-1]) != self.in_channels:
raise ValueError(f"{self.name}: input has {int(t.shape[-1])} channels, expected {self.in_channels}")
weight, bias, grid, cfg = self._device or self._upload(t.device())
y = ttnn.linear(t, weight, bias=bias, activation=self.activation, compute_kernel_config=cfg, core_grid=grid,
memory_config=ttnn.DRAM_MEMORY_CONFIG, dtype=ttnn_dtype(self.output_dtype))
return x.with_tensor(y, self.out_channels) if isinstance(x, FeatureMap) else y
def release(self) -> None:
if self._device is not None:
free_maps(self._device[0], self._device[1])
self._device = None
def describe(self) -> Dict[str, Any]:
return {"name": self.name, "cin": self.in_channels, "cout": self.out_channels, "activation": self.activation,
"precision": self.precision.label, "output_dtype": self.output_dtype}
def _precision(policy: Optional[PrecisionPolicy], name: str) -> Precision:
return policy.resolve(name) if policy is not None else Precision.parse(DEFAULT_PRECISION)
class Second:
"""SECOND backbone: ``blocks[i]`` is the list of :class:`ConvSpec` of block ``i`` (k x k convs, BN folded)."""
def __init__(self, blocks: Sequence[Sequence[ConvSpec]], *, policy: Optional[PrecisionPolicy] = None,
prefix: str = "backbone"):
self.prefix = prefix
self.blocks: List[List[Conv2d]] = []
for i, block in enumerate(blocks):
layers = []
for j, spec in enumerate(block):
if spec.kind != "conv":
raise ValueError(f"{prefix}.block{i}.conv{j}: SECOND layers are convs, got {spec.kind}")
name = f"{prefix}.block{i}.conv{j}"
pad = spec.kernel // 2 if spec.padding is None else int(spec.padding)
prec = _precision(policy, name)
layers.append(Conv2d(spec.weight, spec.bias, stride=spec.stride, padding=pad,
activation="relu" if spec.relu else None, precision=prec,
output_dtype=prec.activations, name=name))
self.blocks.append(layers)
def out_shapes(self, height: int, width: int) -> List[tuple]:
"""``(H, W, C)`` of every block output for an input of ``height x width``."""
shapes = []
for block in self.blocks:
for layer in block:
height, width = layer.out_hw(height, width)
shapes.append((height, width, block[-1].out_channels))
return shapes
def __call__(self, x: FeatureMap) -> List[FeatureMap]:
"""Run every block on ``x`` (kept); returns the block outputs (DRAM). Inner layer outputs are freed."""
outs: List[FeatureMap] = []
cur = x
for block in self.blocks:
for layer in block:
nxt = _to_dram(layer(cur))
if cur is not x and all(cur is not o for o in outs):
free_maps(cur)
cur = nxt
outs.append(cur)
return outs
def release(self) -> None:
for block in self.blocks:
for layer in block:
layer.release()
def describe(self) -> List[List[Dict[str, Any]]]:
return [[layer.describe() for layer in block] for block in self.blocks]
class SecondFPN:
"""SECONDFPN neck: ``deblocks[i]`` (a :class:`ConvSpec`, ReLU after) up-samples block output ``i``."""
def __init__(self, deblocks: Sequence[ConvSpec], *, policy: Optional[PrecisionPolicy] = None,
prefix: str = "neck"):
self.prefix = prefix
self.deblocks: List[Any] = []
self.strides: List[int] = []
for i, spec in enumerate(deblocks):
name = f"{prefix}.deblock{i}"
act = "relu" if spec.relu else None
prec = _precision(policy, name)
if spec.kind == "conv":
if spec.kernel != 1 or spec.stride != 1:
raise ValueError(f"{name}: a conv deblock must be 1x1 stride 1 (k={spec.kernel}, s={spec.stride})")
layer: Any = RowLinear(spec.weight, spec.bias, activation=act, precision=prec,
output_dtype=prec.activations, name=name)
self.strides.append(1)
elif spec.kind == "conv_transpose":
if spec.kernel != spec.stride:
raise ValueError(f"{name}: ConvTranspose deblocks need kernel == stride (k={spec.kernel}, "
f"s={spec.stride})")
if spec.stride == 1:
layer = RowLinear(np.asarray(spec.weight, np.float32)[:, :, 0, 0], spec.bias, activation=act,
precision=prec, output_dtype=prec.activations, name=name)
else:
layer = ConvTranspose2d(spec.weight, spec.bias, stride=spec.stride, activation=act,
precision=prec, output_dtype=prec.activations, name=name)
self.strides.append(spec.stride)
else:
raise ValueError(f"{name}: unknown kind {spec.kind!r}")
self.deblocks.append(layer)
def __call__(self, blocks: Sequence[FeatureMap]) -> List[FeatureMap]:
"""The deblock outputs (DRAM interleaved TILE maps of one grid); the inputs are kept."""
if len(blocks) != len(self.deblocks):
raise ValueError(f"{self.prefix}: {len(blocks)} inputs for {len(self.deblocks)} deblocks")
outs = []
for layer, x in zip(self.deblocks, blocks):
outs.append(_to_dram(layer(x)))
return outs
def release(self) -> None:
for d in self.deblocks:
d.release()
def describe(self) -> List[Dict[str, Any]]:
return [d.describe() for d in self.deblocks]
# ------------------------------------------------------------------------------------------- host oracles
def _apply_torch(spec: ConvSpec, x):
"""One layer in torch fp32 on an NCHW tensor."""
import torch
import torch.nn.functional as F
w = torch.from_numpy(np.ascontiguousarray(spec.weight, dtype=np.float32))
b = None if spec.bias is None else torch.from_numpy(np.ascontiguousarray(spec.bias, dtype=np.float32))
if spec.kind == "conv":
pad = spec.kernel // 2 if spec.padding is None else spec.padding
y = F.conv2d(x, w, b, stride=spec.stride, padding=pad)
else:
y = F.conv_transpose2d(x, w, b, stride=spec.stride)
return torch.relu(y) if spec.relu else y
def second_numpy(blocks: Sequence[Sequence[ConvSpec]], x_nchw: Any) -> List[np.ndarray]:
"""fp32 oracle of :class:`Second` (torch): NCHW in, NCHW block outputs out."""
import torch
with torch.no_grad():
x = torch.as_tensor(np.asarray(x_nchw, dtype=np.float32))
outs = []
for block in blocks:
for spec in block:
x = _apply_torch(spec, x)
outs.append(x.numpy().copy())
return outs
def second_fpn_numpy(deblocks: Sequence[ConvSpec], blocks_nchw: Sequence[Any]) -> List[np.ndarray]:
"""fp32 oracle of :class:`SecondFPN` (torch): NCHW in and out."""
import torch
with torch.no_grad():
return [_apply_torch(s, torch.as_tensor(np.asarray(x, dtype=np.float32))).numpy().copy()
for s, x in zip(deblocks, blocks_nchw)]