#include #include #include #include #include #include #include #include #include #include #include #include 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(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(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 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(samples.size()); const double p50 = samples[samples.size() / 2]; const double p95 = samples[std::min( samples.size() - 1, static_cast(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(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; }