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;
}
|