File size: 4,169 Bytes
13b4736
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// SPDX-License-Identifier: Apache-2.0
// Fused residual add + LayerNorm on 2 cores per tile row (moge-2 opt kernel, round 7): NCRISC.
// 1. Reads this core's half row of x (HALF tiles) into c_0, RB tiles per NoC barrier (116 cores with 32 reads
//    each in flight saturate the NoC: 21 us vs 7 us for the same bytes). 
// 2. Two partial-statistic exchanges with the partner core of the same tile row: round k (1: mean, 2:
//    variance): wait for the compute's partial tile in c_18, copy it to slot 0 of the local pair CB (c_19 /
//    c_20) and to slot 1 of the partner's pair CB (same L1 address on both cores), increment the partner's
//    semaphore, wait until the own semaphore reaches k, push the pair (2 tiles). The semaphore is reset to 0
//    before exit (the partner has finished both increments by then).
// 3. Writes the odd BLK-blocks of LN(z) (c_27); the writer (BRISC) writes the even ones.
// Runtime args: x_addr, out_addr, row, col0, partner_noc_x, partner_noc_y.
// Compile args: HALF, WT, BLK, SEM_ID, TensorAccessorArgs x, out.
#include "api/dataflow/dataflow_api.h"
#include "api/dataflow/noc.h"
#include "api/dataflow/dataflow_buffer.h"
#include "api/tensor/noc_traits.h"

#ifdef ADDLN_PROF
#include "tools/profiler/kernel_profiler.hpp"
#define PZONE(n) DeviceZoneScopedN(n)
#else
#define PZONE(n)
#endif
#ifndef RB
#define RB 2
#endif

void kernel_main() {
    const uint32_t x_addr = get_arg_val<uint32_t>(0);
    const uint32_t o_addr = get_arg_val<uint32_t>(1);
    const uint32_t row = get_arg_val<uint32_t>(2);
    const uint32_t col0 = get_arg_val<uint32_t>(3);
    const uint32_t px = get_arg_val<uint32_t>(4);
    const uint32_t py = get_arg_val<uint32_t>(5);
    constexpr uint32_t HALF = get_compile_time_arg_val(0);
    constexpr uint32_t WT = get_compile_time_arg_val(1);
    constexpr uint32_t BLK = get_compile_time_arg_val(2);
    constexpr uint32_t SEM_ID = get_compile_time_arg_val(3);
    constexpr auto x_args = TensorAccessorArgs<4>();
    constexpr auto o_args = TensorAccessorArgs<x_args.next_compile_time_args_offset()>();
    constexpr uint32_t cb_x = 0, cb_scaler = 2, cb_send = 18, cb_pair1 = 19, cb_pair2 = 20, cb_out2 = 27;
    constexpr uint32_t PAGE = 2048;       // bf16 tile
    constexpr uint32_t FPAGE = 4096;      // fp32 tile

    const auto sx = TensorAccessor(x_args, x_addr);
    const auto so = TensorAccessor(o_args, o_addr);
    DataflowBuffer dx(cb_x);
    const uint32_t base = row * WT + col0;
    {
        PZONE("RREAD");
        for (uint32_t b = 0; b < HALF; b += RB) {
            dx.reserve_back(RB);
            const uint32_t lx = dx.get_write_ptr();
            for (uint32_t j = 0; j < RB; ++j) {
                noc_async_read(sx.get_noc_addr(base + b + j), lx + j * PAGE, PAGE);
            }
            noc_async_read_barrier();
            dx.push_back(RB);
        }
    }
    const uint32_t sem_addr = get_semaphore(SEM_ID);
    volatile tt_l1_ptr uint32_t* sem_ptr = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(sem_addr);
    DataflowBuffer send(cb_send);
    for (uint32_t k = 1; k <= 2; ++k) {
        DataflowBuffer pair(k == 1 ? cb_pair1 : cb_pair2);
        {
            PZONE("RWAIT");
            send.wait_front(1);
        }
        PZONE("REXCH");
        pair.reserve_back(2);
        const uint32_t src = send.get_read_ptr();
        const uint32_t dst = pair.get_write_ptr();
        noc_async_write(src, get_noc_addr(dst), FPAGE);
        noc_async_write(src, get_noc_addr(px, py, dst + FPAGE), FPAGE);
        noc_async_write_barrier();
        send.pop_front(1);
        noc_semaphore_inc(get_noc_addr(px, py, sem_addr), 1);
        noc_semaphore_wait_min(sem_ptr, k);
        pair.push_back(2);
    }

    DataflowBuffer dout(cb_out2);
    for (uint32_t b = BLK; b < HALF; b += 2 * BLK) {
        dout.wait_front(BLK);
        const uint32_t l = dout.get_read_ptr();
        for (uint32_t j = 0; j < BLK; ++j) {
            noc_async_write(l + j * PAGE, so.get_noc_addr(base + b + j), PAGE);
        }
        noc_async_writes_flushed();
        dout.pop_front(BLK);
    }
    noc_async_write_barrier();
    noc_async_atomic_barrier();
    *sem_ptr = 0;
}