File size: 4,972 Bytes
3fd1a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
#include <rknn_api.h>
#include <rknn_matmul_api.h>

#include <algorithm>
#include <chrono>
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <numeric>
#include <stdexcept>
#include <string>
#include <vector>

namespace {

using Clock = std::chrono::steady_clock;

void Check(int status, const char * operation) {
    if (status != RKNN_SUCC) {
        throw std::runtime_error(
            std::string(operation) + " failed with status " + std::to_string(status));
    }
}

rknn_core_mask CoreMask(int core) {
    if (core == 0) return RKNN_NPU_CORE_0;
    if (core == 1) return RKNN_NPU_CORE_1;
    if (core == 2) return RKNN_NPU_CORE_2;
    throw std::invalid_argument("core must be in [0, 2]");
}

double Microseconds(Clock::time_point begin, Clock::time_point end) {
    return std::chrono::duration<double, std::micro>(end - begin).count();
}

void PrintAttribute(const char * name, const rknn_matmul_tensor_attr & value) {
    std::printf("%s bytes=%u dims=[", name, value.size);
    for (std::uint32_t index = 0; index < value.n_dims; ++index) {
        std::printf("%s%u", index == 0 ? "" : ",", value.dims[index]);
    }
    std::printf("] type=%d\n", static_cast<int>(value.type));
}

} // namespace

int main(int argc, char ** argv) try {
    if (argc != 6) {
        std::fprintf(stderr, "usage: %s TOKENS K N CORE ITERATIONS\n", argv[0]);
        return 2;
    }
    const int tokens = std::atoi(argv[1]);
    const int k = std::atoi(argv[2]);
    const int n = std::atoi(argv[3]);
    const int core = std::atoi(argv[4]);
    const int iterations = std::atoi(argv[5]);
    if (tokens < 1 || k < 1 || n < 1 || iterations < 1) {
        throw std::invalid_argument("tokens, dimensions, and iterations must be positive");
    }

    // One INT8 activation row is represented by two INT4 rows: high and low nibble.
    rknn_matmul_info info {};
    info.M = 2 * tokens;
    info.K = k;
    info.N = n;
    info.type = RKNN_INT4_MM_INT4_TO_INT16;
    info.B_layout = RKNN_MM_LAYOUT_NATIVE;
    info.AC_layout = RKNN_MM_LAYOUT_NATIVE;
    info.AC_quant_type = RKNN_QUANT_TYPE_PER_LAYER_SYM;
    info.B_quant_type = RKNN_QUANT_TYPE_PER_LAYER_SYM;
    info.iommu_domain_id = 0;

    rknn_matmul_ctx context = 0;
    rknn_matmul_io_attr io {};
    Check(rknn_matmul_create(&context, &info, &io), "create W4 matmul");
    try {
        Check(rknn_matmul_set_core_mask(context, CoreMask(core)), "set W4 core");
        PrintAttribute("A", io.A);
        PrintAttribute("B", io.B);
        PrintAttribute("C", io.C);
        auto * a = rknn_create_mem2(context, io.A.size, RKNN_FLAG_MEMORY_CACHEABLE);
        auto * b = rknn_create_mem2(context, io.B.size, RKNN_FLAG_MEMORY_CACHEABLE);
        auto * c = rknn_create_mem2(context, io.C.size, RKNN_FLAG_MEMORY_CACHEABLE);
        if (a == nullptr || b == nullptr || c == nullptr) {
            throw std::runtime_error("cannot allocate W4 DMA memory");
        }
        std::memset(a->virt_addr, 0x11, io.A.size);
        std::memset(b->virt_addr, 0x11, io.B.size);
        std::memset(c->virt_addr, 0, io.C.size);
        Check(rknn_mem_sync(context, a, RKNN_MEMORY_SYNC_TO_DEVICE), "sync A");
        Check(rknn_mem_sync(context, b, RKNN_MEMORY_SYNC_TO_DEVICE), "sync B");
        Check(rknn_matmul_set_io_mem(context, a, &io.A), "bind A");
        Check(rknn_matmul_set_io_mem(context, b, &io.B), "bind B");
        Check(rknn_matmul_set_io_mem(context, c, &io.C), "bind C");
        for (int index = 0; index < 20; ++index) {
            Check(rknn_matmul_run(context), "warm W4 matmul");
        }
        std::vector<double> samples;
        samples.reserve(iterations);
        for (int index = 0; index < iterations; ++index) {
            const auto begin = Clock::now();
            Check(rknn_matmul_run(context), "run W4 matmul");
            samples.push_back(Microseconds(begin, Clock::now()));
        }
        std::sort(samples.begin(), samples.end());
        const double mean = std::accumulate(samples.begin(), samples.end(), 0.0) /
            static_cast<double>(samples.size());
        const double p50 = samples[samples.size() / 2];
        const double p95 = samples[std::min(
            samples.size() - 1, static_cast<std::size_t>(samples.size() * 0.95))];
        std::printf(
            "RESULT tokens=%d M=%d K=%d N=%d core=%d iterations=%d "
            "mean_us=%.3f p50_us=%.3f p95_us=%.3f token_us=%.3f\n",
            tokens, info.M, k, n, core, iterations, mean, p50, p95,
            mean / static_cast<double>(tokens));
        rknn_destroy_mem(context, c);
        rknn_destroy_mem(context, b);
        rknn_destroy_mem(context, a);
        Check(rknn_matmul_destroy(context), "destroy W4 matmul");
        context = 0;
    } catch (...) {
        if (context != 0) rknn_matmul_destroy(context);
        throw;
    }
    return 0;
} catch (const std::exception & error) {
    std::fprintf(stderr, "error: %s\n", error.what());
    return 1;
}