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