Download tools/rknn_w4_batch_probe.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 4.97 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/rknn_w4_batch_probe.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/tools/rknn_w4_batch_probe.cpp
-
curl -L -o rknn_w4_batch_probe.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/tools/rknn_w4_batch_probe.cpp
4.97 kB
| 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; | |
| } | |