Download code/tt_diffusion_planner/ttaw/models/resnet.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 13.4 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/models/resnet.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/models/resnet.py
-
curl -L -o resnet.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/models/resnet.py
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 | |
| 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 | |