Download code/tt_moge/kernels/addln_compute.cpp from changh95/moge-2-p150: direct link, hf CLI and curl.
- Browser
- Download file 7.22 kB
-
https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/kernels/addln_compute.cpp
- Command line
-
hf download hf://changh95/moge-2-p150/code/tt_moge/kernels/addln_compute.cpp
-
curl -L -o addln_compute.cpp https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/kernels/addln_compute.cpp
7.22 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Fused residual add + LayerNorm on 2 cores per tile row (moge-2 opt kernel, round 7): compute. | |
| // Per core (HALF tiles of one tile row, the partner core holds the other half): | |
| // z = x + y (bf16, packed to c_16 for the LN and to c_17 for the writer) | |
| // partial mean = rowsum(z) / W -> c_18 -> exchange -> c_19 [own, partner] -> mean (fp32 SFPU add) | |
| // xmm = z - mean (fp32 c_22); partial var = rowsum(xmm^2) / W -> c_18 -> exchange -> c_20 -> var | |
| // rstd = rsqrt(var + eps); out = xmm * rstd (bf16 c_26) | |
| // Mirrors ttnn's interleaved LayerNorm kernel (layernorm.cpp: two-pass variance, fp32 intermediates, | |
| // reduce scaler 1 + scale by 1/W, eps truncated to bf16, rsqrt non-legacy); the halves' partial sums are | |
| // added in fp32 on the SFPU (exact unpack to DEST). | |
| // Compile args: HALF, BLK, W, EPS_BITS (fp32 bits, bf16-truncated). | |
| namespace numeric = norm::kernel_util::compute::numeric; | |
| namespace policies = norm::kernel_util::compute::policies; | |
| void kernel_main() { | |
| constexpr uint32_t HALF = get_compile_time_arg_val(0); | |
| constexpr uint32_t BLK = get_compile_time_arg_val(1); | |
| constexpr uint32_t W = get_compile_time_arg_val(2); | |
| constexpr uint32_t EPS_BITS = get_compile_time_arg_val(3); | |
| constexpr uint32_t cb_x = 0, cb_y = 1, cb_scaler = 2, cb_z = 16, cb_zout = 17, cb_send = 18, cb_pair1 = 19, | |
| cb_pair2 = 20, cb_mean = 21, cb_xmm = 22, cb_xmm2 = 23, cb_rstd = 25, cb_out = 26, cb_out2 = 27; | |
| DataflowBuffer dx(cb_x), dy(cb_y), dscaler(cb_scaler), dz(cb_z), dzout(cb_zout), dsend(cb_send), | |
| dpair1(cb_pair1), dpair2(cb_pair2), dmean(cb_mean), dxmm(cb_xmm), dxmm2(cb_xmm2), drstd(cb_rstd), | |
| dout(cb_out); | |
| compute_kernel_hw_startup(cb_x, cb_y, cb_z); | |
| // mean = own + partner (fp32, SFPU) | |
| auto pair_sum = [&](DataflowBuffer& dpair, uint32_t cb_pair, uint32_t cb_res, bool add_eps) { | |
| dpair.wait_front(2); | |
| reconfig_data_format_srca(cb_pair); | |
| copy_tile_to_dst_init_short(cb_pair); | |
| tile_regs_acquire(); | |
| copy_tile(cb_pair, 0, 0); | |
| copy_tile(cb_pair, 1, 1); | |
| add_binary_tile_init(); | |
| add_binary_tile(0, 1, 0); | |
| if (add_eps) { | |
| binop_with_scalar_tile_init(); | |
| add_unary_tile(0, EPS_BITS); | |
| rsqrt_tile_init<false>(); | |
| rsqrt_tile<false>(0); | |
| } | |
| tile_regs_commit(); | |
| dpair.pop_front(2); | |
| DataflowBuffer dres(cb_res); | |
| dres.reserve_back(1); | |
| pack_reconfig_data_format(cb_res); | |
| tile_regs_wait(); | |
| pack_tile(0, cb_res); | |
| tile_regs_release(); | |
| dres.push_back(1); | |
| }; | |
| { | |
| PZONE("ZADD"); | |
| // ---- z = x + y | |
| reconfig_data_format(cb_x, cb_y); | |
| pack_reconfig_data_format(cb_z); | |
| add_init(cb_x, cb_y); | |
| for (uint32_t b = 0; b < HALF; b += BLK) { | |
| dx.wait_front(BLK); | |
| dy.wait_front(BLK); | |
| tile_regs_acquire(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| add_tiles(cb_x, cb_y, i, i, i); | |
| } | |
| tile_regs_commit(); | |
| dx.pop_front(BLK); | |
| dy.pop_front(BLK); | |
| dz.reserve_back(BLK); | |
| dzout.reserve_back(BLK); | |
| tile_regs_wait(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| pack_tile(i, cb_z); | |
| } | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| pack_tile(i, cb_zout); | |
| } | |
| tile_regs_release(); | |
| dz.push_back(BLK); | |
| dzout.push_back(BLK); | |
| } | |
| } | |
| { | |
| PZONE("ZMEAN"); | |
| // ---- partial mean -> exchange | |
| numeric::row_wise_mean<PoolType::SUM, ReduceDim::REDUCE_ROW, true, policies::FullBlockWithoutPopPolicy, | |
| policies::WaitAtEndPolicy::NO_WAIT>(dz, dscaler, dsend, W, HALF, BLK); | |
| } | |
| { | |
| PZONE("ZPAIR1"); | |
| pair_sum(dpair1, cb_pair1, cb_mean, false); | |
| } | |
| { | |
| PZONE("ZSUB"); | |
| // ---- xmm = z - mean | |
| reconfig_data_format(cb_z, cb_mean); | |
| pack_reconfig_data_format(cb_xmm); | |
| dmean.wait_front(1); | |
| sub_bcast_cols_init(cb_z, cb_mean); | |
| for (uint32_t b = 0; b < HALF; b += BLK) { | |
| tile_regs_acquire(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| sub_tiles_bcast_cols(cb_z, cb_mean, b + i, 0, i); | |
| } | |
| tile_regs_commit(); | |
| dxmm.reserve_back(BLK); | |
| tile_regs_wait(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| pack_tile(i, cb_xmm); | |
| } | |
| tile_regs_release(); | |
| dxmm.push_back(BLK); | |
| } | |
| dz.pop_front(HALF); | |
| dmean.pop_front(1); | |
| } | |
| { | |
| PZONE("ZSQ"); | |
| // ---- xmm^2 -> partial var -> exchange | |
| reconfig_data_format(cb_xmm, cb_xmm); | |
| pack_reconfig_data_format(cb_xmm2); | |
| mul_init(cb_xmm, cb_xmm); | |
| for (uint32_t b = 0; b < HALF; b += BLK) { | |
| dxmm.wait_front(b + BLK); | |
| tile_regs_acquire(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| mul_tiles(cb_xmm, cb_xmm, b + i, b + i, i); | |
| } | |
| tile_regs_commit(); | |
| dxmm2.reserve_back(BLK); | |
| tile_regs_wait(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| pack_tile(i, cb_xmm2); | |
| } | |
| tile_regs_release(); | |
| dxmm2.push_back(BLK); | |
| } | |
| } | |
| { | |
| PZONE("ZVAR"); | |
| numeric::row_wise_mean<PoolType::SUM, ReduceDim::REDUCE_ROW, true, policies::FullBlockWithPopPolicy, | |
| policies::WaitAtEndPolicy::NO_WAIT>(dxmm2, dscaler, dsend, W, HALF, BLK); | |
| } | |
| { | |
| PZONE("ZPAIR2"); | |
| // rstd = rsqrt(own + partner + eps) | |
| pair_sum(dpair2, cb_pair2, cb_rstd, true); | |
| } | |
| { | |
| PZONE("ZOUT"); | |
| // ---- out = xmm * rstd | |
| reconfig_data_format(cb_xmm, cb_rstd); | |
| pack_reconfig_data_format(cb_out); | |
| drstd.wait_front(1); | |
| mul_bcast_cols_init(cb_xmm, cb_rstd); | |
| for (uint32_t b = 0; b < HALF; b += BLK) { | |
| tile_regs_acquire(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| mul_tiles_bcast_cols(cb_xmm, cb_rstd, b + i, 0, i); | |
| } | |
| tile_regs_commit(); | |
| // even blocks -> c_26 (BRISC writes them), odd blocks -> c_27 (NCRISC) | |
| const uint32_t cbo = ((b / BLK) & 1) ? cb_out2 : cb_out; | |
| DataflowBuffer dcbo(cbo); | |
| dcbo.reserve_back(BLK); | |
| tile_regs_wait(); | |
| for (uint32_t i = 0; i < BLK; ++i) { | |
| pack_tile(i, cbo); | |
| } | |
| tile_regs_release(); | |
| dcbo.push_back(BLK); | |
| } | |
| } | |
| dxmm.pop_front(HALF); | |
| drstd.pop_front(1); | |
| dscaler.pop_front(1); | |
| } | |