File size: 8,642 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
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
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Descriptor head in one op (models/tt/desc_head.py): 1x1 conv (256 -> 256) + bias, then the L2 norm and
// untilize of DescNormRM (dn_compute.cpp), per tile row of the height-sharded [rows, 256] TILE input.
//  1) z_f32 = x @ W, fp32 DST, K in ttnn order (k = 0..7), at HiFi2 (the fidelity of the ttnn 1x1 conv;
//     the kernel itself is compiled at HiFi4 for the norm steps) -> CB_P (fp32)
//  2) z = bf16(z_f32 + bias) (row broadcast; ttnn's fused-bias epilogue) -> CB_Z (bf16, the ttnn conv output)
//  3) S = sum_c z^2 (fp32), r = rsqrt(S), y = z * r, pack-untilized -> CB_OUT (exactly dn_compute.cpp)
// CB_W tile t = nb * 32 + k * 4 + j holds W tile (k, 4 nb + j); tiles 64 + n hold the bias (row 0) of
// output tile n.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/compute_kernel_hw_startup.h"
#include "api/compute/matmul.h"
#include "api/compute/pack.h"
#include "api/compute/eltwise_binary.h"
#include "api/compute/bcast.h"
#include "api/compute/reduce.h"
#include "api/compute/eltwise_unary/rsqrt.h"
#include "api/compute/pack_untilize.h"
#include "api/compute/reconfig_data_format.h"

#ifndef DH_SKIP
#define DH_SKIP 0
#endif
#ifndef DH_KIN
#define DH_KIN 8  // input tiles per tile row (16: merged head conv output, descriptor channels first)
#endif
#ifndef DH_NW
#define DH_NW 72
#endif
#ifndef MM_FID
#define MM_FID ckernel::MathFidelity::HiFi2
#endif

// matmul_block_init / matmul_block with an explicit math fidelity (no dynamic throttle)
ALWI void mmf_init(uint32_t in0, uint32_t in1, uint32_t ct) {
    state_configure(in1, in0, __builtin_LINE());
    UNPACK((llk_unpack_AB_matmul_init(in0, in1, 0, ct, 1, 1)));
    MATH((llk_math_matmul_init<MM_FID, MM_THROTTLE>(in0, in1, 0, ct, 1)));
}
ALWI void mmf_block(uint32_t in0, uint32_t in1, uint32_t i0, uint32_t i1, uint32_t d, uint32_t ct) {
    state_configure(in1, in0, __builtin_LINE());
    UNPACK((llk_unpack_AB_matmul(in0, in1, i0, i1, ct, 1, 1)));
    MATH((llk_math_matmul<MM_FID, MM_THROTTLE>(d, ct, 1)));
}

