Download src/gdn_step.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 40 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/gdn_step.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/gdn_step.cpp
-
curl -L -o gdn_step.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/gdn_step.cpp
40 kB
| namespace ling3 { | |
| namespace { | |
| using Clock = std::chrono::steady_clock; | |
| 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]); | |
| } | |
| } | |
| 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)); | |
| } | |
| } | |
| } | |
| } // namespace | |
| struct GdnStep::Impl { | |
| std::vector<std::unique_ptr<Lane>> lanes; | |
| std::vector<std::unique_ptr<Lane>> batch16_lanes; | |
| 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; | |
| std::size_t total_state_bytes = 0; | |
| Impl(std::span<const std::byte> heads6, std::span<const std::byte> heads5) { | |
| 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(); | |
| (void)heads6; | |
| (void)heads5; | |
| throw std::runtime_error("GdnStep requires a build with RKNN support"); | |
| } | |
| ~Impl() = default; | |
| 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"); | |
| } | |
| } | |
| 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; | |
| } | |
| 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; | |
| } | |
| 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"); | |
| } | |
| } | |
| 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"); | |
| } | |
| (void)query; | |
| (void)key; | |
| (void)value; | |
| (void)decay; | |
| (void)beta; | |
| (void)output; | |
| throw std::runtime_error("GdnStep requires a build with RKNN support"); | |
| 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), | |
| }; | |
| } | |
| 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"); | |
| } | |
| (void)query; | |
| (void)key; | |
| (void)value; | |
| (void)decay; | |
| (void)beta; | |
| (void)output; | |
| throw std::runtime_error("GdnStep requires a build with RKNN support"); | |
| 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), | |
| }; | |
| } | |
| 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) { | |
| (void)query; | |
| (void)key; | |
| (void)value; | |
| (void)decay; | |
| (void)beta; | |
| (void)output; | |
| throw std::runtime_error("CPU GDN prefill requires an AArch64 RKNN build"); | |
| 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), | |
| }; | |
| } | |
| }; | |
| 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 (!impl_->cpu_state_valid) impl_->CaptureDeviceState(); | |
| return {impl_->cpu_state, impl_->cpu_state_fp32, impl_->cpu_state_fp32_valid}; | |
| throw std::runtime_error("GDN checkpoints require the RKNN aarch64 backend"); | |
| } | |
| void GdnStep::RestoreState(const GdnState & state) { | |
| 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; | |
| (void)state; | |
| throw std::runtime_error("GDN checkpoints require the RKNN aarch64 backend"); | |
| } | |
| GdnStep::GdnStep(GdnStep &&) noexcept = default; | |
| GdnStep & GdnStep::operator=(GdnStep &&) noexcept = default; | |
| void GdnStep::Reset() { | |
| impl_->Reset(); | |
| throw std::runtime_error("GdnStep requires a build with RKNN support"); | |
| } | |
| 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 { | |
| return impl_->batch16_lanes.size() == 3; | |
| return false; | |
| } | |
| std::size_t GdnStep::state_bytes() const noexcept { return impl_->total_state_bytes; } | |
| } // namespace ling3 | |