Ling-3.0-tiny-RKNN / tools /gdn_cpu_benchmark.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
7.3 kB
#include <algorithm>
#include <array>
#include <chrono>
#include <cmath>
#include <cstddef>
#include <cstdio>
#include <cstdlib>
#include <numeric>
#include <span>
#include <stdexcept>
#include <thread>
#include <vector>
#if !defined(__aarch64__)
#error "gdn_cpu_benchmark requires AArch64 FP16 NEON"
#endif
#include <arm_neon.h>
#include <pthread.h>
#include <sched.h>
namespace {
constexpr int kHeads = 16;
constexpr int kDimension = 128;
constexpr int kWorkers = 4;
void PinWorker(int worker) {
cpu_set_t cores;
CPU_ZERO(&cores);
CPU_SET(4 + worker, &cores);
pthread_setaffinity_np(pthread_self(), sizeof(cores), &cores);
}
float Horizontal(float32x4_t value) {
return vaddvq_f32(value);
}
void RunHeadFp16(
int head,
int tokens,
std::span<const __fp16> query,
std::span<const __fp16> key,
std::span<const __fp16> value,
std::span<const __fp16> factor,
std::span<__fp16> state,
std::span<float> output) {
const auto vector_offset = static_cast<std::size_t>(head) * kDimension;
const auto state_offset = static_cast<std::size_t>(head) * kDimension * kDimension;
std::array<float, kDimension> delta {};
for (int token = 0; token < tokens; ++token) {
const auto token_offset = static_cast<std::size_t>(token) * kHeads * kDimension;
const auto * q = query.data() + token_offset + vector_offset;
const auto * k = key.data() + token_offset + vector_offset;
const auto * v = value.data() + token_offset + vector_offset;
const auto * d = factor.data() + token_offset + vector_offset;
auto * s = state.data() + state_offset;
auto * o = output.data() + token_offset + vector_offset;
for (int column = 0; column < kDimension; ++column) {
auto * row = s + static_cast<std::size_t>(column) * kDimension;
float32x4_t sum0 = vdupq_n_f32(0.0F);
float32x4_t sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kDimension; index += 8) {
const float16x8_t decayed = vmulq_f16(
vld1q_f16(row + index), vld1q_f16(d + index));
vst1q_f16(row + index, decayed);
const float16x8_t key8 = vld1q_f16(k + index);
sum0 = vfmaq_f32(
sum0, vcvt_f32_f16(vget_low_f16(decayed)),
vcvt_f32_f16(vget_low_f16(key8)));
sum1 = vfmaq_f32(
sum1, vcvt_f32_f16(vget_high_f16(decayed)),
vcvt_f32_f16(vget_high_f16(key8)));
}
delta[column] = static_cast<float>(v[column]) - Horizontal(vaddq_f32(sum0, sum1));
}
const float beta = 0.5F;
for (int column = 0; column < kDimension; ++column) {
auto * row = s + static_cast<std::size_t>(column) * kDimension;
const __fp16 update = static_cast<__fp16>(beta * delta[column]);
float32x4_t sum0 = vdupq_n_f32(0.0F);
float32x4_t sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kDimension; index += 8) {
const float16x8_t updated = vfmaq_n_f16(
vld1q_f16(row + index), vld1q_f16(k + index), update);
vst1q_f16(row + index, updated);
const float16x8_t query8 = vld1q_f16(q + index);
sum0 = vfmaq_f32(
sum0, vcvt_f32_f16(vget_low_f16(updated)),
vcvt_f32_f16(vget_low_f16(query8)));
sum1 = vfmaq_f32(
sum1, vcvt_f32_f16(vget_high_f16(updated)),
vcvt_f32_f16(vget_high_f16(query8)));
}
o[column] = Horizontal(vaddq_f32(sum0, sum1));
}
}
}
void RunFp16(
int tokens,
std::span<const __fp16> query,
std::span<const __fp16> key,
std::span<const __fp16> value,
std::span<const __fp16> factor,
std::span<__fp16> state,
std::span<float> output) {
std::array<std::thread, kWorkers> workers;
for (int worker = 0; worker < kWorkers; ++worker) {
workers[worker] = std::thread([=, &state, &output]() {
PinWorker(worker);
for (int head = worker; head < kHeads; head += kWorkers) {
RunHeadFp16(head, tokens, query, key, value, factor, state, output);
}
});
}
for (auto & worker : workers) worker.join();
}
double Compare(std::span<const float> left, std::span<const float> right) {
double dot = 0.0;
double left_norm = 0.0;
double right_norm = 0.0;
for (std::size_t index = 0; index < left.size(); ++index) {
dot += static_cast<double>(left[index]) * right[index];
left_norm += static_cast<double>(left[index]) * left[index];
right_norm += static_cast<double>(right[index]) * right[index];
}
return dot / std::sqrt(left_norm * right_norm);
}
} // namespace
int main(int argc, char ** argv) try {
const int tokens = argc > 1 ? std::atoi(argv[1]) : 128;
const int iterations = argc > 2 ? std::atoi(argv[2]) : 20;
if (tokens < 1 || iterations < 1) throw std::invalid_argument("invalid benchmark size");
const std::size_t vectors = static_cast<std::size_t>(tokens) * kHeads * kDimension;
const std::size_t states = static_cast<std::size_t>(kHeads) * kDimension * kDimension;
std::vector<__fp16> query(vectors), key(vectors), value(vectors), factor(vectors);
std::vector<__fp16> initial_state(states), state(states);
std::vector<float> output(vectors), first_output(vectors);
for (std::size_t index = 0; index < vectors; ++index) {
query[index] = static_cast<__fp16>(
std::sin(static_cast<float>(index + 1) * 0.013F) / std::sqrt(128.0F));
key[index] = static_cast<__fp16>(
std::cos(static_cast<float>(index + 3) * 0.017F) / std::sqrt(128.0F));
value[index] = static_cast<__fp16>(0.2F * std::sin(static_cast<float>(index) * 0.021F));
factor[index] = static_cast<__fp16>(0.92F + 0.04F *
std::cos(static_cast<float>(index) * 0.009F));
}
std::fill(initial_state.begin(), initial_state.end(), static_cast<__fp16>(0));
state = initial_state;
RunFp16(tokens, query, key, value, factor, state, first_output);
std::vector<double> samples;
samples.reserve(iterations);
for (int iteration = 0; iteration < iterations; ++iteration) {
state = initial_state;
const auto begin = std::chrono::steady_clock::now();
RunFp16(tokens, query, key, value, factor, state, output);
const auto end = std::chrono::steady_clock::now();
samples.push_back(std::chrono::duration<double, std::milli>(end - begin).count());
}
const double mean = std::accumulate(samples.begin(), samples.end(), 0.0) / samples.size();
std::printf(
"tokens=%d heads=%d workers=%d mean_ms=%.6f min_ms=%.6f max_ms=%.6f "
"tokens_per_second=%.3f repeat_cosine=%.9f\n",
tokens, kHeads, kWorkers, mean,
*std::min_element(samples.begin(), samples.end()),
*std::max_element(samples.begin(), samples.end()),
1000.0 * tokens / mean, Compare(first_output, output));
return 0;
} catch (const std::exception & error) {
std::fprintf(stderr, "error: %s\n", error.what());
return 1;
}