File size: 2,821 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
53
54
55
56
57
58
59
60
61
62
63
// SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
// SPDX-License-Identifier: Apache-2.0
//
// fp32 TILE [KMAX, C] -> ROW_MAJOR rows of ONE bucket tensor [BSTEP * (b + 1), C] (data movement
// only), b = ceil(n / BSTEP) - 1 with n = HDR[2]: the host then reads back a single right-sized
// tensor (one D2H) instead of several fixed-size chunks.
// Unit u = (tile row tr, tile column tc): the tile's 32 rows x 32 columns are written straight from
// the tile buffer as 2 x 64 B segments per row (face rows) into bucket row 32*tr + r.
// Tile rows at or beyond n are skipped (never read back). Runtime args: s_addr, hdr_addr, nunits,
// NB bucket addresses, then the nunits unit ids.
#include <stdint.h>
#include "api/dataflow/dataflow_api.h"

void kernel_main() {
    const uint32_t s_addr = get_arg_val<uint32_t>(0);
    const uint32_t hdr_addr = get_arg_val<uint32_t>(1);
    const uint32_t nunits = get_arg_val<uint32_t>(2);

    constexpr uint32_t cb_scratch = get_compile_time_arg_val(0);
    constexpr uint32_t C = get_compile_time_arg_val(1);
    constexpr uint32_t KMAX = get_compile_time_arg_val(2);
    constexpr uint32_t BSTEP = get_compile_time_arg_val(3);
    constexpr uint32_t NB = KMAX / BSTEP;
    constexpr uint32_t TC = C / 32;
    constexpr auto s_args = TensorAccessorArgs<4>();
    constexpr auto h_args = TensorAccessorArgs<s_args.next_compile_time_args_offset()>();
    constexpr auto o_args = TensorAccessorArgs<h_args.next_compile_time_args_offset()>();
    const auto sacc = TensorAccessor(s_args, s_addr, 4096);
    const auto hacc = TensorAccessor(h_args, hdr_addr, (16 + 4 * KMAX) * 4);

    const uint32_t base = get_write_ptr(cb_scratch);
    const uint32_t hb = base;
    const uint32_t tbuf = base + 64;
    noc_async_read(hacc.get_noc_addr(0), hb, 64);
    noc_async_read_barrier();
    uint32_t n = reinterpret_cast<const uint32_t*>(hb)[2];
    if (n == 0) {
        return;
    }
    if (n > KMAX) {
        n = KMAX;
    }
    const uint32_t o_addr = get_arg_val<uint32_t>(3 + (n + BSTEP - 1) / BSTEP - 1);
    const auto oacc = TensorAccessor(o_args, o_addr, C * 4);
    for (uint32_t ui = 0; ui < nunits; ++ui) {
        const uint32_t u = get_arg_val<uint32_t>(3 + NB + ui);
        const uint32_t tr = u / TC, tc = u % TC;
        if (tr * 32 >= n) {
            continue;
        }
        noc_async_read(sacc.get_noc_addr(tr * TC + tc), tbuf, 4096);
        noc_async_read_barrier();
        const uint32_t r0 = tr * 32;
        for (uint32_t r = 0; r < 32; ++r) {
            const uint32_t src = tbuf + (((r >> 4) * 2) * 256 + (r & 15) * 16) * 4;
            const uint64_t dst = oacc.get_noc_addr(r0 + r) + tc * 128;
            noc_async_write(src, dst, 64);
            noc_async_write(src + 1024, dst + 64, 64);
        }
        noc_async_write_barrier();
    }
}