#include "ling3/gdn_step.h" #include "core_workers.h" #include #include #include #include #include #include #include #include #include #include #include #include #include #include #if defined(__aarch64__) #include #endif #if LING3_WITH_RKNN #include #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(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 input_attributes; std::vector output_attributes; std::vector input_memory; std::vector 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(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(index); } throw std::runtime_error("GDN model is missing output " + std::string(name)); } }; std::vector 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 bytes(static_cast(end)); stream.seekg(0); stream.read(reinterpret_cast(bytes.data()), end); if (!stream) throw std::runtime_error("cannot read GDN prefill model " + path); return bytes; } void ConvertToFp16(std::span 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 output) { const auto * input = static_cast(memory->virt_addr); for (std::size_t index = 0; index < output.size(); ++index) { output[index] = static_cast(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(head) * kHeadDimension; const std::size_t state_offset = static_cast(head) * kHeadDimension * kHeadDimension; std::array 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(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(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(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 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(head) * kHeadDimension; const std::size_t state_offset = static_cast(head) * kHeadDimension * kHeadDimension; std::array delta {}; std::array query_fp32 {}; std::array key_fp32 {}; std::array value_fp32 {}; std::array 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) { 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(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> lanes; std::vector> batch16_lanes; #if defined(__aarch64__) std::vector cpu_query; std::vector cpu_key; std::vector cpu_value; std::vector cpu_factor; std::vector cpu_state; std::vector cpu_state_fp32; std::vector 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 heads6, std::span 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 InitializeLane( int core, int head_offset, int heads, int tokens, std::span model) { if (model.empty()) throw std::invalid_argument("GDN RKNN model is empty"); auto lane = std::make_unique(); lane->core = core; lane->head_offset = head_offset; lane->heads = heads; lane->tokens = tokens; CheckRknn( rknn_init( &lane->context, const_cast(model.data()), static_cast(model.size()), 0, nullptr), "rknn_init GDN"); CheckRknn( rknn_set_core_mask(lane->context, static_cast(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(heads) * kHeadDimension; const auto packed_vector_elements = static_cast(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(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(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(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(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(cpu_state.data()); for (const auto & lane : lanes) { auto * state_memory = lane->input_memory[lane->state]; const std::size_t offset = static_cast(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 query_values, std::span key_values, std::span value_values, std::span decay_values, std::span beta_values) { const auto begin = static_cast(lane.head_offset) * kHeadDimension; const auto count = static_cast(lane.heads) * kHeadDimension; const auto beta_begin = static_cast(lane.head_offset); const auto stage = [&lane](int index, std::span 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 query_values, std::span key_values, std::span value_values, std::span decay_values, std::span beta_values) { constexpr std::size_t global_width = kHeads * kHeadDimension; const auto head_begin = static_cast(lane.head_offset) * kHeadDimension; const auto lane_width = static_cast(lane.heads) * kHeadDimension; const auto stage_vectors = [&](int index, std::span 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(token) * global_width + head_begin, lane_width); for (std::size_t element = 0; element < lane_width; ++element) { destination[static_cast(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(token) * kHeads + lane.head_offset, lane.heads); for (int head = 0; head < lane.heads; ++head) { beta_destination[static_cast(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> & active_lanes) { constexpr std::array 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 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(lane.head_offset) * kHeadDimension; const auto count = static_cast(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 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(lane.head_offset) * kHeadDimension; const auto lane_width = static_cast(lane.heads) * kHeadDimension; const auto * source = static_cast(output->virt_addr); for (int token = 0; token < lane.tokens; ++token) { auto destination = output_values.subspan( static_cast(token) * global_width + head_begin, lane_width); for (std::size_t element = 0; element < lane_width; ++element) { destination[element] = static_cast( source[static_cast(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 query, std::span key, std::span value, std::span decay, std::span beta, std::span output) { constexpr auto vector_elements = static_cast(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 query, std::span key, std::span value, std::span decay, std::span beta, std::span 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 query, std::span key, std::span value, std::span decay, std::span beta, std::span 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(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 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(worker) / 4; const std::size_t last = query.size() * static_cast(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(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 heads6_model, std::span heads5_model) : impl_(std::make_unique(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 query, std::span key, std::span value, std::span decay, std::span beta, std::span output) { return impl_->Run(query, key, value, decay, beta, output); } GdnRunTimings GdnStep::RunBatch16( std::span query, std::span key, std::span value, std::span decay, std::span beta, std::span output) { return impl_->RunBatch16(query, key, value, decay, beta, output); } GdnRunTimings GdnStep::RunBatchCpu( std::span query, std::span key, std::span value, std::span decay, std::span beta, std::span 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