File size: 3,672 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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// Descriptor L2 normalisation + untilize in one op (models/tt/desc_norm.py). Per tile row of the
// height-sharded [rows, 256] TILE descriptor map: S = sum_c x^2 (x*x on the FPU, packed fp32, row-reduced
// with a ones scaler), r = rsqrt(S) (SFPU), y = x * r (column broadcast), pack-untilized into a
// [32 rows, 256] row-major block for the writer.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/compute_kernel_hw_startup.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"

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 WT = 8, HB = 4;                    // tiles per row (256 ch), tiles per DST block

    compute_kernel_hw_startup(cb_in, cb_in, cb_sq);
    cb_wait_front(cb_in, TR * WT);
    cb_wait_front(cb_one, 1);
    for (uint32_t r = 0; r < TR; ++r) {
        // ---- x^2 -> CB_SQ (fp32)
        reconfig_data_format(cb_in, cb_in);
        pack_reconfig_data_format(cb_sq);
        mul_init(cb_in, cb_in);
        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) {
                const uint32_t t = r * WT + b * HB + k;
                mul_tiles(cb_in, cb_in, t, t, 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) {
            reduce_tile<PoolType::SUM, ReduceDim::REDUCE_ROW>(cb_sq, cb_one, t, 0, 0);
        }
        rsqrt_tile_init();
        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_in, cb_rs);
        pack_reconfig_data_format(cb_out);
        mul_bcast_cols_init(cb_in, 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) {
                mul_tiles_bcast_cols(cb_in, cb_rs, r * WT + 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);
    }
}