File size: 5,460 Bytes
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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SuperPoint descriptor sampling, gather step (data movement only, ttnn.generic_op).
// Unit u = (tap t, tile row tr, column half h): keypoints k = 32*tr .. 32*tr+31 (from HDR, see
// kp_compact.cpp: position, cell rows/columns; HDR[2] = n, tile rows >= n are skipped):
//   G_t[k, h*C/2 .. ] = D[cell_t(k), h*C/2 .. ]   (bf16 descriptor-map row half, TILE pages)
//   W_t[k, 0]         = WTAB[y_k, 4*x_k + t]      (fp32 tap weight, [KV, 1] TILE; written by h == 0)
// D: [N_CELLS, C] ROW_MAJOR bf16 interleaved. WTAB [H, 4W] fp32 (exact host products,
// postprocess.SampleTables.weight_table).
// The device then computes sum_t G_t * W_t in fp32 (IEEE-exact on Tensix). Pure copies -> exact.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

void kernel_main() {
    const uint32_t d_addr = get_arg_val<uint32_t>(0);
    const uint32_t hdr_addr = get_arg_val<uint32_t>(1);
    const uint32_t wtab_addr = get_arg_val<uint32_t>(2);
    const uint32_t nunits = get_arg_val<uint32_t>(3);

    constexpr uint32_t cb_scratch = get_compile_time_arg_val(0);
    constexpr uint32_t C = get_compile_time_arg_val(1);
    constexpr uint32_t KV = get_compile_time_arg_val(2);
    constexpr uint32_t W = get_compile_time_arg_val(3);
    constexpr uint32_t H = get_compile_time_arg_val(4);
    constexpr uint32_t CH = C / 2;   // channels per unit
    constexpr uint32_t TCH = CH / 32;  // tiles per unit
    constexpr uint32_t TC = C / 32;
    constexpr auto d_args = TensorAccessorArgs<5>();
    constexpr auto hdr_args = TensorAccessorArgs<d_args.next_compile_time_args_offset()>();
    constexpr auto wtab_args = TensorAccessorArgs<hdr_args.next_compile_time_args_offset()>();
    constexpr auto g_args = TensorAccessorArgs<wtab_args.next_compile_time_args_offset()>();
    constexpr auto w_args = TensorAccessorArgs<g_args.next_compile_time_args_offset()>();
    const auto dacc = TensorAccessor(d_args, d_addr, C * 2);
    const auto hacc = TensorAccessor(hdr_args, hdr_addr, (16 + 4 * KV) * 4);
    const auto wtacc = TensorAccessor(wtab_args, wtab_addr, W * 16);

    const uint32_t base = get_write_ptr(cb_scratch);
    const uint32_t hdr0 = base;                       // HDR[0..15] (64 B)
    const uint32_t kps = base + 64;                    // 32 HDR keypoint entries (512 B)
    const uint32_t wblk = kps + 512;                   // 32 x 64 B weight blocks
    const uint32_t rows = wblk + 32 * 64;              // 32 rows x CH bf16
    const uint32_t tiles = rows + 32 * CH * 2;         // TCH bf16 tiles
    const uint32_t wtile = tiles + TCH * 2048;         // one fp32 tile

    noc_async_read(hacc.get_noc_addr(0), hdr0, 64);
    noc_async_read_barrier();
    const uint32_t n = reinterpret_cast<const uint32_t*>(hdr0)[2];
    for (uint32_t ui = 0; ui < nunits; ++ui) {
        const uint32_t u = get_arg_val<uint32_t>(4 + 3 * ui);
        const uint32_t g_addr = get_arg_val<uint32_t>(5 + 3 * ui);
        const uint32_t w_addr = get_arg_val<uint32_t>(6 + 3 * ui);
        const uint32_t t = u & 3, h = (u >> 2) & 1, tr = u >> 3;
        if (tr * 32 >= n) {
            continue;
        }
        const auto gacc = TensorAccessor(g_args, g_addr, 2048);
        const auto wacc = TensorAccessor(w_args, w_addr, 4096);
        noc_async_read(hacc.get_noc_addr(0) + (16 + 128 * tr) * 4, kps, 512);
        noc_async_read_barrier();
        const uint32_t* kp = reinterpret_cast<const uint32_t*>(kps);
        const uint32_t xs = (t & 1) ? 0 : 16, ys = (t >> 1) ? 0 : 16;
        for (uint32_t r = 0; r < 32; ++r) {
            const uint32_t yx = kp[4 * r];
            const uint32_t y = yx >> 16, x = yx & 0xFFFF;
            const uint32_t cell = ((kp[4 * r + 2] >> ys) & 0xFFFF) + ((kp[4 * r + 3] >> xs) & 0xFFFF);
            noc_async_read(dacc.get_noc_addr(cell) + h * CH * 2, rows + r * CH * 2, CH * 2);
            if (h == 0) {
                noc_async_read(wtacc.get_noc_addr(y) + ((x * 16) & ~63u), wblk + r * 64, 64);
            }
        }
        noc_async_read_barrier();
        if (h == 0) {
            uint32_t* wt = reinterpret_cast<uint32_t*>(wtile);
            for (uint32_t r = 0; r < 32; ++r) {
                const uint32_t x = kp[4 * r] & 0xFFFF;
                const uint32_t v = reinterpret_cast<const uint32_t*>(wblk + r * 64 + ((x * 16) & 63))[t];
                uint32_t* f0 = wt + ((r >> 4) * 2) * 256 + (r & 15) * 16;
                f0[0] = v;  // column 0 is all the broadcast multiply reads
            }
            noc_async_write(wtile, wacc.get_noc_addr(tr), 4096);
        }
        for (uint32_t r = 0; r < 32; ++r) {
            const uint32_t* src = reinterpret_cast<const uint32_t*>(rows + r * CH * 2);
            const uint32_t fr = (r >> 4) * 2;
            const uint32_t ro = (r & 15) * 16;
            for (uint32_t j = 0; j < CH / 16; ++j) {
                uint32_t* dst = reinterpret_cast<uint32_t*>(tiles + (j >> 1) * 2048 + ((fr + (j & 1)) * 256 + ro) * 2);
                const uint32_t* s = src + j * 8;
                dst[0] = s[0]; dst[1] = s[1]; dst[2] = s[2]; dst[3] = s[3];
                dst[4] = s[4]; dst[5] = s[5]; dst[6] = s[6]; dst[7] = s[7];
            }
        }
        for (uint32_t tc = 0; tc < TCH; ++tc) {
            noc_async_write(tiles + tc * 2048, gacc.get_noc_addr(tr * TC + h * TCH + tc), 2048);
        }
        noc_async_write_barrier();
    }
}