#include #include #include #include #include #include #include #include #include #include #include #include #if !defined(__aarch64__) #error "gdn_cpu_benchmark requires AArch64 FP16 NEON" #endif #include #include #include 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 query, std::span key, std::span value, std::span factor, std::span<__fp16> state, std::span output) { const auto vector_offset = static_cast(head) * kDimension; const auto state_offset = static_cast(head) * kDimension * kDimension; std::array delta {}; for (int token = 0; token < tokens; ++token) { const auto token_offset = static_cast(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(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(v[column]) - Horizontal(vaddq_f32(sum0, sum1)); } const float beta = 0.5F; for (int column = 0; column < kDimension; ++column) { auto * row = s + static_cast(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 query, std::span key, std::span value, std::span factor, std::span<__fp16> state, std::span output) { std::array 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 left, std::span 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(left[index]) * right[index]; left_norm += static_cast(left[index]) * left[index]; right_norm += static_cast(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(tokens) * kHeads * kDimension; const std::size_t states = static_cast(kHeads) * kDimension * kDimension; std::vector<__fp16> query(vectors), key(vectors), value(vectors), factor(vectors); std::vector<__fp16> initial_state(states), state(states); std::vector output(vectors), first_output(vectors); for (std::size_t index = 0; index < vectors; ++index) { query[index] = static_cast<__fp16>( std::sin(static_cast(index + 1) * 0.013F) / std::sqrt(128.0F)); key[index] = static_cast<__fp16>( std::cos(static_cast(index + 3) * 0.017F) / std::sqrt(128.0F)); value[index] = static_cast<__fp16>(0.2F * std::sin(static_cast(index) * 0.021F)); factor[index] = static_cast<__fp16>(0.92F + 0.04F * std::cos(static_cast(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 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(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; }