void kernel_main() {
    constexpr uint32_t cb_in = get_compile_time_arg_val(0);
    constexpr uint32_t cb_one = get_compile_time_arg_val(1);
    constexpr uint32_t cb_sq = get_compile_time_arg_val(2);
    constexpr uint32_t cb_rs = get_compile_time_arg_val(3);
    constexpr uint32_t cb_out = get_compile_time_arg_val(4);
    constexpr uint32_t TR = get_compile_time_arg_val(5);  // tile rows per core
    constexpr uint32_t cb_w = get_compile_time_arg_val(6);
    constexpr uint32_t cb_p = get_compile_time_arg_val(7);
    constexpr uint32_t cb_z = get_compile_time_arg_val(8);
    constexpr uint32_t WT = 8, HB = 4, KT = 8, NW = DH_NW, BIAS0 = 64, KIN = DH_KIN;

    compute_kernel_hw_startup<SrcOrder::Reverse>(cb_in, cb_w, cb_p);
    cb_wait_front(cb_in, TR * KIN);
    constexpr uint32_t PS = 0;
    cb_wait_front(cb_w, NW);
    cb_wait_front(cb_one, 1);
#ifdef DH_SCORE_CB
    // ---- merged head: score 1x1 first (s = bf16(x[:, 256:] @ Ws + bs), the ttnn 1x1 conv steps: fp32 DST over K
    // in order at MM_FID, bias row-broadcast on the fp32 partials); one tile row at a time packed in place into the
    // logits shard (CB bound to it) and pushed, so the writer can signal the softmax cores early
    {
        constexpr uint32_t cb_s = DH_SCORE_CB, cb_sp = DH_SP_CB, ST = 3, SW0 = 72, SB0 = 96;
        for (uint32_t r = 0; r < TR; ++r) {
            reconfig_data_format<SrcOrder::Reverse>(cb_in, cb_w);
            pack_reconfig_data_format(cb_sp);
            mmf_init(cb_in, cb_w, ST);
            cb_reserve_back(cb_sp, ST);
            tile_regs_acquire();
            for (uint32_t k = 0; k < KT; ++k) {
                mmf_block(cb_in, cb_w, r * KIN + KT + k, SW0 + k * ST, 0, ST);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t j = 0; j < ST; ++j) {
                pack_tile<true>(j, cb_sp, j);
            }
            tile_regs_release();
            cb_push_back(cb_sp, ST);
            cb_wait_front(cb_sp, ST);
            reconfig_data_format(cb_sp, cb_w);
            pack_reconfig_data_format(cb_s);
            add_bcast_rows_init(cb_sp, cb_w);
            cb_reserve_back(cb_s, ST);
            tile_regs_acquire();
            for (uint32_t j = 0; j < ST; ++j) {
                add_tiles_bcast_rows(cb_sp, cb_w, j, SB0 + j, j);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t j = 0; j < ST; ++j) {
                pack_tile(j, cb_s);
            }
            tile_regs_release();
            cb_push_back(cb_s, ST);
            cb_pop_front(cb_sp, ST);
        }
    }
#endif
    for (uint32_t r = 0; r < TR; ++r) {
        // ---- 1) matmul -> CB_P (fp32)
        reconfig_data_format<SrcOrder::Reverse>(cb_in, cb_w);
        pack_reconfig_data_format(cb_p);
        mmf_init(cb_in, cb_w, HB);
        cb_reserve_back(cb_p, WT + PS);
        for (uint32_t nb = 0; nb < WT / HB; ++nb) {
            tile_regs_acquire();
            for (uint32_t k = 0; k < KT; ++k) {
                if constexpr (!(DH_SKIP & 1)) mmf_block(cb_in, cb_w, r * KIN + k, nb * (KT * HB) + k * HB, 0, HB);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t j = 0; j < HB; ++j) {
                pack_tile<true>(j, cb_p, nb * HB + j);
            }
            tile_regs_release();
        }
        cb_push_back(cb_p, WT + PS);
        // ---- 2) + bias -> CB_Z (bf16)
        cb_wait_front(cb_p, WT + PS);
        reconfig_data_format(cb_p, cb_w);
        pack_reconfig_data_format(cb_z);
        add_bcast_rows_init(cb_p, cb_w);
        cb_reserve_back(cb_z, WT);
        for (uint32_t b = 0; b < WT / HB; ++b) {
            tile_regs_acquire();
            for (uint32_t k = 0; k < HB; ++k) {
                if constexpr (!(DH_SKIP & 2)) add_tiles_bcast_rows(cb_p, cb_w, b * HB + k, BIAS0 + b * HB + k, k);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t k = 0; k < HB; ++k) {
                pack_tile(k, cb_z);
            }
            tile_regs_release();
        }
        cb_push_back(cb_z, WT);
        cb_pop_front(cb_p, WT + PS);
        cb_wait_front(cb_z, WT);
        // ---- 3) x^2 -> CB_SQ (fp32)
        reconfig_data_format(cb_z, cb_z);
        pack_reconfig_data_format(cb_sq);
        mul_init(cb_z, cb_z);
        cb_reserve_back(cb_sq, WT);
        for (uint32_t b = 0; b < WT / HB; ++b) {
            tile_regs_acquire();
            for (uint32_t k = 0; k < HB; ++k) {
                if constexpr (!(DH_SKIP & 4)) mul_tiles(cb_z, cb_z, b * HB + k, b * HB + k, k);
            }
            tile_regs_commit();
            tile_regs_wait();
            for (uint32_t k = 0; k < HB; ++k) {
                pack_tile(k, cb_sq);
            }
            tile_regs_release();
        }
        cb_push_back(cb_sq, WT);
        // ---- row sum, rsqrt -> CB_RS (fp32, column 0)
        cb_wait_front(cb_sq, WT);
        reconfig_data_format(cb_one, cb_sq);
        pack_reconfig_data_format(cb_rs);
        reduce_init<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, cb_rs);
        cb_reserve_back(cb_rs, 1);
        tile_regs_acquire();
        for (uint32_t t = 0; t < WT; ++t) {
            if constexpr (!(DH_SKIP & 8)) reduce_tile<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, t, 0, 0);
        }
        rsqrt_tile_init();
        if constexpr (!(DH_SKIP & 16)) rsqrt_tile(0);
        tile_regs_commit();
        tile_regs_wait();
        pack_tile(0, cb_rs);
        tile_regs_release();
        reduce_uninit(cb_sq);
        cb_push_back(cb_rs, 1);
        cb_pop_front(cb_sq, WT);
        // ---- x * r (column broadcast), pack-untilize -> CB_OUT
        cb_wait_front(cb_rs, 1);
        reconfig_data_format(cb_z, cb_rs);
        pack_reconfig_data_format(cb_out);
        mul_bcast_cols_init(cb_z, cb_rs);
        pack_untilize_dest_init<HB, WT>(cb_out);
        cb_reserve_back(cb_out, WT);
        for (uint32_t b = 0; b < WT / HB; ++b) {
            tile_regs_acquire();
            for (uint32_t k = 0; k < HB; ++k) {
                if constexpr (!(DH_SKIP & 32)) mul_tiles_bcast_cols(cb_z, cb_rs, b * HB + k, 0, k);
            }
            tile_regs_commit();
            tile_regs_wait();
            pack_untilize_dest<HB, WT>(cb_out, 1, b);
            tile_regs_release();
        }
        pack_untilize_uninit(cb_out);
        cb_push_back(cb_out, WT);
        cb_pop_front(cb_rs, 1);
        cb_pop_front(cb_z, WT);
    }
}