File size: 2,923 Bytes
be62f78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
// 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].
#include <cstdint>

#include "api/compute/common.h"
#include "api/compute/compute_kernel_api.h"
#include "api/compute/eltwise_binary_sfpu.h"
#include "api/compute/eltwise_unary/bitwise.h"
#include "api/compute/eltwise_unary/eltwise_unary.h"
#include "api/compute/tile_move_copy.h"

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);
    }
}