Download code/tt_diffusion_planner/tt/kernels/smask_compute.cpp from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 2.92 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smask_compute.cpp
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tt/kernels/smask_compute.cpp
-
curl -L -o smask_compute.cpp https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tt/kernels/smask_compute.cpp
2.92 kB
| // SPDX-License-Identifier: Apache-2.0 | |
| // Attention score scale + mask (tt/smask_kernel.py), compute: out = s * scale + mask, the two stock binary_ng fp32 | |
| // programs `ttnn.multiply(s, scale)` (SFPU mul_binary_tile against a tile filled with the fp32 scale) and | |
| // `ttnn.add(., mask)`, in one pass. With an fp32 mask the stock add is the SFPU add_binary_tile<NearestEven>; with a | |
| // bf16 mask (trunc_tf32 = 1) binary_ng takes its FPU path, whose SrcA truncates the fp32 scores to TF32 (measured: | |
| // logs/diffusion-planner/opt_r2/smask_diag.log, the stock result == (exact & 0xFFFFE000) for every finite element), | |
| // so the scaled scores are truncated the same way (SFPU bitwise AND on the raw bits) before the exact add of the | |
| // 0 / -inf mask. Bit-identical to the two programs in both cases. | |
| // CT args: [0] H (even), [1] per-core RT-arg count (P3), [2] trunc_tf32. Per-core RT args: [n, 0]. | |
| using namespace ckernel; | |
| constexpr uint32_t H = get_compile_time_arg_val(0); | |
| constexpr uint32_t trunc_tf32 = get_compile_time_arg_val(2); | |
| constexpr uint32_t cb_s = 0, cb_m = 1, cb_c = 2, cb_out = 16; | |
| constexpr auto RNE = ckernel::DstRoundingMode::NearestEven; | |
| void kernel_main() { | |
| const uint32_t n = get_arg_val<uint32_t>(0); | |
| if (n == 0) { | |
| return; | |
| } | |
| compute_kernel_hw_startup(cb_s, cb_out); | |
| cb_wait_front(cb_c, 1); | |
| for (uint32_t t = 0; t < n; ++t) { | |
| cb_wait_front(cb_m, 1); | |
| cb_wait_front(cb_s, H); | |
| for (uint32_t h = 0; h < H; h += 2) { | |
| tile_regs_acquire(); | |
| copy_init(cb_s); | |
| copy_tile(cb_s, h, 0); | |
| copy_tile(cb_s, h + 1, 1); | |
| copy_init(cb_c); | |
| copy_tile(cb_c, 0, 2); | |
| mul_binary_tile_init(); | |
| mul_binary_tile(0, 2, 0); | |
| mul_binary_tile(1, 2, 1); | |
| if constexpr (trunc_tf32) { | |
| bitwise_and_tile_init(); | |
| bitwise_and_tile<DataFormat::Int32>(0, 0xFFFFE000u); | |
| bitwise_and_tile<DataFormat::Int32>(1, 0xFFFFE000u); | |
| } | |
| reconfig_data_format_srca(cb_c, cb_m); | |
| copy_init(cb_m); | |
| copy_tile(cb_m, 0, 3); | |
| reconfig_data_format_srca(cb_m, cb_s); | |
| add_binary_tile_init(); | |
| add_binary_tile<RNE>(0, 3, 0); | |
| add_binary_tile<RNE>(1, 3, 1); | |
| cb_reserve_back(cb_out, 2); | |
| tile_regs_commit(); | |
| tile_regs_wait(); | |
| pack_tile(0, cb_out); | |
| pack_tile(1, cb_out); | |
| tile_regs_release(); | |
| cb_push_back(cb_out, 2); | |
| } | |
| cb_pop_front(cb_s, H); | |
| cb_pop_front(cb_m, 1); | |
| } | |
| } | |