changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
13.4 kB
# 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