Download code/tt_moge/kernels/addln_reader.cpp from changh95/moge-2-p150: direct link, hf CLI and curl.
- Browser
- Download file 4.17 kB
-
https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/kernels/addln_reader.cpp
- Command line
-
hf download hf://changh95/moge-2-p150/code/tt_moge/kernels/addln_reader.cpp
-
curl -L -o addln_reader.cpp https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/kernels/addln_reader.cpp
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. | |
| 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; | |
| } | |