Ling-3.0-tiny-RKNN / src /gdn_step.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
40 kB
#include "ling3/gdn_step.h"
#include "core_workers.h"
#include <algorithm>
#include <array>
#include <chrono>
#include <cmath>
#include <cstdlib>
#include <cstdint>
#include <cstring>
#include <fstream>
#include <memory>
#include <stdexcept>
#include <string>
#include <type_traits>
#include <utility>
#include <vector>
#if defined(__aarch64__)
#include <arm_neon.h>
#endif
#if LING3_WITH_RKNN
#include <rknn_api.h>
#endif
namespace ling3 {
namespace {
using Clock = std::chrono::steady_clock;
#if LING3_WITH_RKNN
constexpr int kHeads = 16;
constexpr int kHeadDimension = 128;
double Milliseconds(Clock::time_point begin, Clock::time_point end) {
return std::chrono::duration<double, std::milli>(end - begin).count();
}
void CheckRknn(int status, const char * operation) {
if (status != RKNN_SUCC) {
throw std::runtime_error(
std::string(operation) + " failed with RKNN status " + std::to_string(status));
}
}
std::size_t ElementCount(const rknn_tensor_attr & attribute) {
if (attribute.n_elems != 0) return attribute.n_elems;
std::size_t result = 1;
for (std::uint32_t index = 0; index < attribute.n_dims; ++index) {
result *= attribute.dims[index];
}
return result;
}
struct Lane {
int core = 0;
int head_offset = 0;
int heads = 0;
int tokens = 1;
rknn_context context = 0;
std::vector<rknn_tensor_attr> input_attributes;
std::vector<rknn_tensor_attr> output_attributes;
std::vector<rknn_tensor_mem *> input_memory;
std::vector<rknn_tensor_mem *> output_memory;
int query = -1;
int key = -1;
int value = -1;
int decay = -1;
int beta = -1;
int state = -1;
int output = -1;
int new_state = -1;
Lane() = default;
Lane(const Lane &) = delete;
Lane & operator=(const Lane &) = delete;
~Lane() {
if (context == 0) return;
for (auto * memory : input_memory) if (memory != nullptr) rknn_destroy_mem(context, memory);
for (auto * memory : output_memory) if (memory != nullptr) rknn_destroy_mem(context, memory);
rknn_destroy(context);
}
int Input(std::string_view name) const {
for (std::size_t index = 0; index < input_attributes.size(); ++index) {
if (name == input_attributes[index].name) return static_cast<int>(index);
}
throw std::runtime_error("GDN model is missing input " + std::string(name));
}
int Output(std::string_view name) const {
for (std::size_t index = 0; index < output_attributes.size(); ++index) {
if (name == output_attributes[index].name) return static_cast<int>(index);
}
throw std::runtime_error("GDN model is missing output " + std::string(name));
}
};
std::vector<std::byte> ReadModel(const std::string & path) {
std::ifstream stream(path, std::ios::binary | std::ios::ate);
if (!stream) throw std::runtime_error("cannot open GDN prefill model " + path);
const auto end = stream.tellg();
if (end <= 0) throw std::runtime_error("GDN prefill model is empty: " + path);
std::vector<std::byte> bytes(static_cast<std::size_t>(end));
stream.seekg(0);
stream.read(reinterpret_cast<char *>(bytes.data()), end);
if (!stream) throw std::runtime_error("cannot read GDN prefill model " + path);
return bytes;
}
void ConvertToFp16(std::span<const float> input, rknn_tensor_mem * memory) {
auto * output = static_cast<__fp16 *>(memory->virt_addr);
for (std::size_t index = 0; index < input.size(); ++index) {
output[index] = static_cast<__fp16>(input[index]);
}
}
void ConvertFromFp16(rknn_tensor_mem * memory, std::span<float> output) {
const auto * input = static_cast<const __fp16 *>(memory->virt_addr);
for (std::size_t index = 0; index < output.size(); ++index) {
output[index] = static_cast<float>(input[index]);
}
}
#if defined(__aarch64__)
void RunCpuHead(
int head,
std::size_t tokens,
const __fp16 * query,
const __fp16 * key,
const __fp16 * value,
const __fp16 * factor,
const float * beta,
__fp16 * state,
float * output) {
constexpr std::size_t global_width = kHeads * kHeadDimension;
const std::size_t vector_offset = static_cast<std::size_t>(head) * kHeadDimension;
const std::size_t state_offset =
static_cast<std::size_t>(head) * kHeadDimension * kHeadDimension;
std::array<float, kHeadDimension> delta {};
for (std::size_t token = 0; token < tokens; ++token) {
const std::size_t token_offset = token * global_width;
const auto * q = query + token_offset + vector_offset;
const auto * k = key + token_offset + vector_offset;
const auto * v = value + token_offset + vector_offset;
const auto * d = factor + token_offset + vector_offset;
auto * head_state = state + state_offset;
auto * token_output = output + token_offset + vector_offset;
for (int column = 0; column < kHeadDimension; ++column) {
auto * row = head_state + static_cast<std::size_t>(column) * kHeadDimension;
float32x4_t sum0 = vdupq_n_f32(0.0F);
float32x4_t sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kHeadDimension; 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]) - vaddvq_f32(vaddq_f32(sum0, sum1));
}
const float token_beta = beta[token * kHeads + head];
for (int column = 0; column < kHeadDimension; ++column) {
auto * row = head_state + static_cast<std::size_t>(column) * kHeadDimension;
const __fp16 update = static_cast<__fp16>(token_beta * delta[column]);
float32x4_t sum0 = vdupq_n_f32(0.0F);
float32x4_t sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kHeadDimension; 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)));
}
token_output[column] = vaddvq_f32(vaddq_f32(sum0, sum1));
}
}
}
template <typename Input>
void RunCpuHeadFp32State(
int head,
std::size_t tokens,
const Input * query,
const Input * key,
const Input * value,
const Input * factor,
const float * beta,
float * state,
float * output) {
constexpr std::size_t global_width = kHeads * kHeadDimension;
const std::size_t vector_offset = static_cast<std::size_t>(head) * kHeadDimension;
const std::size_t state_offset =
static_cast<std::size_t>(head) * kHeadDimension * kHeadDimension;
std::array<float, kHeadDimension> delta {};
std::array<float, kHeadDimension> query_fp32 {};
std::array<float, kHeadDimension> key_fp32 {};
std::array<float, kHeadDimension> value_fp32 {};
std::array<float, kHeadDimension> factor_fp32 {};
for (std::size_t token = 0; token < tokens; ++token) {
const std::size_t token_offset = token * global_width;
const auto * q = query + token_offset + vector_offset;
const auto * k = key + token_offset + vector_offset;
const auto * v = value + token_offset + vector_offset;
const auto * d = factor + token_offset + vector_offset;
auto * head_state = state + state_offset;
auto * token_output = output + token_offset + vector_offset;
if constexpr (std::is_same_v<Input, float>) {
std::copy_n(q, kHeadDimension, query_fp32.data());
std::copy_n(k, kHeadDimension, key_fp32.data());
std::copy_n(v, kHeadDimension, value_fp32.data());
std::copy_n(d, kHeadDimension, factor_fp32.data());
} else for (int index = 0; index < kHeadDimension; index += 8) {
const float16x8_t query8 = vld1q_f16(q + index);
const float16x8_t key8 = vld1q_f16(k + index);
const float16x8_t value8 = vld1q_f16(v + index);
const float16x8_t factor8 = vld1q_f16(d + index);
vst1q_f32(query_fp32.data() + index, vcvt_f32_f16(vget_low_f16(query8)));
vst1q_f32(
query_fp32.data() + index + 4, vcvt_f32_f16(vget_high_f16(query8)));
vst1q_f32(key_fp32.data() + index, vcvt_f32_f16(vget_low_f16(key8)));
vst1q_f32(key_fp32.data() + index + 4, vcvt_f32_f16(vget_high_f16(key8)));
vst1q_f32(value_fp32.data() + index, vcvt_f32_f16(vget_low_f16(value8)));
vst1q_f32(
value_fp32.data() + index + 4, vcvt_f32_f16(vget_high_f16(value8)));
vst1q_f32(factor_fp32.data() + index, vcvt_f32_f16(vget_low_f16(factor8)));
vst1q_f32(
factor_fp32.data() + index + 4, vcvt_f32_f16(vget_high_f16(factor8)));
}
for (int column = 0; column < kHeadDimension; ++column) {
auto * row = head_state + static_cast<std::size_t>(column) * kHeadDimension;
float32x4_t sum0 = vdupq_n_f32(0.0F);
float32x4_t sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kHeadDimension; index += 8) {
const float32x4_t decay0 = vld1q_f32(factor_fp32.data() + index);
const float32x4_t decay1 = vld1q_f32(factor_fp32.data() + index + 4);
const float32x4_t decayed0 = vmulq_f32(vld1q_f32(row + index), decay0);
const float32x4_t decayed1 = vmulq_f32(vld1q_f32(row + index + 4), decay1);
vst1q_f32(row + index, decayed0);
vst1q_f32(row + index + 4, decayed1);
sum0 = vfmaq_f32(sum0, decayed0, vld1q_f32(key_fp32.data() + index));
sum1 = vfmaq_f32(
sum1, decayed1, vld1q_f32(key_fp32.data() + index + 4));
}
delta[column] = value_fp32[column] - vaddvq_f32(vaddq_f32(sum0, sum1));
// This output column depends only on its own state row. Update it
// while resident in L1 rather than rereading the entire 64 KiB head.
const float token_beta = beta[token * kHeads + head];
const float update = token_beta * delta[column];
sum0 = vdupq_n_f32(0.0F);
sum1 = vdupq_n_f32(0.0F);
for (int index = 0; index < kHeadDimension; index += 8) {
const float32x4_t key0 = vld1q_f32(key_fp32.data() + index);
const float32x4_t key1 = vld1q_f32(key_fp32.data() + index + 4);
const float32x4_t updated0 =
vfmaq_n_f32(vld1q_f32(row + index), key0, update);
const float32x4_t updated1 =
vfmaq_n_f32(vld1q_f32(row + index + 4), key1, update);
vst1q_f32(row + index, updated0);
vst1q_f32(row + index + 4, updated1);
sum0 = vfmaq_f32(
sum0, updated0, vld1q_f32(query_fp32.data() + index));
sum1 = vfmaq_f32(
sum1, updated1, vld1q_f32(query_fp32.data() + index + 4));
}
token_output[column] = vaddvq_f32(vaddq_f32(sum0, sum1));
}
}
}
#endif
#endif
} // namespace
struct GdnStep::Impl {
#if LING3_WITH_RKNN
std::vector<std::unique_ptr<Lane>> lanes;
std::vector<std::unique_ptr<Lane>> batch16_lanes;
#if defined(__aarch64__)
std::vector<std::uint16_t> cpu_query;
std::vector<std::uint16_t> cpu_key;
std::vector<std::uint16_t> cpu_value;
std::vector<std::uint16_t> cpu_factor;
std::vector<std::uint16_t> cpu_state;
std::vector<float> cpu_state_fp32;
std::vector<float> cpu_query_fp32, cpu_factor_fp32;
bool cpu_state_valid = false;
bool cpu_state_fp32_valid = false;
bool device_state_dirty = false;
#endif
#endif
std::size_t total_state_bytes = 0;
Impl(std::span<const std::byte> heads6, std::span<const std::byte> heads5) {
#if LING3_WITH_RKNN
lanes.push_back(InitializeLane(0, 0, 6, 1, heads6));
lanes.push_back(InitializeLane(1, 6, 5, 1, heads5));
lanes.push_back(InitializeLane(2, 11, 5, 1, heads5));
for (const auto & lane : lanes) {
total_state_bytes += lane->input_memory[lane->state]->size;
}
if (const char * directory = std::getenv("LING3_GDN_PREFILL_DIR")) {
std::string prefix(directory);
if (!prefix.empty() && prefix.back() != '/') prefix.push_back('/');
const auto heads6_batch = ReadModel(
prefix + "heads6/ling3_gdn_prefill_t16_h6_fp16_rk3588.rknn");
const auto heads5_batch = ReadModel(
prefix + "heads5/ling3_gdn_prefill_t16_h5_fp16_rk3588.rknn");
batch16_lanes.push_back(InitializeLane(0, 0, 6, 16, heads6_batch));
batch16_lanes.push_back(InitializeLane(1, 6, 5, 16, heads5_batch));
batch16_lanes.push_back(InitializeLane(2, 11, 5, 16, heads5_batch));
}
Reset();
#else
(void)heads6;
(void)heads5;
throw std::runtime_error("GdnStep requires a build with RKNN support");
#endif
}
~Impl() = default;
#if LING3_WITH_RKNN
std::unique_ptr<Lane> InitializeLane(
int core,
int head_offset,
int heads,
int tokens,
std::span<const std::byte> model) {
if (model.empty()) throw std::invalid_argument("GDN RKNN model is empty");
auto lane = std::make_unique<Lane>();
lane->core = core;
lane->head_offset = head_offset;
lane->heads = heads;
lane->tokens = tokens;
CheckRknn(
rknn_init(
&lane->context,
const_cast<std::byte *>(model.data()),
static_cast<std::uint32_t>(model.size()),
0,
nullptr),
"rknn_init GDN");
CheckRknn(
rknn_set_core_mask(lane->context, static_cast<rknn_core_mask>(1U << core)),
"rknn_set_core_mask GDN");
rknn_input_output_num counts {};
CheckRknn(
rknn_query(lane->context, RKNN_QUERY_IN_OUT_NUM, &counts, sizeof(counts)),
"query GDN counts");
if (counts.n_input != 6 || counts.n_output != 2) {
throw std::runtime_error("GDN RKNN model has unexpected I/O counts");
}
lane->input_attributes.resize(counts.n_input);
lane->output_attributes.resize(counts.n_output);
lane->input_memory.resize(counts.n_input, nullptr);
lane->output_memory.resize(counts.n_output, nullptr);
for (std::uint32_t index = 0; index < counts.n_input; ++index) {
auto & attribute = lane->input_attributes[index];
attribute.index = index;
CheckRknn(
rknn_query(lane->context, RKNN_QUERY_INPUT_ATTR, &attribute, sizeof(attribute)),
"query GDN input");
if (attribute.type != RKNN_TENSOR_FLOAT16) {
throw std::runtime_error("GDN input is not FP16");
}
const auto bytes = std::max(attribute.size, attribute.size_with_stride);
lane->input_memory[index] = rknn_create_mem2(
lane->context, bytes, RKNN_FLAG_MEMORY_CACHEABLE);
if (lane->input_memory[index] == nullptr) {
throw std::runtime_error("cannot allocate GDN input memory");
}
auto binding = attribute;
binding.pass_through = 1;
CheckRknn(
rknn_set_io_mem(lane->context, lane->input_memory[index], &binding),
"bind GDN input");
}
for (std::uint32_t index = 0; index < counts.n_output; ++index) {
auto & attribute = lane->output_attributes[index];
attribute.index = index;
CheckRknn(
rknn_query(lane->context, RKNN_QUERY_OUTPUT_ATTR, &attribute, sizeof(attribute)),
"query GDN output");
if (attribute.type != RKNN_TENSOR_FLOAT16) {
throw std::runtime_error("GDN output is not FP16");
}
const auto bytes = std::max(attribute.size, attribute.size_with_stride);
lane->output_memory[index] = rknn_create_mem2(
lane->context, bytes, RKNN_FLAG_MEMORY_CACHEABLE);
if (lane->output_memory[index] == nullptr) {
throw std::runtime_error("cannot allocate GDN output memory");
}
auto binding = attribute;
binding.pass_through = 1;
CheckRknn(
rknn_set_io_mem(lane->context, lane->output_memory[index], &binding),
"bind GDN output");
}
lane->query = lane->Input("query");
lane->key = lane->Input("key");
lane->value = lane->Input("value");
lane->decay = lane->Input("decay");
lane->beta = lane->Input("beta");
lane->state = lane->Input("state");
lane->output = lane->Output("output");
lane->new_state = lane->Output("new_state");
const auto vector_elements = static_cast<std::size_t>(heads) * kHeadDimension;
const auto packed_vector_elements = static_cast<std::size_t>(tokens) * vector_elements;
const auto state_elements = vector_elements * kHeadDimension;
for (int index : {lane->query, lane->key, lane->value, lane->decay}) {
if (ElementCount(lane->input_attributes[index]) != packed_vector_elements) {
const auto & attribute = lane->input_attributes[index];
throw std::runtime_error(
"GDN vector input " + std::string(attribute.name) + " has " +
std::to_string(ElementCount(attribute)) + " elements; expected " +
std::to_string(packed_vector_elements));
}
}
if (ElementCount(lane->input_attributes[lane->beta]) !=
static_cast<std::size_t>(tokens * heads) ||
ElementCount(lane->input_attributes[lane->state]) != state_elements ||
ElementCount(lane->output_attributes[lane->output]) != packed_vector_elements ||
ElementCount(lane->output_attributes[lane->new_state]) != state_elements) {
throw std::runtime_error("GDN state or output shape is incompatible");
}
return lane;
}
void Reset() {
for (auto * collection : {&lanes, &batch16_lanes}) {
for (const auto & lane : *collection) {
auto * state = lane->input_memory[lane->state];
std::memset(state->virt_addr, 0, state->size);
CheckRknn(
rknn_mem_sync(lane->context, state, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync reset GDN state");
}
}
#if defined(__aarch64__)
cpu_state.assign(
static_cast<std::size_t>(kHeads) * kHeadDimension * kHeadDimension, 0);
cpu_state_fp32.clear();
cpu_state_valid = true;
cpu_state_fp32_valid = false;
device_state_dirty = false;
#endif
}
#if defined(__aarch64__)
void CaptureDeviceState() {
cpu_state.resize(
static_cast<std::size_t>(kHeads) * kHeadDimension * kHeadDimension);
auto * state16 = reinterpret_cast<__fp16 *>(cpu_state.data());
for (const auto & lane : lanes) {
const auto * state_memory = lane->input_memory[lane->state];
const std::size_t offset =
static_cast<std::size_t>(lane->head_offset) * kHeadDimension * kHeadDimension;
std::memcpy(state16 + offset, state_memory->virt_addr, state_memory->size);
}
cpu_state_fp32.clear();
cpu_state_valid = true;
cpu_state_fp32_valid = false;
device_state_dirty = false;
}
void CommitCpuStateToDevice() {
if (!device_state_dirty) return;
const auto * state16 = reinterpret_cast<const __fp16 *>(cpu_state.data());
for (const auto & lane : lanes) {
auto * state_memory = lane->input_memory[lane->state];
const std::size_t offset =
static_cast<std::size_t>(lane->head_offset) * kHeadDimension * kHeadDimension;
std::memcpy(state_memory->virt_addr, state16 + offset, state_memory->size);
CheckRknn(
rknn_mem_sync(lane->context, state_memory, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync CPU-prefill GDN state to device");
}
device_state_dirty = false;
}
#endif
void StageLane(
Lane & lane,
std::span<const float> query_values,
std::span<const float> key_values,
std::span<const float> value_values,
std::span<const float> decay_values,
std::span<const float> beta_values) {
const auto begin = static_cast<std::size_t>(lane.head_offset) * kHeadDimension;
const auto count = static_cast<std::size_t>(lane.heads) * kHeadDimension;
const auto beta_begin = static_cast<std::size_t>(lane.head_offset);
const auto stage = [&lane](int index, std::span<const float> values) {
auto * memory = lane.input_memory[index];
ConvertToFp16(values, memory);
CheckRknn(
rknn_mem_sync(lane.context, memory, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync GDN input");
};
stage(lane.query, query_values.subspan(begin, count));
stage(lane.key, key_values.subspan(begin, count));
stage(lane.value, value_values.subspan(begin, count));
stage(lane.decay, decay_values.subspan(begin, count));
stage(lane.beta, beta_values.subspan(beta_begin, lane.heads));
}
void StageBatchLane(
Lane & lane,
std::span<const float> query_values,
std::span<const float> key_values,
std::span<const float> value_values,
std::span<const float> decay_values,
std::span<const float> beta_values) {
constexpr std::size_t global_width = kHeads * kHeadDimension;
const auto head_begin = static_cast<std::size_t>(lane.head_offset) * kHeadDimension;
const auto lane_width = static_cast<std::size_t>(lane.heads) * kHeadDimension;
const auto stage_vectors = [&](int index, std::span<const float> values) {
auto * memory = lane.input_memory[index];
auto * destination = static_cast<__fp16 *>(memory->virt_addr);
for (int token = 0; token < lane.tokens; ++token) {
const auto source = values.subspan(
static_cast<std::size_t>(token) * global_width + head_begin, lane_width);
for (std::size_t element = 0; element < lane_width; ++element) {
destination[static_cast<std::size_t>(token) * lane_width + element] =
static_cast<__fp16>(source[element]);
}
}
CheckRknn(
rknn_mem_sync(lane.context, memory, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync GDN batch vector input");
};
stage_vectors(lane.query, query_values);
stage_vectors(lane.key, key_values);
stage_vectors(lane.value, value_values);
stage_vectors(lane.decay, decay_values);
auto * beta_memory = lane.input_memory[lane.beta];
auto * beta_destination = static_cast<__fp16 *>(beta_memory->virt_addr);
for (int token = 0; token < lane.tokens; ++token) {
const auto source = beta_values.subspan(
static_cast<std::size_t>(token) * kHeads + lane.head_offset, lane.heads);
for (int head = 0; head < lane.heads; ++head) {
beta_destination[static_cast<std::size_t>(token) * lane.heads + head] =
static_cast<__fp16>(source[head]);
}
}
CheckRknn(
rknn_mem_sync(lane.context, beta_memory, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync GDN batch beta input");
}
void RunWorkers(std::vector<std::unique_ptr<Lane>> & active_lanes) {
constexpr std::array<int, 3> cores {0, 1, 2};
CoreWorkers::Instance().Run(cores, [&active_lanes](int core) {
CheckRknn(rknn_run(active_lanes[core]->context, nullptr), "rknn_run GDN");
});
}
void CollectLane(Lane & lane, std::span<float> output_values) {
auto * output = lane.output_memory[lane.output];
auto * new_state = lane.output_memory[lane.new_state];
CheckRknn(
rknn_mem_sync(lane.context, output, RKNN_MEMORY_SYNC_FROM_DEVICE),
"sync GDN output");
CheckRknn(
rknn_mem_sync(lane.context, new_state, RKNN_MEMORY_SYNC_FROM_DEVICE),
"sync GDN new state");
const auto begin = static_cast<std::size_t>(lane.head_offset) * kHeadDimension;
const auto count = static_cast<std::size_t>(lane.heads) * kHeadDimension;
ConvertFromFp16(output, output_values.subspan(begin, count));
auto * state = lane.input_memory[lane.state];
if (state->size != new_state->size) {
throw std::runtime_error("GDN state buffers have different sizes");
}
std::memcpy(state->virt_addr, new_state->virt_addr, state->size);
CheckRknn(
rknn_mem_sync(lane.context, state, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync persistent GDN state");
}
void CopyState(Lane & source, Lane & destination) {
auto * source_state = source.input_memory[source.state];
auto * destination_state = destination.input_memory[destination.state];
if (source_state->size != destination_state->size) {
throw std::runtime_error("GDN single and batch state buffers differ");
}
std::memcpy(destination_state->virt_addr, source_state->virt_addr, source_state->size);
CheckRknn(
rknn_mem_sync(
destination.context, destination_state, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync copied GDN state");
}
void CollectBatchLane(Lane & lane, Lane & single_lane, std::span<float> output_values) {
auto * output = lane.output_memory[lane.output];
auto * new_state = lane.output_memory[lane.new_state];
CheckRknn(
rknn_mem_sync(lane.context, output, RKNN_MEMORY_SYNC_FROM_DEVICE),
"sync GDN batch output");
CheckRknn(
rknn_mem_sync(lane.context, new_state, RKNN_MEMORY_SYNC_FROM_DEVICE),
"sync GDN batch new state");
constexpr std::size_t global_width = kHeads * kHeadDimension;
const auto head_begin = static_cast<std::size_t>(lane.head_offset) * kHeadDimension;
const auto lane_width = static_cast<std::size_t>(lane.heads) * kHeadDimension;
const auto * source = static_cast<const __fp16 *>(output->virt_addr);
for (int token = 0; token < lane.tokens; ++token) {
auto destination = output_values.subspan(
static_cast<std::size_t>(token) * global_width + head_begin, lane_width);
for (std::size_t element = 0; element < lane_width; ++element) {
destination[element] = static_cast<float>(
source[static_cast<std::size_t>(token) * lane_width + element]);
}
}
auto * batch_state = lane.input_memory[lane.state];
auto * single_state = single_lane.input_memory[single_lane.state];
if (batch_state->size != new_state->size || single_state->size != new_state->size) {
throw std::runtime_error("GDN batch state buffers have different sizes");
}
for (const auto [context, destination] :
{std::pair {lane.context, batch_state}, std::pair {single_lane.context, single_state}}) {
std::memcpy(destination->virt_addr, new_state->virt_addr, new_state->size);
CheckRknn(
rknn_mem_sync(context, destination, RKNN_MEMORY_SYNC_TO_DEVICE),
"sync committed GDN batch state");
}
}
#endif
GdnRunTimings Run(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
constexpr auto vector_elements = static_cast<std::size_t>(16 * 128);
if (query.size() != vector_elements || key.size() != vector_elements ||
value.size() != vector_elements || decay.size() != vector_elements ||
beta.size() != 16 || output.size() != vector_elements) {
throw std::invalid_argument("GDN step received an incompatible tensor size");
}
#if !LING3_WITH_RKNN
(void)query;
(void)key;
(void)value;
(void)decay;
(void)beta;
(void)output;
throw std::runtime_error("GdnStep requires a build with RKNN support");
#else
if (std::getenv("LING3_GDN_CPU_DECODE") != nullptr) {
return RunBatchCpu(query, key, value, decay, beta, output);
}
CommitCpuStateToDevice();
const auto begin = Clock::now();
for (const auto & lane : lanes) {
StageLane(*lane, query, key, value, decay, beta);
}
const auto staged = Clock::now();
RunWorkers(lanes);
const auto executed = Clock::now();
for (const auto & lane : lanes) CollectLane(*lane, output);
CaptureDeviceState();
const auto end = Clock::now();
return {
Milliseconds(begin, staged),
Milliseconds(staged, executed),
Milliseconds(executed, end),
Milliseconds(begin, end),
};
#endif
}
GdnRunTimings RunBatch16(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
constexpr std::size_t vector_elements = 16 * 16 * 128;
if (query.size() != vector_elements || key.size() != vector_elements ||
value.size() != vector_elements || decay.size() != vector_elements ||
beta.size() != 16 * 16 || output.size() != vector_elements) {
throw std::invalid_argument("GDN batch16 received an incompatible tensor size");
}
#if !LING3_WITH_RKNN
(void)query;
(void)key;
(void)value;
(void)decay;
(void)beta;
(void)output;
throw std::runtime_error("GdnStep requires a build with RKNN support");
#else
if (batch16_lanes.size() != 3) {
throw std::runtime_error("GDN batch16 models are not configured");
}
CommitCpuStateToDevice();
const auto begin = Clock::now();
for (std::size_t index = 0; index < batch16_lanes.size(); ++index) {
CopyState(*lanes[index], *batch16_lanes[index]);
StageBatchLane(
*batch16_lanes[index], query, key, value, decay, beta);
}
const auto staged = Clock::now();
RunWorkers(batch16_lanes);
const auto executed = Clock::now();
for (std::size_t index = 0; index < batch16_lanes.size(); ++index) {
CollectBatchLane(*batch16_lanes[index], *lanes[index], output);
}
CaptureDeviceState();
const auto end = Clock::now();
return {
Milliseconds(begin, staged),
Milliseconds(staged, executed),
Milliseconds(executed, end),
Milliseconds(begin, end),
};
#endif
}
GdnRunTimings RunBatchCpu(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
#if !LING3_WITH_RKNN || !defined(__aarch64__)
(void)query;
(void)key;
(void)value;
(void)decay;
(void)beta;
(void)output;
throw std::runtime_error("CPU GDN prefill requires an AArch64 RKNN build");
#else
constexpr std::size_t width = kHeads * kHeadDimension;
if (query.empty() || query.size() % width != 0 || key.size() != query.size() ||
value.size() != query.size() || decay.size() != query.size() ||
beta.size() != query.size() / kHeadDimension || output.size() != query.size()) {
throw std::invalid_argument("CPU GDN batch received incompatible tensor sizes");
}
const std::size_t tokens = query.size() / width;
const auto begin = Clock::now();
const bool use_fp32_state = std::getenv("LING3_GDN_CPU_FP32_STATE") != nullptr;
const bool full_fp32 = std::getenv("LING3_GDN_FULL_FP32") != nullptr;
if (full_fp32 && !use_fp32_state)
throw std::invalid_argument("LING3_GDN_FULL_FP32 requires LING3_GDN_CPU_FP32_STATE");
if (full_fp32) {
cpu_query_fp32.resize(query.size());
cpu_factor_fp32.resize(decay.size());
}
cpu_query.resize(query.size());
cpu_key.resize(key.size());
cpu_value.resize(value.size());
cpu_factor.resize(decay.size());
const std::size_t state_elements =
static_cast<std::size_t>(kHeads) * kHeadDimension * kHeadDimension;
if (!cpu_state_valid || cpu_state.size() != state_elements) {
throw std::runtime_error("CPU GDN shadow state is unavailable");
}
auto * q16 = reinterpret_cast<__fp16 *>(cpu_query.data());
auto * k16 = reinterpret_cast<__fp16 *>(cpu_key.data());
auto * v16 = reinterpret_cast<__fp16 *>(cpu_value.data());
auto * d16 = reinterpret_cast<__fp16 *>(cpu_factor.data());
constexpr std::array<int, 4> cores {0, 1, 2, 3};
constexpr float query_scale = 1.0F / std::sqrt(128.0F);
CoreWorkers::Instance().Run(cores, [&](int worker) {
const std::size_t first = query.size() * static_cast<std::size_t>(worker) / 4;
const std::size_t last = query.size() * static_cast<std::size_t>(worker + 1) / 4;
for (std::size_t index = first; index < last; ++index) {
if (full_fp32) {
cpu_query_fp32[index] = query[index] * query_scale;
cpu_factor_fp32[index] = std::exp(decay[index]);
continue;
}
q16[index] = static_cast<__fp16>(query[index] * query_scale);
k16[index] = static_cast<__fp16>(key[index]);
v16[index] = static_cast<__fp16>(value[index]);
d16[index] = static_cast<__fp16>(std::exp(decay[index]));
}
});
auto * state16 = reinterpret_cast<__fp16 *>(cpu_state.data());
if (use_fp32_state && !cpu_state_fp32_valid) {
cpu_state_fp32.resize(state_elements);
for (std::size_t index = 0; index < state_elements; ++index) {
cpu_state_fp32[index] = static_cast<float>(state16[index]);
}
}
const auto staged = Clock::now();
CoreWorkers::Instance().Run(cores, [&](int worker) {
for (int head = worker; head < kHeads; head += 4) {
if (use_fp32_state) {
if (full_fp32) RunCpuHeadFp32State(
head, tokens, cpu_query_fp32.data(), key.data(), value.data(),
cpu_factor_fp32.data(), beta.data(), cpu_state_fp32.data(), output.data());
else RunCpuHeadFp32State(
head, tokens, q16, k16, v16, d16, beta.data(),
cpu_state_fp32.data(), output.data());
} else {
RunCpuHead(
head, tokens, q16, k16, v16, d16, beta.data(), state16, output.data());
}
}
});
const auto executed = Clock::now();
if (use_fp32_state) {
for (std::size_t index = 0; index < cpu_state_fp32.size(); ++index) {
state16[index] = static_cast<__fp16>(cpu_state_fp32[index]);
}
cpu_state_fp32_valid = true;
} else {
cpu_state_fp32_valid = false;
}
device_state_dirty = true;
const auto end = Clock::now();
return {
Milliseconds(begin, staged),
Milliseconds(staged, executed),
Milliseconds(executed, end),
Milliseconds(begin, end),
};
#endif
}
};
GdnStep::GdnStep(
std::span<const std::byte> heads6_model,
std::span<const std::byte> heads5_model)
: impl_(std::make_unique<Impl>(heads6_model, heads5_model)) {}
GdnStep::~GdnStep() = default;
GdnState GdnStep::SaveState() {
#if LING3_WITH_RKNN && defined(__aarch64__)
if (!impl_->cpu_state_valid) impl_->CaptureDeviceState();
return {impl_->cpu_state, impl_->cpu_state_fp32, impl_->cpu_state_fp32_valid};
#else
throw std::runtime_error("GDN checkpoints require the RKNN aarch64 backend");
#endif
}
void GdnStep::RestoreState(const GdnState & state) {
#if LING3_WITH_RKNN && defined(__aarch64__)
constexpr std::size_t count = kHeads * kHeadDimension * kHeadDimension;
if (state.fp16.size() != count || (state.fp32_valid && state.fp32.size() != count))
throw std::invalid_argument("incompatible GDN checkpoint");
impl_->cpu_state = state.fp16;
impl_->cpu_state_fp32 = state.fp32;
impl_->cpu_state_valid = true;
impl_->cpu_state_fp32_valid = state.fp32_valid;
impl_->device_state_dirty = true;
#else
(void)state;
throw std::runtime_error("GDN checkpoints require the RKNN aarch64 backend");
#endif
}
GdnStep::GdnStep(GdnStep &&) noexcept = default;
GdnStep & GdnStep::operator=(GdnStep &&) noexcept = default;
void GdnStep::Reset() {
#if LING3_WITH_RKNN
impl_->Reset();
#else
throw std::runtime_error("GdnStep requires a build with RKNN support");
#endif
}
GdnRunTimings GdnStep::Run(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
return impl_->Run(query, key, value, decay, beta, output);
}
GdnRunTimings GdnStep::RunBatch16(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
return impl_->RunBatch16(query, key, value, decay, beta, output);
}
GdnRunTimings GdnStep::RunBatchCpu(
std::span<const float> query,
std::span<const float> key,
std::span<const float> value,
std::span<const float> decay,
std::span<const float> beta,
std::span<float> output) {
return impl_->RunBatchCpu(query, key, value, decay, beta, output);
}
bool GdnStep::has_batch16() const noexcept {
#if LING3_WITH_RKNN
return impl_->batch16_lanes.size() == 3;
#else
return false;
#endif
}
std::size_t GdnStep::state_bytes() const noexcept { return impl_->total_state_bytes; }
} // namespace ling3