File size: 10,472 Bytes
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6ffd3f8
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""2x2 / stride-2 max pool as ONE custom generic_op (no halo), for the 64-channel encoder blocks.

Input: L1 height-sharded ROW_MAJOR bf16 [1, 1, N*H*W, 64] whose shards hold whole, even numbers of
image rows (block 0: 4 rows of 640 px per core on 120 cores; block 1: 2 rows of 320). Output: L1
height-sharded TILE bf16 [1, 1, N*H/2*W/2, 64] on the same cores (shard = the pooled rows). The
compute kernel tilizes 64-pixel row chunks read as 32 "pixel pairs" x 128 channels (even/odd pixels
land in different tiles) and takes the SFPU max of the 4 tiles of each 2x2 window: exact.
"""
from __future__ import annotations

import os

import ttnn

_KDIR = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "kernels", "sp_pool")


def _src(name):
    with open(os.path.join(_KDIR, name)) as f:
        return f.read()


class Pool2x2:
    C = 64

    def __init__(self, device):
        self.device = device
        self._src = [_src(f"pool2x2_{k}.cpp") for k in ("reader", "compute", "writer")]

    def supports(self, x: ttnn.Tensor, h: int, w: int, c: int) -> bool:
        if c != self.C or w % 64 or h % 2 or not x.is_sharded() or x.layout != ttnn.ROW_MAJOR_LAYOUT:
            return False
        mc = x.memory_config()
        if mc.memory_layout != ttnn.TensorMemoryLayout.HEIGHT_SHARDED or mc.buffer_type != ttnn.BufferType.L1:
            return False
        rows = self._rows(x)
        n = x.shape[-2]
        ncores = mc.shard_spec.grid.num_cores()
        return x.shape[-1] == c and rows % (2 * w) == 0 and rows * ncores == n

    def _rows(self, x: ttnn.Tensor) -> int:
        """Pixels (C-channel rows) per shard; also right for a zero-copy view of a cell-conv
        output whose shard spec still says [rows / cell, cell * C]."""
        sh = x.memory_config().shard_spec.shape
        return sh[0] * sh[1] // self.C

    def __call__(self, x: ttnn.Tensor, h: int, w: int) -> ttnn.Tensor:
        mc = x.memory_config()
        grid = mc.shard_spec.grid
        rows = self._rows(x)
        rows_out_img = rows // (2 * w)  # output image rows per core
        ch = w // 64  # 64-pixel chunks per image row
        n_out = x.shape[-2] // 4
        out_rows = rows // 4
        omc = ttnn.MemoryConfig(
            ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1,
            ttnn.ShardSpec(grid, [out_rows, self.C], mc.shard_spec.orientation),
        )
        out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, n_out, self.C]), ttnn.bfloat16, ttnn.TILE_LAYOUT, self.device, omc)
        CB_IN, CB_T, CB_OUT = 0, 1, 2
        n_in_pages = rows * self.C * 2 // 8192
        n_out_pages = out_rows * self.C * 2 // 2048
        cb_in = ttnn.cb_descriptor_from_sharded_tensor(CB_IN, x)
        cb_in.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_IN, data_format=ttnn.bfloat16, page_size=8192)]
        cb_out = ttnn.cb_descriptor_from_sharded_tensor(CB_OUT, out)
        cb_out.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_OUT, data_format=ttnn.bfloat16, page_size=2048)]
        f = ttnn.CBFormatDescriptor(buffer_index=CB_T, data_format=ttnn.bfloat16, page_size=2048)
        cb_t = ttnn.CBDescriptor(total_size=2048 * 8 * ch, core_ranges=grid, format_descriptors=[f])
        SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE
        r, c, wr = self._src
        ks = [
            ttnn.KernelDescriptor(kernel_source=r, source_type=SC, core_ranges=grid, compile_time_args=[CB_IN, n_in_pages],
                                  runtime_args=[], config=ttnn.ReaderConfigDescriptor()),
            ttnn.KernelDescriptor(kernel_source=c, source_type=SC, core_ranges=grid,
                                  compile_time_args=[CB_IN, CB_T, CB_OUT, rows_out_img, ch], runtime_args=[],
                                  config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4)),
            ttnn.KernelDescriptor(kernel_source=wr, source_type=SC, core_ranges=grid, compile_time_args=[CB_OUT, n_out_pages],
                                  runtime_args=[], config=ttnn.WriterConfigDescriptor()),
        ]
        return ttnn.generic_op([x, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=[cb_in, cb_t, cb_out]))


class U8ToBf16:
    """uint8 image upload -> the bf16 cell input of the ``l1`` encoder, ONE data-movement generic_op
    (kernels/sp_input/u8_to_bf16.cpp): table lookup bf16(fp32(u) / 255) on both RISCs of every core,
    local L1 only. Bit-identical to the host fp32 /255 + bf16 cast (fused_host.U8_TO_BF16)."""

    def __init__(self, device):
        import torch

        from .fused_host import U8_TO_BF16

        self.device = device
        kdir = os.path.join(os.path.dirname(_KDIR), "sp_input")
        with open(os.path.join(kdir, "u8_to_bf16.cpp")) as f:
            self._src = f.read()
        bits = (U8_TO_BF16.view(torch.int16).to(torch.int32) & 0xFFFF).tolist()
        self._src = "#define SP_U8_LUT " + ",".join(f"0x{b:04x}" for b in bits) + "\n" + self._src

    def __call__(self, x: ttnn.Tensor, out_mc: ttnn.MemoryConfig, out_shape) -> ttnn.Tensor:
        out = ttnn.allocate_tensor_on_device(ttnn.Shape(list(out_shape)), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, self.device, out_mc)
        grid = out_mc.shard_spec.grid
        sh = out_mc.shard_spec.shape
        n_px = sh[0] * sh[1]
        xs = x.memory_config().shard_spec
        if xs.grid != grid or xs.shape[0] * xs.shape[1] != n_px or n_px % 8:
            raise ValueError("u8 input and bf16 cell input must split the pixels identically")
        CB_IN, CB_OUT = 0, 1
        cb_in = ttnn.cb_descriptor_from_sharded_tensor(CB_IN, x)
        cb_in.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_IN, data_format=ttnn.bfloat16, page_size=n_px)]
        cb_out = ttnn.cb_descriptor_from_sharded_tensor(CB_OUT, out)
        cb_out.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_OUT, data_format=ttnn.bfloat16, page_size=2 * n_px)]
        SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE
        ks = []
        for proc, cfg in ((0, ttnn.ReaderConfigDescriptor()), (1, ttnn.WriterConfigDescriptor())):
            ks.append(ttnn.KernelDescriptor(kernel_source=self._src, source_type=SC, core_ranges=grid,
                                            compile_time_args=[CB_IN, CB_OUT, n_px, proc], runtime_args=[], config=cfg))
        return ttnn.generic_op([x, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=[cb_in, cb_out]))


class Pool2x2Rows:
    """2x2 / stride-2 max pool for an activation sharded as ONE image row per core (block 2: 160 px x
    128 ch on 120 cores), ONE generic_op instead of halo + Move + max_pool2d
    (kernels/sp_pool/pool_rows_{reader,compute}.cpp). Input: L1 height-sharded ROW_MAJOR bf16
    [1, 1, H*W, C], shard [W, C] on H cores; output: L1 height-sharded ROW_MAJOR bf16
    [1, 1, H/2*W/2, C], shard [W/4, C] on the same cores (what ttnn.max_pool2d produces): output core
    c holds half c % 2 of output row c // 2. Each RISC gathers the even / odd pixels of one of the two
    input rows (NoC reads, one row is local), the compute kernel takes the SFPU max of the four
    2 KB pseudo tiles: exact."""

    def __init__(self, device):
        self.device = device
        self._r = _src("pool_rows_reader.cpp")
        self._c = _src("pool_rows_compute.cpp")

    def supports(self, x: ttnn.Tensor, h: int, w: int, c: int) -> bool:
        if not x.is_sharded() or x.layout != ttnn.ROW_MAJOR_LAYOUT or x.dtype != ttnn.bfloat16 or h % 2 or w % 4:
            return False
        mc = x.memory_config()
        if mc.memory_layout != ttnn.TensorMemoryLayout.HEIGHT_SHARDED or mc.buffer_type != ttnn.BufferType.L1:
            return False
        sh = mc.shard_spec.shape
        return (x.shape[-1] == c and sh[1] == c and sh[0] == w and mc.shard_spec.grid.num_cores() == h
                and (w // 4 * c * 2) % 2048 == 0 and 2048 % (c * 2) == 0 and x.shape[-2] == h * w)

    def __call__(self, x: ttnn.Tensor, h: int, w: int) -> ttnn.Tensor:
        mc = x.memory_config()
        grid = mc.shard_spec.grid
        c = x.shape[-1]
        npx = w // 4
        omc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1,
                                ttnn.ShardSpec(grid, [npx, c], mc.shard_spec.orientation))
        out = ttnn.allocate_tensor_on_device(ttnn.Shape([1, 1, h * w // 4, c]), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT, self.device, omc)
        cores = ttnn.corerange_to_cores(grid, row_wise=mc.shard_spec.orientation == ttnn.ShardOrientation.ROW_MAJOR)
        pages = npx * c * 2 // 2048
        CB_E0, CB_O0, CB_E1, CB_O1, CB_OUT = 0, 1, 2, 3, 4
        cbs = []
        for i in (CB_E0, CB_O0, CB_E1, CB_O1):
            f = ttnn.CBFormatDescriptor(buffer_index=i, data_format=ttnn.bfloat16, page_size=2048)
            cbs.append(ttnn.CBDescriptor(total_size=2048 * min(pages, 4), core_ranges=grid, format_descriptors=[f]))
        cb_out = ttnn.cb_descriptor_from_sharded_tensor(CB_OUT, out)
        cb_out.format_descriptors = [ttnn.CBFormatDescriptor(buffer_index=CB_OUT, data_format=ttnn.bfloat16, page_size=2048)]
        cbs.append(cb_out)
        addr = x.buffer_address()
        ks = []
        SC = ttnn.KernelDescriptor.SourceType.SOURCE_CODE
        for proc, cfg, ce, co in ((0, ttnn.ReaderConfigDescriptor(), CB_E0, CB_O0), (1, ttnn.WriterConfigDescriptor(), CB_E1, CB_O1)):
            rt = ttnn.RuntimeArgs()
            for i, core in enumerate(cores):
                hh = i % 2
                src = cores[i - hh + proc]
                nc = self.device.worker_core_from_logical_core(src)
                rt[core.x][core.y] = [nc.x, nc.y, addr, hh * (w // 2)]
            ks.append(ttnn.KernelDescriptor(kernel_source=self._r, source_type=SC, core_ranges=grid,
                                            compile_time_args=[ce, co, npx, c * 2], runtime_args=rt, config=cfg))
        ks.append(ttnn.KernelDescriptor(kernel_source=self._c, source_type=SC, core_ranges=grid,
                                        compile_time_args=[CB_E0, CB_O0, CB_E1, CB_O1, CB_OUT, pages], runtime_args=[],
                                        config=ttnn.ComputeConfigDescriptor(math_fidelity=ttnn.MathFidelity.HiFi4)))
        return ttnn.generic_op([x, out], ttnn.ProgramDescriptor(kernels=ks, semaphores=[], cbs=cbs))