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