File size: 3,259 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
#include "ling3/rknn_backend.h"

#include <cstdint>
#include <cstring>
#include <iostream>

#if LING3_WITH_RKNN
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#endif

namespace ling3 {

bool RknnBackendAvailable() noexcept {
#if LING3_WITH_RKNN
    return true;
#else
    return false;
#endif
}

int RunRknnMatmulSmoke(std::size_t iterations) {
#if !LING3_WITH_RKNN
    (void)iterations;
    std::cerr << "RKNN backend was disabled at build time\n";
    return 2;
#else
    if (iterations == 0) iterations = 1;

    rknn_matmul_info info {};
    info.M = 1;
    info.K = 32;
    info.N = 32;
    info.type = RKNN_INT8_MM_INT8_TO_INT32;
    info.B_layout = RKNN_MM_LAYOUT_NORM;
    info.AC_layout = RKNN_MM_LAYOUT_NORM;

    rknn_matmul_ctx context = 0;
    rknn_matmul_io_attr attributes {};
    int status = rknn_matmul_create(&context, &info, &attributes);
    if (status < 0) {
        std::cerr << "rknn_matmul_create failed: " << status << '\n';
        return 3;
    }

    rknn_tensor_mem * a = rknn_create_mem(context, attributes.A.size);
    rknn_tensor_mem * b = rknn_create_mem(context, attributes.B.size);
    rknn_tensor_mem * c = rknn_create_mem(context, attributes.C.size);
    if (a == nullptr || b == nullptr || c == nullptr) {
        std::cerr << "rknn_create_mem failed\n";
        if (a != nullptr) rknn_destroy_mem(context, a);
        if (b != nullptr) rknn_destroy_mem(context, b);
        if (c != nullptr) rknn_destroy_mem(context, c);
        rknn_matmul_destroy(context);
        return 4;
    }

    auto * a_values = static_cast<std::int8_t *>(a->virt_addr);
    auto * b_values = static_cast<std::int8_t *>(b->virt_addr);
    for (int k = 0; k < info.K; ++k) a_values[k] = static_cast<std::int8_t>((k % 7) - 3);
    for (int k = 0; k < info.K; ++k) {
        for (int n = 0; n < info.N; ++n) {
            b_values[k * info.N + n] = static_cast<std::int8_t>(((k + n) % 5) - 2);
        }
    }
    std::memset(c->virt_addr, 0, attributes.C.size);

    status = rknn_matmul_set_io_mem(context, a, &attributes.A);
    if (status == 0) status = rknn_matmul_set_io_mem(context, b, &attributes.B);
    if (status == 0) status = rknn_matmul_set_io_mem(context, c, &attributes.C);
    for (std::size_t index = 0; status == 0 && index < iterations; ++index) {
        status = rknn_matmul_run(context);
    }

    bool correct = status == 0;
    const auto * result = static_cast<const std::int32_t *>(c->virt_addr);
    for (int n = 0; correct && n < info.N; ++n) {
        std::int32_t expected = 0;
        for (int k = 0; k < info.K; ++k) {
            expected += static_cast<std::int32_t>(a_values[k]) *
                        static_cast<std::int32_t>(b_values[k * info.N + n]);
        }
        correct = result[n] == expected;
    }

    rknn_destroy_mem(context, a);
    rknn_destroy_mem(context, b);
    rknn_destroy_mem(context, c);
    rknn_matmul_destroy(context);

    if (status < 0) {
        std::cerr << "rknn_matmul_run failed: " << status << '\n';
        return 5;
    }
    if (!correct) {
        std::cerr << "RKNN matmul returned an incorrect result\n";
        return 6;
    }
    std::cout << "RKNN matmul smoke passed (" << iterations << " iterations)\n";
    return 0;
#endif
}

} // namespace ling3