# SPDX-License-Identifier: Apache-2.0 """C27: ResNet builders on device feature maps, cameras as batch (PLAN.md section 1.2 row C27; first built by the BEVDet port for its ResNet-50 image backbone and CustomResNet BEV encoder; BEVFormer and METEOR reuse it). - :class:`Bottleneck`: ResNet-50/101, pytorch style (the stride sits on the 3x3 ``conv2``; 1x1 projection ``downsample`` on the first block of a stage): ``relu(conv3(relu(conv2(relu(conv1(x))))) + idt)``; - :class:`BasicBlock`: ResNet-18/34 and BEVDet's CustomResNet (two 3x3 convs; the projection may be any conv, CustomResNet's is a 3x3 stride-2 one): ``relu(conv2(relu(conv1(x))) + idt)``; - :class:`ResNetStem`: ``maxpool 3x3 s2 p1 (relu(conv1 7x7 s2 p3))`` (the max pool runs DRAM-sliced); - :class:`ResNetStages`: ``len(layers)`` stages of one block kind; ``__call__(x, keep=...)`` returns the stage outputs a caller wants (FPN levels, taps) and frees the others as soon as the next stage consumed them. **Weights** come from a parameter provider ``provider(module) -> ConvParams`` (BatchNorm folded, conv attributes as the source graph defines them: BEVDet's reads ``OnnxWeights`` by consuming node, so no kernel size or stride is re-typed here). **Names** come from a ``names(stage, block)`` callable: :func:`mmdet_names` (``layer1.0.conv1``, ``downsample.0``) and :func:`custom_resnet_names` (``layers.0.0.conv1``, ``downsample``). **Convs** come from a ``factory(module, params, activation)`` (default :func:`conv_factory`: one C17 ``Conv2d`` per module, with a ``ttaw.precision`` spec or a ``PrecisionPolicy`` resolved per module name); a port that needs deformable convs (BEVFormer / METEOR DCN, PLAN.md C23 / C27 "DCN hook") returns its own callable for the modules it names. Placement: every activation is a DRAM-interleaved TILE ``[1, 1, N*H*W, C]`` map (C17's rule: DRAM inputs run ttnn's automatic DRAM slicing, 1x1 convs are matmuls placed back in DRAM), so the 6 x 256 x 704 cameras' 34.6 MB stem and layer1 maps fit; intermediates are freed as soon as they are consumed. Every op is traceable; a layer prepares its weights on its first (eager) call. :func:`resnet_numpy` is the fp32 torch oracle of the same blocks. """ from __future__ import annotations from dataclasses import dataclass from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union import numpy as np from ..ops.conv import CONV_PRECISION, Conv2d, FeatureMap, residual_add from ..precision import Precision, PrecisionPolicy __all__ = ["ConvParams", "DEFAULT_PRECISION", "conv_factory", "mmdet_names", "custom_resnet_names", "Bottleneck", "BasicBlock", "ResNetStem", "ResNetStages", "maxpool_out_hw", "resnet_numpy", "identity_like"] DEFAULT_PRECISION = CONV_PRECISION @dataclass(frozen=True) class ConvParams: """One BN-folded convolution as a provider returns it (numpy float32).""" weight: np.ndarray # [Cout, Cin / groups, kh, kw] bias: Optional[np.ndarray] = None # [Cout] stride: Tuple[int, int] = (1, 1) padding: Tuple[int, int, int, int] = (0, 0, 0, 0) # top, bottom, left, right dilation: Tuple[int, int] = (1, 1) groups: int = 1 Provider = Callable[[str], ConvParams] ConvFactory = Callable[[str, ConvParams, Optional[str]], Callable[[FeatureMap], FeatureMap]] Names = Callable[[int, int], Dict[str, str]] def conv_factory(precision: Union[str, Precision, PrecisionPolicy] = DEFAULT_PRECISION, **conv_kwargs: Any ) -> ConvFactory: """The standard factory: one C17 ``Conv2d`` per module. ``precision``: a spec / :class:`Precision` for every conv, or a :class:`PrecisionPolicy` resolved per module name; other keywords go to every ``Conv2d``.""" def make(module: str, p: ConvParams, activation: Optional[str]) -> Conv2d: prec = precision.resolve(module) if isinstance(precision, PrecisionPolicy) else precision return Conv2d(p.weight, p.bias, stride=p.stride, padding=p.padding, dilation=p.dilation, groups=p.groups, activation=activation, precision=prec, name=module, **conv_kwargs) return make def mmdet_names(prefix: str) -> Names: """mmdet / torchvision names: ``{prefix}.layer{stage + 1}.{block}.conv{1,2,3}`` and ``.downsample.0``.""" def names(stage: int, block: int) -> Dict[str, str]: p = f"{prefix}.layer{stage + 1}.{block}" return {"conv1": p + ".conv1", "conv2": p + ".conv2", "conv3": p + ".conv3", "downsample": p + ".downsample.0", "block": p} return names def custom_resnet_names(prefix: str) -> Names: """BEVDet CustomResNet names: ``{prefix}.layers.{stage}.{block}.conv{1,2}`` and ``.downsample``.""" def names(stage: int, block: int) -> Dict[str, str]: p = f"{prefix}.layers.{stage}.{block}" return {"conv1": p + ".conv1", "conv2": p + ".conv2", "downsample": p + ".downsample", "block": p} return names def identity_like(x: FeatureMap, like: FeatureMap) -> FeatureMap: """The identity shortcut ``x`` in the tensor shape of the residual branch output ``like`` (same rows and channels). ``ttnn.max_pool2d`` on a DRAM input (:class:`ResNetStem`) returns a 4-D ``[N, OH, OW, C]`` TILE map (``generic_pools.cpp`` ``pool2d_DRAM``) while a conv returns the flattened ``[1, 1, N*OH*OW, C]`` rows, and ``ttnn.add`` of the two raises ``Invalid subtile broadcast type``: the first block of a ResNet-18 / 34 stage 1 (no projection) adds the stem output itself (METEOR port, ttaw 0.17.2). A no-op when the shapes already agree.""" if tuple(x.tensor.shape) == tuple(like.tensor.shape): return x import ttnn return x.with_tensor(ttnn.reshape(x.tensor, tuple(like.tensor.shape))) class Bottleneck: """``relu(conv3(relu(conv2(relu(conv1(x))))) + idt)``; ``idt = downsample(x)`` when the block has a projection.""" def __init__(self, provider: Provider, names: Dict[str, str], *, has_downsample: bool, factory: ConvFactory): self.name = names["block"] self.conv1 = factory(names["conv1"], provider(names["conv1"]), "relu") self.conv2 = factory(names["conv2"], provider(names["conv2"]), "relu") self.conv3 = factory(names["conv3"], provider(names["conv3"]), None) self.downsample = (factory(names["downsample"], provider(names["downsample"]), None) if has_downsample else None) def __call__(self, x: FeatureMap) -> FeatureMap: y = self.conv1(x) z = self.conv2(y) y.deallocate() y = self.conv3(z) z.deallocate() if self.downsample is not None: idt = self.downsample(x) out = residual_add(y, idt, "relu") idt.deallocate() else: out = residual_add(y, identity_like(x, y), "relu") y.deallocate() return out def convs(self) -> List[Any]: return [c for c in (self.conv1, self.conv2, self.conv3, self.downsample) if c is not None] class BasicBlock: """``relu(conv2(relu(conv1(x))) + idt)``; ``idt = downsample(x)`` when the block has a projection.""" def __init__(self, provider: Provider, names: Dict[str, str], *, has_downsample: bool, factory: ConvFactory): self.name = names["block"] self.conv1 = factory(names["conv1"], provider(names["conv1"]), "relu") self.conv2 = factory(names["conv2"], provider(names["conv2"]), None) self.downsample = (factory(names["downsample"], provider(names["downsample"]), None) if has_downsample else None) def __call__(self, x: FeatureMap) -> FeatureMap: y = self.conv1(x) z = self.conv2(y) y.deallocate() if self.downsample is not None: idt = self.downsample(x) out = residual_add(z, idt, "relu") idt.deallocate() else: out = residual_add(z, identity_like(x, z), "relu") z.deallocate() return out def convs(self) -> List[Any]: return [c for c in (self.conv1, self.conv2, self.downsample) if c is not None] def maxpool_out_hw(h: int, w: int, k: int = 3, s: int = 2, p: int = 1) -> Tuple[int, int]: return (h + 2 * p - k) // s + 1, (w + 2 * p - k) // s + 1 class ResNetStem: """``maxpool 3x3 s2 p1 (relu(conv1))`` (``conv_module`` names the stem conv; ``pool=False`` skips the pool). The max pool reads the DRAM conv output with automatic DRAM slicing and writes a TILE map for the first block; on that DRAM path its tensor is 4-D ``[N, OH, OW, C]`` (not the flattened rows of a conv output): convs take it as is, and the blocks' identity shortcuts reshape it (:func:`identity_like`).""" def __init__(self, provider: Provider, conv_module: str, *, factory: Optional[ConvFactory] = None, pool: bool = True): factory = factory or conv_factory() self.conv1 = factory(conv_module, provider(conv_module), "relu") self.pool = bool(pool) def __call__(self, x: FeatureMap) -> FeatureMap: import ttnn y = self.conv1(x) if not self.pool: return y out = ttnn.max_pool2d(y.tensor, batch_size=y.batch, input_h=y.height, input_w=y.width, channels=y.channels, kernel_size=[3, 3], stride=[2, 2], padding=[1, 1], dilation=[1, 1], ceil_mode=False, memory_config=ttnn.DRAM_MEMORY_CONFIG, output_layout=ttnn.TILE_LAYOUT) y.deallocate() oh, ow = maxpool_out_hw(y.height, y.width) return FeatureMap(out, y.batch, oh, ow, y.channels) def convs(self) -> List[Any]: return [self.conv1] class ResNetStages: """``len(layers)`` stages of ``block`` (``"bottleneck"`` or ``"basic"``); block 0 of stage s has the projection when ``downsample_first[s]`` (default: every stage). ``__call__(x, keep)`` returns the outputs of the stages in ``keep`` (others are freed once consumed; ``x`` itself is freed only with ``free_input=True``).""" def __init__(self, provider: Provider, names: Names, layers: Sequence[int], *, block: str = "bottleneck", factory: Optional[ConvFactory] = None, downsample_first: Sequence[bool] = ()): if block not in ("bottleneck", "basic"): raise ValueError("block must be 'bottleneck' or 'basic'") factory = factory or conv_factory() cls = Bottleneck if block == "bottleneck" else BasicBlock ds = list(downsample_first) or [True] * len(layers) if len(ds) != len(layers): raise ValueError("downsample_first needs one entry per stage") self.block = block self.stages: List[List[Any]] = [] for s, n in enumerate(layers): self.stages.append([cls(provider, names(s, b), has_downsample=(b == 0 and bool(ds[s])), factory=factory) for b in range(int(n))]) def __call__(self, x: FeatureMap, keep: Sequence[int] = (), *, free_input: bool = False) -> List[FeatureMap]: outs: List[FeatureMap] = [] cur = x for s, blocks in enumerate(self.stages): for blk in blocks: nxt = blk(cur) if (cur is not x or free_input) and not any(cur is o for o in outs): cur.deallocate() cur = nxt if s in keep: outs.append(cur) if cur is not x and not any(cur is o for o in outs): cur.deallocate() return outs def convs(self) -> List[Any]: return [c for blocks in self.stages for blk in blocks for c in blk.convs()] def resnet_numpy(x: np.ndarray, provider: Provider, names: Names, layers: Sequence[int], *, block: str = "bottleneck", stem: Optional[str] = None, keep: Sequence[int] = (), downsample_first: Sequence[bool] = ()) -> List[np.ndarray]: """fp32 torch oracle of :class:`ResNetStem` (if ``stem`` names its conv) + :class:`ResNetStages` on an NCHW array -> the kept stage outputs (NCHW float32).""" import torch import torch.nn.functional as F def conv(t, module, relu): p = provider(module) top, bottom, left, right = p.padding t = F.pad(t, (left, right, top, bottom)) y = F.conv2d(t, torch.from_numpy(np.asarray(p.weight, np.float32)), None if p.bias is None else torch.from_numpy(np.asarray(p.bias, np.float32)), stride=tuple(p.stride), dilation=tuple(p.dilation), groups=p.groups) return torch.relu(y) if relu else y ds = list(downsample_first) or [True] * len(layers) outs = [] with torch.no_grad(): t = torch.from_numpy(np.asarray(x, np.float32)) if stem is not None: t = F.max_pool2d(conv(t, stem, True), 3, 2, 1) for s, n in enumerate(layers): for b in range(int(n)): nm = names(s, b) idt = conv(t, nm["downsample"], False) if (b == 0 and ds[s]) else t y = conv(t, nm["conv1"], True) if block == "bottleneck": y = conv(conv(y, nm["conv2"], True), nm["conv3"], False) else: y = conv(y, nm["conv2"], False) t = torch.relu(y + idt) if s in keep: outs.append(t.numpy()) return outs