Ling-3.0-tiny-RKNN / src /rknn_backend.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
3.26 kB
#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