File size: 2,732 Bytes
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// SP_RSZ5: the three TRISCs take 3 of the 5 shares of the core's resize work (resize_r8_fn.inc, prepended; plain scalar
// code, no FPU/SFPU). They start when the data-movement PROC 0 pushes the scratch CB (coefficients and source rows in
// L1, output shard address stored in word 8 of the flag area). RT args as resize_r8.cpp.
#include <cstdint>
#include "api/compute/compute_kernel_api.h"
#include "api/compute/common.h"
#include "api/compute/cb_api.h"

void kernel_main() {
    const uint32_t W = get_arg_val<uint32_t>(2);
    const uint32_t kv = get_arg_val<uint32_t>(4);
    const uint32_t vt_off = get_arg_val<uint32_t>(5);
    const uint32_t Y0 = get_arg_val<uint32_t>(6);
    const uint32_t r0 = get_arg_val<uint32_t>(7);
    const uint32_t nrows = get_arg_val<uint32_t>(8);
    constexpr uint32_t cb = get_compile_time_arg_val(0);
    constexpr uint32_t OUT_W = get_compile_time_arg_val(2);
    constexpr uint32_t OUT_ROWS = get_compile_time_arg_val(3);
    constexpr uint32_t KHMAX = get_compile_time_arg_val(4);
    constexpr uint32_t KVMAX = get_compile_time_arg_val(5);
    constexpr uint32_t NRMAX = get_compile_time_arg_val(6);
    constexpr uint32_t RBMAX = get_compile_time_arg_val(7);
    constexpr uint32_t PAGE = 256;
    constexpr uint32_t CORE_ROWS = 2 * OUT_ROWS;
    constexpr uint32_t HT_BYTES = ((OUT_W * (KHMAX + 1) * 4 + PAGE - 1) / PAGE) * PAGE;
    constexpr uint32_t VT_BYTES = ((CORE_ROWS * (KVMAX + 1) * 4 + 2 * PAGE - 1) / PAGE + 1) * PAGE;
    constexpr uint32_t TMP_BYTES = (NRMAX + KVMAX + 1) * OUT_W;
#if defined(TRISC_UNPACK)
    constexpr uint32_t w = 2;
#elif defined(TRISC_MATH)
    constexpr uint32_t w = 3;
#else
    constexpr uint32_t w = 4;
#endif
    cb_wait_front(cb, 1);
    const uint32_t base = get_tile_address(cb, 0);
    const uint32_t vt_l1 = base + HT_BYTES;  // PROC 0's copy of the vertical entries
    const uint32_t rows_l1 = base + HT_BYTES + 2 * VT_BYTES;
    uint8_t* tmp = reinterpret_cast<uint8_t*>(rows_l1 + RBMAX);
    volatile uint32_t* flags = reinterpret_cast<volatile uint32_t*>(rows_l1 + RBMAX + TMP_BYTES);
    asm volatile("fence" ::: "memory");
    const uint32_t out_l1 = flags[8];  // output shard address, stored by PROC 0 before its push
    const uint32_t rb = ((W + 63) & ~63u) + 64;
    const uint32_t vw0 = vt_off + Y0 * (kv + 1);
    const int32_t* vt = reinterpret_cast<const int32_t*>(vt_l1) + (vw0 - (vw0 / 64) * 64);
    rsz5_work<RSZ_KH, RSZ_KV, OUT_W, CORE_ROWS>(w, reinterpret_cast<const int32_t*>(base), vt, kv + 1,
                                                 reinterpret_cast<const uint8_t*>(rows_l1), rb, nrows, tmp, r0, out_l1);
}