File size: 7,300 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 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 | #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;
}
|