moge-2-p150 / code /tt_moge /kernels /addln_reader.cpp
changh95's picture
Optimized build (2026-10-03): model call 19.0 ms, trace 17.7 ms
13b4736 verified
Raw History Blame Contribute Delete
4.17 kB
// 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;
}