moge-2-p150 / code /tt_moge /kernels /heads_local.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
3.85 kB
// SPDX-License-Identifier: Apache-2.0
// Head split / merge as a tile permutation with LOCAL reads (moge-2 opt kernel, round 7, ttnn.generic_op).
// The source is L1-interleaved: each core handles exactly the source pages that live in its own L1 bank and
// writes them straight from L1 to their destination page (no reader, no CB, no read round trip). Two RISCs
// per core (PART 0 / 1 on their own NoC) take alternate pages of the bank.
// MODE 0 (create heads): src [S, 3*H*2 tiles], page p = r*(6H) + c, c = which*2H + h*2 + j
// -> dst[which] [H, S, 64]: page (h*RT + r)*2 + j
// MODE 1 (concat heads): src [H, S, 64]: page p = (h*RT + r)*2 + j -> dst [S, H*64]: page r*2H + h*2 + j
// Runtime args: src_addr, dst0_addr, dst1_addr, dst2_addr, num_pages.
// Compile args: MODE, RT, H, PART, TensorAccessorArgs src, dst0, dst1, dst2.
#include "api/dataflow/dataflow_api.h"
#include "api/tensor/noc_traits.h"
void kernel_main() {
const uint32_t src_addr = get_arg_val<uint32_t>(0);
const uint32_t a0 = get_arg_val<uint32_t>(1);
const uint32_t a1 = get_arg_val<uint32_t>(2);
const uint32_t a2 = get_arg_val<uint32_t>(3);
const uint32_t num_pages = get_arg_val<uint32_t>(4);
constexpr uint32_t MODE = get_compile_time_arg_val(0);
constexpr uint32_t RT = get_compile_time_arg_val(1);
constexpr uint32_t H = get_compile_time_arg_val(2);
constexpr uint32_t PART = get_compile_time_arg_val(3);
constexpr uint32_t PAGE = 2048;
constexpr auto src_args = TensorAccessorArgs<4>();
constexpr auto d0_args = TensorAccessorArgs<src_args.next_compile_time_args_offset()>();
constexpr auto d1_args = TensorAccessorArgs<d0_args.next_compile_time_args_offset()>();
constexpr auto d2_args = TensorAccessorArgs<d1_args.next_compile_time_args_offset()>();
const auto s = TensorAccessor(src_args, src_addr);
const auto s0 = TensorAccessor(d0_args, a0);
const auto s1 = TensorAccessor(d1_args, a1);
const auto s2 = TensorAccessor(d2_args, a2);
const uint32_t mx = my_x[noc_index], my = my_y[noc_index];
auto is_local = [&](uint32_t p) -> bool {
const uint64_t a = s.get_noc_addr(p, 0, noc_index);
return NOC_UNICAST_ADDR_X(a) == mx && NOC_UNICAST_ADDR_Y(a) == my;
};
// find this core's bank: first local page b, bank stride nb = distance to the next local page
uint32_t b = 0xFFFFFFFF, nb = 0;
for (uint32_t p = 0; p < num_pages && p < 512; ++p) {
if (is_local(p)) {
if (b == 0xFFFFFFFF) {
b = p;
} else {
nb = p - b;
break;
}
}
}
if (b == 0xFFFFFFFF) {
return; // no source page in this core's bank
}
if (nb == 0) {
nb = num_pages; // only one page lives here
}
uint32_t k = 0;
for (uint32_t p = b; p < num_pages; p += nb, ++k) {
if ((k & 1) != PART) {
continue;
}
const uint32_t src_l1 = (uint32_t)(s.get_noc_addr(p, 0, noc_index) & NOC_LOCAL_ADDR_MASK);
if constexpr (MODE == 0) {
const uint32_t r = p / (6 * H), c = p % (6 * H);
const uint32_t which = c / (2 * H), hj = c % (2 * H);
const uint32_t dst = ((hj >> 1) * RT + r) * 2 + (hj & 1);
const uint64_t d = which == 0 ? s0.get_noc_addr(dst, 0, noc_index)
: which == 1 ? s1.get_noc_addr(dst, 0, noc_index)
: s2.get_noc_addr(dst, 0, noc_index);
noc_async_write(src_l1, d, PAGE);
} else {
const uint32_t j = p & 1, hr = p >> 1;
const uint32_t h = hr / RT, r = hr % RT;
noc_async_write(src_l1, s0.get_noc_addr(r * 2 * H + h * 2 + j, 0, noc_index), PAGE);
}
}
noc_async_write_barrier();
}