Download src/w4_linear.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 48.8 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/w4_linear.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/w4_linear.cpp
-
curl -L -o w4_linear.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/w4_linear.cpp
48.8 kB
| namespace ling3 { | |
| namespace { | |
| using Clock = std::chrono::steady_clock; | |
| double Milliseconds(Clock::time_point begin, Clock::time_point end) { | |
| return std::chrono::duration<double, std::milli>(end - begin).count(); | |
| } | |
| std::vector<std::pair<int, int>> SplitAligned(int value, int parts, int alignment) { | |
| if (parts < 1 || value % alignment != 0 || value / alignment < parts) { | |
| throw std::invalid_argument("cannot split tensor dimension into aligned ranges"); | |
| } | |
| const int blocks = value / alignment; | |
| const int base = blocks / parts; | |
| const int extra = blocks % parts; | |
| std::vector<std::pair<int, int>> ranges; | |
| ranges.reserve(parts); | |
| int offset = 0; | |
| for (int index = 0; index < parts; ++index) { | |
| const int size = (base + (index < extra ? 1 : 0)) * alignment; | |
| ranges.emplace_back(offset, size); | |
| offset += size; | |
| } | |
| return ranges; | |
| } | |
| 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)); | |
| } | |
| } | |
| void PutNativeInt4(std::uint8_t * output, std::size_t index, std::int8_t value) { | |
| const std::uint8_t nibble = static_cast<std::uint8_t>(value) & 0x0FU; | |
| if ((index & 1U) == 0) { | |
| output[index / 2] = static_cast<std::uint8_t>(nibble << 4U); | |
| } else { | |
| output[index / 2] = static_cast<std::uint8_t>(output[index / 2] | nibble); | |
| } | |
| } | |
| void PackActivation16( | |
| int8x16_t values, | |
| std::uint8_t * high_output, | |
| std::uint8_t * low_output) { | |
| const uint8x16_t mask = vdupq_n_u8(0x0F); | |
| const uint8x16_t high = vandq_u8( | |
| vreinterpretq_u8_s8(vshrq_n_s8(values, 4)), mask); | |
| const uint8x16_t low = veorq_u8( | |
| vandq_u8(vreinterpretq_u8_s8(values), mask), vdupq_n_u8(0x08)); | |
| const uint8x8_t high_even = vget_low_u8(vuzp1q_u8(high, high)); | |
| const uint8x8_t high_odd = vget_low_u8(vuzp2q_u8(high, high)); | |
| const uint8x8_t low_even = vget_low_u8(vuzp1q_u8(low, low)); | |
| const uint8x8_t low_odd = vget_low_u8(vuzp2q_u8(low, low)); | |
| vst1_u8(high_output, vorr_u8(vshl_n_u8(high_even, 4), high_odd)); | |
| vst1_u8(low_output, vorr_u8(vshl_n_u8(low_even, 4), low_odd)); | |
| } | |
| void PackActivationRow( | |
| const std::int8_t * input, | |
| int count, | |
| std::uint8_t * high_output, | |
| std::uint8_t * low_output) { | |
| if ((count & 1) != 0) throw std::invalid_argument("INT4 row width must be even"); | |
| int column = 0; | |
| for (; column + 16 <= count; column += 16) { | |
| PackActivation16( | |
| vld1q_s8(input + column), | |
| high_output + column / 2, | |
| low_output + column / 2); | |
| } | |
| for (; column < count; column += 2) { | |
| auto nibble = [input, column](int lane, bool high) { | |
| const int value = input[column + lane]; | |
| const int code = high | |
| ? (value < 0 ? -((-value + 15) / 16) : value / 16) | |
| : ((value & 0x0F) ^ 0x08); | |
| return static_cast<std::uint8_t>(code) & 0x0FU; | |
| }; | |
| high_output[column / 2] = static_cast<std::uint8_t>( | |
| (nibble(0, true) << 4U) | nibble(1, true)); | |
| low_output[column / 2] = static_cast<std::uint8_t>( | |
| (nibble(0, false) << 4U) | nibble(1, false)); | |
| } | |
| } | |
| void PackSignedInt4Row( | |
| const std::int8_t * input, | |
| int count, | |
| std::uint8_t * output) { | |
| if ((count & 1) != 0) throw std::invalid_argument("INT4 row width must be even"); | |
| for (int column = 0; column < count; column += 2) { | |
| output[column / 2] = static_cast<std::uint8_t>( | |
| (static_cast<std::uint8_t>(input[column]) & 0x0FU) << 4U | | |
| (static_cast<std::uint8_t>(input[column + 1]) & 0x0FU)); | |
| } | |
| } | |
| struct Segment { | |
| int part = 0; | |
| int offset_k = 0; | |
| int size_k = 0; | |
| int offset_n = 0; | |
| int size_n = 0; | |
| int core = 0; | |
| rknn_matmul_ctx context = 0; | |
| rknn_matmul_info info {}; | |
| rknn_matmul_io_attr attributes {}; | |
| rknn_tensor_mem * a = nullptr; | |
| rknn_tensor_mem * b = nullptr; | |
| rknn_tensor_mem * c = nullptr; | |
| bool owns_a = true; | |
| bool owns_b = true; | |
| Segment() = default; | |
| Segment(const Segment &) = delete; | |
| Segment & operator=(const Segment &) = delete; | |
| Segment(Segment && other) noexcept { *this = std::move(other); } | |
| Segment & operator=(Segment && other) noexcept { | |
| if (this == &other) return *this; | |
| Release(); | |
| part = other.part; | |
| offset_k = other.offset_k; | |
| size_k = other.size_k; | |
| offset_n = other.offset_n; | |
| size_n = other.size_n; | |
| core = other.core; | |
| context = std::exchange(other.context, 0); | |
| info = other.info; | |
| attributes = other.attributes; | |
| a = std::exchange(other.a, nullptr); | |
| b = std::exchange(other.b, nullptr); | |
| c = std::exchange(other.c, nullptr); | |
| owns_a = std::exchange(other.owns_a, true); | |
| owns_b = std::exchange(other.owns_b, true); | |
| return *this; | |
| } | |
| ~Segment() { Release(); } | |
| void Release() noexcept { | |
| if (context == 0) return; | |
| if (a != nullptr && owns_a) rknn_destroy_mem(context, a); | |
| if (b != nullptr && owns_b) rknn_destroy_mem(context, b); | |
| if (c != nullptr) rknn_destroy_mem(context, c); | |
| rknn_matmul_destroy(context); | |
| context = 0; | |
| a = nullptr; | |
| b = nullptr; | |
| c = nullptr; | |
| owns_a = true; | |
| owns_b = true; | |
| } | |
| }; | |
| struct BatchWorkspace { | |
| int rows = 0; | |
| bool w4a4 = false; | |
| std::vector<Segment> segments; | |
| std::vector<std::int8_t> quantized; | |
| std::vector<std::int32_t> accumulator; | |
| std::vector<float> scales; | |
| }; | |
| // All batch shapes of one segment execute sequentially. Each keeps its own | |
| // exact-sized FD views, backed by these maximum-sized, same-domain owners. | |
| // Parallel cores and K segments never share these buffers with each other. | |
| struct BatchMemory { | |
| rknn_matmul_ctx context = 0; // Borrowed from the live decode segment. | |
| rknn_tensor_mem * a = nullptr; | |
| rknn_tensor_mem * c = nullptr; | |
| ~BatchMemory() { | |
| if (a) rknn_destroy_mem(context, a); | |
| if (c) rknn_destroy_mem(context, c); | |
| } | |
| }; | |
| } // namespace | |
| struct DynamicW4Linear::Impl { | |
| W4LinearConfig config; | |
| std::vector<float> scales; | |
| std::vector<std::int32_t> correction; | |
| std::vector<std::int8_t> quantized; | |
| std::vector<std::int32_t> accumulator; | |
| std::size_t resident_bytes = 0; | |
| bool prefill_w4a4 = false; | |
| std::vector<Segment> segments; | |
| // Reverse destruction order: batch views, backing memory, source contexts. | |
| std::vector<std::unique_ptr<BatchMemory>> batch_memory; | |
| std::array<std::unique_ptr<BatchWorkspace>, 8> batch_workspaces; | |
| Impl( | |
| W4LinearConfig value, | |
| std::span<const std::byte> weights, | |
| std::span<const float> weight_scales, | |
| std::span<const std::int32_t> weight_correction) | |
| : config(std::move(value)), | |
| scales(weight_scales.begin(), weight_scales.end()), | |
| correction(weight_correction.begin(), weight_correction.end()), | |
| prefill_w4a4(std::getenv("LING3_PREFILL_W4A4") != nullptr) { | |
| Validate(weights); | |
| quantized.resize(config.k); | |
| accumulator.resize(config.n); | |
| Initialize(weights); | |
| (void)weights; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| } | |
| ~Impl() = default; | |
| void Validate(std::span<const std::byte> weights) const { | |
| if (config.k < 1 || config.n < 1 || config.k_splits < 1 || | |
| config.k % 32 != 0 || config.n % 64 != 0) { | |
| throw std::invalid_argument("W4 linear requires positive K/N, K%32=0 and N%64=0"); | |
| } | |
| if (config.cores.empty() || config.cores.size() > 3) { | |
| throw std::invalid_argument("W4 linear requires one to three NPU cores"); | |
| } | |
| if (config.iommu_domain_id < 0 || config.iommu_domain_id > 15) { | |
| throw std::invalid_argument("W4 IOMMU domain must be in [0, 15]"); | |
| } | |
| auto sorted = config.cores; | |
| std::sort(sorted.begin(), sorted.end()); | |
| if (sorted.front() < 0 || sorted.back() > 2 || | |
| std::adjacent_find(sorted.begin(), sorted.end()) != sorted.end()) { | |
| throw std::invalid_argument("NPU cores must be unique values in [0, 2]"); | |
| } | |
| SplitAligned(config.k, config.k_splits, 32); | |
| SplitAligned(config.n, static_cast<int>(config.cores.size()), 64); | |
| const auto expected = static_cast<std::size_t>(config.k) * config.n / 2; | |
| if (weights.size() != expected || scales.size() != static_cast<std::size_t>(config.n) || | |
| correction.size() != static_cast<std::size_t>(config.n)) { | |
| throw std::invalid_argument("W4 weight, scale, or correction size is inconsistent"); | |
| } | |
| for (float scale : scales) { | |
| if (!(scale > 0.0F) || !std::isfinite(scale)) { | |
| throw std::invalid_argument("W4 weight scale must be finite and positive"); | |
| } | |
| } | |
| } | |
| void Initialize(std::span<const std::byte> weights) { | |
| const auto k_ranges = SplitAligned(config.k, config.k_splits, 32); | |
| const auto n_ranges = SplitAligned(config.n, static_cast<int>(config.cores.size()), 64); | |
| segments.reserve(k_ranges.size() * n_ranges.size()); | |
| for (std::size_t part = 0; part < k_ranges.size(); ++part) { | |
| for (std::size_t lane = 0; lane < n_ranges.size(); ++lane) { | |
| Segment segment; | |
| segment.part = static_cast<int>(part); | |
| segment.offset_k = k_ranges[part].first; | |
| segment.size_k = k_ranges[part].second; | |
| segment.offset_n = n_ranges[lane].first; | |
| segment.size_n = n_ranges[lane].second; | |
| segment.core = config.cores[lane]; | |
| InitializeSegment(segment, weights); | |
| resident_bytes += segment.attributes.B.size; | |
| segments.push_back(std::move(segment)); | |
| } | |
| } | |
| } | |
| void InitializeSegment(Segment & segment, std::span<const std::byte> weights) { | |
| segment.info.M = 2; | |
| segment.info.K = segment.size_k; | |
| segment.info.N = segment.size_n; | |
| segment.info.type = RKNN_INT4_MM_INT4_TO_INT16; | |
| segment.info.B_layout = RKNN_MM_LAYOUT_NATIVE; | |
| segment.info.B_quant_type = RKNN_QUANT_TYPE_PER_LAYER_SYM; | |
| segment.info.AC_layout = RKNN_MM_LAYOUT_NATIVE; | |
| segment.info.AC_quant_type = RKNN_QUANT_TYPE_PER_LAYER_SYM; | |
| segment.info.iommu_domain_id = config.iommu_domain_id; | |
| CheckRknn( | |
| rknn_matmul_create(&segment.context, &segment.info, &segment.attributes), | |
| "rknn_matmul_create W4"); | |
| CheckRknn( | |
| rknn_matmul_set_core_mask( | |
| segment.context, static_cast<rknn_core_mask>(1U << segment.core)), | |
| "rknn_matmul_set_core_mask W4"); | |
| const auto expected_a = static_cast<std::size_t>(segment.size_k); | |
| const auto expected_b = static_cast<std::size_t>(segment.size_k) * segment.size_n / 2; | |
| const auto expected_c = static_cast<std::size_t>(2) * segment.size_n * sizeof(std::int16_t); | |
| if (segment.attributes.A.type != RKNN_TENSOR_INT4 || | |
| segment.attributes.B.type != RKNN_TENSOR_INT4 || | |
| segment.attributes.C.type != RKNN_TENSOR_INT16 || | |
| segment.attributes.A.size != expected_a || | |
| segment.attributes.B.size != expected_b || | |
| segment.attributes.C.size != expected_c || | |
| segment.attributes.A.n_dims != 3 || segment.attributes.B.n_dims != 4 || | |
| segment.attributes.C.n_dims != 3) { | |
| throw std::runtime_error("RKNN returned unexpected W4 tensor attributes"); | |
| } | |
| segment.a = rknn_create_mem2( | |
| segment.context, segment.attributes.A.size, RKNN_FLAG_MEMORY_CACHEABLE); | |
| segment.b = rknn_create_mem2( | |
| segment.context, segment.attributes.B.size, RKNN_FLAG_MEMORY_CACHEABLE); | |
| segment.c = rknn_create_mem2( | |
| segment.context, segment.attributes.C.size, RKNN_FLAG_MEMORY_CACHEABLE); | |
| if (segment.a == nullptr || segment.b == nullptr || segment.c == nullptr) { | |
| throw std::runtime_error("rknn_create_mem2 failed for W4 linear"); | |
| } | |
| std::memset(segment.b->virt_addr, 0, segment.attributes.B.size); | |
| const int sub_n = static_cast<int>(segment.attributes.B.dims[2]); | |
| const int sub_k = static_cast<int>(segment.attributes.B.dims[3]); | |
| const int n_blocks = (segment.size_n + sub_n - 1) / sub_n; | |
| const int k_blocks = (segment.size_k + sub_k - 1) / sub_k; | |
| auto * native = static_cast<std::uint8_t *>(segment.b->virt_addr); | |
| for (int n_block = 0; n_block < n_blocks; ++n_block) { | |
| for (int k_block = 0; k_block < k_blocks; ++k_block) { | |
| for (int inner_n = 0; inner_n < sub_n; ++inner_n) { | |
| const int local_n = n_block * sub_n + inner_n; | |
| for (int inner_k = 0; inner_k < sub_k; ++inner_k) { | |
| const int local_k = k_block * sub_k + inner_k; | |
| const auto code = local_k < segment.size_k && local_n < segment.size_n | |
| ? DecodeInt4LowFirst( | |
| weights, | |
| static_cast<std::size_t>(segment.offset_k + local_k) * config.n + | |
| segment.offset_n + local_n) | |
| : static_cast<std::int8_t>(0); | |
| const auto output_index = | |
| ((static_cast<std::size_t>(n_block) * k_blocks + k_block) * sub_n + | |
| inner_n) * sub_k + inner_k; | |
| PutNativeInt4(native, output_index, code); | |
| } | |
| } | |
| } | |
| } | |
| CheckRknn( | |
| rknn_mem_sync(segment.context, segment.b, RKNN_MEMORY_SYNC_TO_DEVICE), | |
| "sync W4 weight"); | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.a, &segment.attributes.A), | |
| "bind W4 input"); | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.b, &segment.attributes.B), | |
| "bind W4 weight"); | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.c, &segment.attributes.C), | |
| "bind W4 output"); | |
| } | |
| static int BatchRows(std::size_t rows) { | |
| for (const int candidate : {1, 2, 4, 8, 16, 32, 64, 128}) { | |
| if (rows <= static_cast<std::size_t>(candidate)) return candidate; | |
| } | |
| throw std::invalid_argument("W4 batch supports at most 128 rows"); | |
| } | |
| static std::size_t BatchSlot(int rows) { | |
| std::size_t slot = 0; | |
| while ((1 << slot) < rows) ++slot; | |
| return slot; | |
| } | |
| void InitializeBatchMemory() { | |
| if (!batch_memory.empty()) return; | |
| std::vector<std::unique_ptr<BatchMemory>> owners; | |
| owners.reserve(segments.size()); | |
| const std::size_t matrix_rows = prefill_w4a4 ? 128 : 256; | |
| for (const auto & segment : segments) { | |
| auto memory = std::make_unique<BatchMemory>(); | |
| memory->context = segment.context; | |
| memory->a = rknn_create_mem2(segment.context, | |
| matrix_rows * segment.size_k / 2, RKNN_FLAG_MEMORY_CACHEABLE); | |
| memory->c = rknn_create_mem2(segment.context, | |
| matrix_rows * segment.size_n * sizeof(std::int16_t), RKNN_FLAG_MEMORY_CACHEABLE); | |
| if (!memory->a || !memory->c) | |
| throw std::runtime_error("rknn_create_mem2 failed for shared W4 batch memory"); | |
| owners.push_back(std::move(memory)); | |
| } | |
| batch_memory = std::move(owners); | |
| } | |
| BatchWorkspace & InitializeBatch(std::size_t active_rows, bool indexed_input) { | |
| const int rows = BatchRows(active_rows); | |
| auto & stored = batch_workspaces[BatchSlot(rows)]; | |
| if (stored) { | |
| if (!indexed_input && stored->quantized.empty()) | |
| stored->quantized.resize(static_cast<std::size_t>(rows) * config.k); | |
| return *stored; | |
| } | |
| auto workspace = std::make_unique<BatchWorkspace>(); | |
| workspace->rows = rows; | |
| workspace->w4a4 = prefill_w4a4 && rows >= 2; | |
| if (!indexed_input) | |
| workspace->quantized.resize(static_cast<std::size_t>(rows) * config.k); | |
| // GatherBatch writes the final K part straight to the float output. | |
| // Only split-K matrices need storage for preceding partial sums. | |
| if (config.k_splits > 1) | |
| workspace->accumulator.resize(static_cast<std::size_t>(rows) * config.n); | |
| workspace->scales.resize(rows, 1.0F); | |
| workspace->segments.reserve(segments.size()); | |
| InitializeBatchMemory(); | |
| for (std::size_t segment_index = 0; segment_index < segments.size(); ++segment_index) { | |
| const auto & source = segments[segment_index]; | |
| Segment segment; | |
| segment.part = source.part; | |
| segment.offset_k = source.offset_k; | |
| segment.size_k = source.size_k; | |
| segment.offset_n = source.offset_n; | |
| segment.size_n = source.size_n; | |
| segment.core = source.core; | |
| segment.info = source.info; | |
| segment.info.M = workspace->w4a4 ? rows : 2 * rows; | |
| CheckRknn( | |
| rknn_matmul_create(&segment.context, &segment.info, &segment.attributes), | |
| "rknn_matmul_create W4 batch"); | |
| CheckRknn( | |
| rknn_matmul_set_core_mask( | |
| segment.context, static_cast<rknn_core_mask>(1U << segment.core)), | |
| "rknn_matmul_set_core_mask W4 batch"); | |
| const auto matrix_rows = workspace->w4a4 ? rows : 2 * rows; | |
| const auto expected_a = static_cast<std::size_t>(matrix_rows) * | |
| segment.size_k / 2; | |
| const auto expected_c = static_cast<std::size_t>(matrix_rows) * | |
| segment.size_n * sizeof(std::int16_t); | |
| if (segment.attributes.A.type != RKNN_TENSOR_INT4 || | |
| segment.attributes.B.type != RKNN_TENSOR_INT4 || | |
| segment.attributes.C.type != RKNN_TENSOR_INT16 || | |
| segment.attributes.A.size != expected_a || | |
| segment.attributes.B.size != source.attributes.B.size || | |
| segment.attributes.C.size != expected_c) { | |
| throw std::runtime_error("RKNN returned unexpected W4 batch attributes"); | |
| } | |
| const auto & owner = *batch_memory[segment_index]; | |
| if (segment.attributes.A.size > owner.a->size || segment.attributes.C.size > owner.c->size) | |
| throw std::runtime_error("W4 batch view exceeds its shared backing memory"); | |
| segment.a = rknn_create_mem_from_fd(segment.context, owner.a->fd, | |
| owner.a->virt_addr, segment.attributes.A.size, 0); | |
| segment.c = rknn_create_mem_from_fd(segment.context, owner.c->fd, | |
| owner.c->virt_addr, segment.attributes.C.size, 0); | |
| if (segment.a == nullptr || segment.c == nullptr) { | |
| throw std::runtime_error("rknn_create_mem_from_fd failed for W4 batch view"); | |
| } | |
| if (segment.a->size != segment.attributes.A.size || segment.c->size != segment.attributes.C.size) | |
| throw std::runtime_error("W4 batch view changed the requested synchronization size"); | |
| segment.b = source.b; | |
| segment.owns_b = false; | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.a, &segment.attributes.A), | |
| "bind W4 batch input"); | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.b, &segment.attributes.B), | |
| "bind shared W4 batch weight"); | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(segment.context, segment.c, &segment.attributes.C), | |
| "bind W4 batch output"); | |
| workspace->segments.push_back(std::move(segment)); | |
| } | |
| stored = std::move(workspace); | |
| return *stored; | |
| } | |
| void StageInput(std::span<const float> input, float & input_scale) { | |
| input_scale = QuantizeSymmetricInt8(input, quantized).scale; | |
| for (int part = 0; part < config.k_splits; ++part) { | |
| const auto first = std::find_if(segments.begin(), segments.end(), [part](const Segment & item) { | |
| return item.part == part; | |
| }); | |
| if (first == segments.end()) throw std::runtime_error("W4 K split is missing"); | |
| const int sub_k = static_cast<int>(first->attributes.A.dims[2]); | |
| if (sub_k < 1 || first->size_k % sub_k != 0) { | |
| throw std::runtime_error("RKNN native W4 input has a partial K block"); | |
| } | |
| auto * packed = static_cast<std::uint8_t *>(first->a->virt_addr); | |
| for (int block = 0; block < first->size_k / sub_k; ++block) { | |
| const auto high_index = static_cast<std::size_t>(block * 2) * sub_k; | |
| const auto low_index = static_cast<std::size_t>(block * 2 + 1) * sub_k; | |
| PackActivationRow( | |
| quantized.data() + first->offset_k + block * sub_k, | |
| sub_k, | |
| packed + high_index / 2, | |
| packed + low_index / 2); | |
| } | |
| for (auto & segment : segments) { | |
| if (segment.part != part || &segment == &*first) continue; | |
| if (segment.attributes.A.size != first->attributes.A.size) { | |
| throw std::runtime_error("RKNN W4 contexts disagree on input size"); | |
| } | |
| std::memcpy(segment.a->virt_addr, first->a->virt_addr, first->attributes.A.size); | |
| } | |
| } | |
| } | |
| void StageBatch( | |
| BatchWorkspace & workspace, | |
| std::span<const float> input, | |
| std::size_t active_rows) { | |
| std::fill_n(workspace.quantized.begin(), active_rows * config.k, 0); | |
| std::fill_n(workspace.scales.begin(), active_rows, 1.0F); | |
| for (std::size_t row = 0; row < active_rows; ++row) { | |
| const auto row_input = input.subspan( | |
| row * config.k, config.k); | |
| auto row_quantized = std::span<std::int8_t>(workspace.quantized).subspan( | |
| row * config.k, config.k); | |
| workspace.scales[row] = workspace.w4a4 | |
| ? QuantizeSymmetricInt4(row_input, row_quantized).scale | |
| : QuantizeSymmetricInt8(row_input, row_quantized).scale; | |
| } | |
| PackBatch(workspace, active_rows, workspace.quantized, {}); | |
| } | |
| void PackBatch( | |
| BatchWorkspace & workspace, | |
| std::size_t active_rows, | |
| std::span<const std::int8_t> quantized, | |
| std::span<const std::size_t> row_indices) { | |
| const int packed_rows = workspace.w4a4 ? workspace.rows : 2 * workspace.rows; | |
| for (int part = 0; part < config.k_splits; ++part) { | |
| const auto first = std::find_if( | |
| workspace.segments.begin(), workspace.segments.end(), | |
| [part](const Segment & item) { return item.part == part; }); | |
| if (first == workspace.segments.end()) { | |
| throw std::runtime_error("W4 batch K split is missing"); | |
| } | |
| const int sub_k = static_cast<int>(first->attributes.A.dims[2]); | |
| if (sub_k < 1 || first->size_k % sub_k != 0) { | |
| throw std::runtime_error("RKNN native W4 batch input has a partial K block"); | |
| } | |
| auto * packed = static_cast<std::uint8_t *>(first->a->virt_addr); | |
| for (int block = 0; block < first->size_k / sub_k; ++block) { | |
| for (int row = 0; row < workspace.rows; ++row) { | |
| const auto element = static_cast<std::size_t>( | |
| block * packed_rows + (workspace.w4a4 ? row : 2 * row)) * sub_k; | |
| if (static_cast<std::size_t>(row) >= active_rows) { | |
| std::memset(packed + element / 2, 0, sub_k / 2); | |
| if (!workspace.w4a4) { | |
| std::memset(packed + (element + sub_k) / 2, 0x88, sub_k / 2); | |
| } | |
| continue; | |
| } | |
| const std::size_t source_row = row_indices.empty() ? row : row_indices[row]; | |
| const auto * source = quantized.data() + | |
| source_row * config.k + first->offset_k + block * sub_k; | |
| if (workspace.w4a4) { | |
| PackSignedInt4Row(source, sub_k, packed + element / 2); | |
| } else { | |
| PackActivationRow( | |
| source, sub_k, packed + element / 2, | |
| packed + (element + sub_k) / 2); | |
| } | |
| } | |
| } | |
| for (auto & segment : workspace.segments) { | |
| if (segment.part != part || &segment == &*first) continue; | |
| if (segment.attributes.A.size != first->attributes.A.size) { | |
| throw std::runtime_error("RKNN W4 batch contexts disagree on input size"); | |
| } | |
| std::memcpy(segment.a->virt_addr, first->a->virt_addr, first->attributes.A.size); | |
| } | |
| } | |
| } | |
| void SyncInputs() { | |
| for (auto & segment : segments) { | |
| CheckRknn( | |
| rknn_mem_sync(segment.context, segment.a, RKNN_MEMORY_SYNC_TO_DEVICE), | |
| "sync W4 input"); | |
| } | |
| } | |
| void SyncBatchInputs(BatchWorkspace & workspace) { | |
| for (auto & segment : workspace.segments) { | |
| CheckRknn( | |
| rknn_mem_sync(segment.context, segment.a, RKNN_MEMORY_SYNC_TO_DEVICE), | |
| "sync W4 batch input"); | |
| } | |
| } | |
| void BindBatchWeights(BatchWorkspace & workspace, const Impl & source) { | |
| if (config.k != source.config.k || config.n != source.config.n || | |
| config.k_splits != source.config.k_splits || | |
| config.iommu_domain_id != source.config.iommu_domain_id || | |
| workspace.segments.size() != source.segments.size()) { | |
| throw std::invalid_argument("W4 batch source has an incompatible shape"); | |
| } | |
| for (std::size_t index = 0; index < workspace.segments.size(); ++index) { | |
| auto & target = workspace.segments[index]; | |
| const auto & weight = source.segments[index]; | |
| if (target.attributes.B.size != weight.attributes.B.size || weight.b == nullptr) { | |
| throw std::invalid_argument("W4 batch source has incompatible native weights"); | |
| } | |
| if (target.b == weight.b) continue; | |
| target.b = weight.b; | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(target.context, target.b, &target.attributes.B), | |
| "rebind W4 batch weight"); | |
| } | |
| } | |
| void RunWorkers(std::vector<Segment> & active_segments) { | |
| if (config.cores.size() == 1) { | |
| for (auto & segment : active_segments) { | |
| CheckRknn(rknn_matmul_run(segment.context), "rknn_matmul_run W4"); | |
| } | |
| return; | |
| } | |
| CoreWorkers::Instance().Run(config.cores, [&active_segments](int core) { | |
| for (auto & segment : active_segments) { | |
| if (segment.core == core) { | |
| CheckRknn(rknn_matmul_run(segment.context), "rknn_matmul_run W4"); | |
| } | |
| } | |
| }); | |
| } | |
| void RunWorkers() { RunWorkers(segments); } | |
| void Gather(float input_scale, std::span<float> output) { | |
| for (auto & segment : segments) { | |
| CheckRknn( | |
| rknn_mem_sync(segment.context, segment.c, RKNN_MEMORY_SYNC_FROM_DEVICE), | |
| "sync W4 output"); | |
| const auto * source = static_cast<const std::int16_t *>(segment.c->virt_addr); | |
| const int sub_n = static_cast<int>(segment.attributes.C.dims[2]); | |
| const int n_blocks = static_cast<int>(segment.attributes.C.dims[0]); | |
| for (int block = 0; block < n_blocks; ++block) { | |
| const int global_n = segment.offset_n + block * sub_n; | |
| auto * destination = accumulator.data() + global_n; | |
| const auto * high = source + static_cast<std::size_t>(block * 2) * sub_n; | |
| const auto * low = source + static_cast<std::size_t>(block * 2 + 1) * sub_n; | |
| const int columns = std::min(sub_n, segment.size_n - block * sub_n); | |
| int column = 0; | |
| for (; column + 8 <= columns; column += 8) { | |
| int32x4_t value0 = segment.part == 0 | |
| ? vld1q_s32(correction.data() + global_n + column) | |
| : vld1q_s32(destination + column); | |
| int32x4_t value1 = segment.part == 0 | |
| ? vld1q_s32(correction.data() + global_n + column + 4) | |
| : vld1q_s32(destination + column + 4); | |
| const int16x8_t high8 = vld1q_s16(high + column); | |
| const int16x8_t low8 = vld1q_s16(low + column); | |
| value0 = vmlal_n_s16(value0, vget_low_s16(high8), 16); | |
| value1 = vmlal_n_s16(value1, vget_high_s16(high8), 16); | |
| value0 = vaddw_s16(value0, vget_low_s16(low8)); | |
| value1 = vaddw_s16(value1, vget_high_s16(low8)); | |
| vst1q_s32(destination + column, value0); | |
| vst1q_s32(destination + column + 4, value1); | |
| } | |
| for (; column < columns; ++column) { | |
| const auto initial = segment.part == 0 | |
| ? correction[global_n + column] | |
| : destination[column]; | |
| destination[column] = initial + 16 * static_cast<std::int32_t>(high[column]) + | |
| static_cast<std::int32_t>(low[column]); | |
| } | |
| } | |
| } | |
| DequantizePerChannel(accumulator, input_scale, scales, output); | |
| } | |
| void GatherBatch( | |
| BatchWorkspace & workspace, | |
| std::size_t active_rows, | |
| const Impl & source_weights, | |
| std::span<float> output) { | |
| const int packed_rows = workspace.w4a4 ? workspace.rows : 2 * workspace.rows; | |
| for (auto & segment : workspace.segments) { | |
| CheckRknn( | |
| rknn_mem_sync(segment.context, segment.c, RKNN_MEMORY_SYNC_FROM_DEVICE), | |
| "sync W4 batch output"); | |
| } | |
| // Expert runners already execute inside CoreWorkers. Only main-thread | |
| // multi-core projections may dispatch the pool; each worker owns rows | |
| // through all K parts, preserving accumulation and multiplication order. | |
| const auto gather_rows = [&](std::size_t first, std::size_t last) { | |
| for (auto & segment : workspace.segments) { | |
| const bool final_part = segment.part == config.k_splits - 1; | |
| const auto * source = static_cast<const std::int16_t *>(segment.c->virt_addr); | |
| const int sub_n = static_cast<int>(segment.attributes.C.dims[2]); | |
| const int n_blocks = static_cast<int>(segment.attributes.C.dims[0]); | |
| for (int block = 0; block < n_blocks; ++block) { | |
| const int global_n = segment.offset_n + block * sub_n; | |
| const int columns = std::min(sub_n, segment.size_n - block * sub_n); | |
| for (std::size_t row = first; row < last; ++row) { | |
| // Avoid offset-pointer arithmetic on empty storage. | |
| auto * destination = config.k_splits > 1 | |
| ? workspace.accumulator.data() + row * config.n + global_n | |
| : nullptr; | |
| auto * result = output.data() + row * config.n + global_n; | |
| const auto source_row = static_cast<std::size_t>(block) * packed_rows + | |
| (workspace.w4a4 ? row : 2 * row); | |
| const auto * high = source + source_row * sub_n; | |
| const auto * low = workspace.w4a4 ? nullptr : high + sub_n; | |
| int column = 0; | |
| for (; column + 8 <= columns; column += 8) { | |
| int32x4_t value0 = segment.part == 0 | |
| ? (workspace.w4a4 | |
| ? vdupq_n_s32(0) | |
| : vld1q_s32( | |
| source_weights.correction.data() + global_n + column)) | |
| : vld1q_s32(destination + column); | |
| int32x4_t value1 = segment.part == 0 | |
| ? (workspace.w4a4 | |
| ? vdupq_n_s32(0) | |
| : vld1q_s32( | |
| source_weights.correction.data() + global_n + column + 4)) | |
| : vld1q_s32(destination + column + 4); | |
| const int16x8_t high8 = vld1q_s16(high + column); | |
| if (workspace.w4a4) { | |
| value0 = vaddw_s16(value0, vget_low_s16(high8)); | |
| value1 = vaddw_s16(value1, vget_high_s16(high8)); | |
| } else { | |
| const int16x8_t low8 = vld1q_s16(low + column); | |
| value0 = vmlal_n_s16(value0, vget_low_s16(high8), 16); | |
| value1 = vmlal_n_s16(value1, vget_high_s16(high8), 16); | |
| value0 = vaddw_s16(value0, vget_low_s16(low8)); | |
| value1 = vaddw_s16(value1, vget_high_s16(low8)); | |
| } | |
| if (final_part) { | |
| // Preserve DequantizePerChannel's multiplication order, | |
| // but avoid writing and rereading the final accumulator. | |
| const float32x4_t activation = vdupq_n_f32(workspace.scales[row]); | |
| vst1q_f32(result + column, vmulq_f32( | |
| vmulq_f32(vcvtq_f32_s32(value0), activation), | |
| vld1q_f32(source_weights.scales.data() + global_n + column))); | |
| vst1q_f32(result + column + 4, vmulq_f32( | |
| vmulq_f32(vcvtq_f32_s32(value1), activation), | |
| vld1q_f32(source_weights.scales.data() + global_n + column + 4))); | |
| } else { | |
| vst1q_s32(destination + column, value0); | |
| vst1q_s32(destination + column + 4, value1); | |
| } | |
| } | |
| for (; column < columns; ++column) { | |
| std::int32_t value; | |
| if (workspace.w4a4) { | |
| const auto initial = segment.part == 0 ? 0 : destination[column]; | |
| value = | |
| initial + static_cast<std::int32_t>(high[column]); | |
| } else { | |
| const auto initial = segment.part == 0 | |
| ? source_weights.correction[global_n + column] | |
| : destination[column]; | |
| value = initial + | |
| 16 * static_cast<std::int32_t>(high[column]) + | |
| static_cast<std::int32_t>(low[column]); | |
| } | |
| if (final_part) { | |
| result[column] = static_cast<float>(value) * workspace.scales[row] * | |
| source_weights.scales[global_n + column]; | |
| } else { | |
| destination[column] = value; | |
| } | |
| } | |
| } | |
| } | |
| } | |
| }; | |
| if (config.cores.size() > 1 && active_rows * config.n >= 65536 && | |
| std::getenv("LING3_DISABLE_PARALLEL_GATHER") == nullptr) { | |
| constexpr std::array<int, 4> workers {0, 1, 2, 3}; | |
| CoreWorkers::Instance().Run(workers, [&](int worker) { | |
| gather_rows(active_rows * worker / 4, active_rows * (worker + 1) / 4); | |
| }); | |
| } else { | |
| gather_rows(0, active_rows); | |
| } | |
| } | |
| W4RunTimings Run(std::span<const float> input, std::span<float> output) { | |
| if (input.size() != static_cast<std::size_t>(config.k) || | |
| output.size() != static_cast<std::size_t>(config.n)) { | |
| throw std::invalid_argument("W4 input or output has the wrong size"); | |
| } | |
| (void)input; | |
| (void)output; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| W4RunTimings timings; | |
| const auto begin = Clock::now(); | |
| float input_scale = 1.0F; | |
| StageInput(input, input_scale); | |
| const auto staged = Clock::now(); | |
| SyncInputs(); | |
| const auto synced = Clock::now(); | |
| RunWorkers(); | |
| const auto executed = Clock::now(); | |
| Gather(input_scale, output); | |
| const auto end = Clock::now(); | |
| timings.quantize_pack_ms = Milliseconds(begin, staged); | |
| timings.input_sync_ms = Milliseconds(staged, synced); | |
| timings.npu_ms = Milliseconds(synced, executed); | |
| timings.gather_ms = Milliseconds(executed, end); | |
| timings.total_ms = Milliseconds(begin, end); | |
| timings.input_scale = input_scale; | |
| return timings; | |
| } | |
| W4RunTimings RunBatch( | |
| std::span<const float> input, | |
| std::size_t rows, | |
| const Impl & source_weights, | |
| std::span<float> output, | |
| std::span<const std::int8_t> quantized = {}, | |
| std::span<const float> input_scales = {}, | |
| std::span<const std::size_t> row_indices = {}) { | |
| const bool indexed = !row_indices.empty(); | |
| if (rows < 1 || rows > 128 || (!indexed && input.size() != rows * config.k) || | |
| output.size() != rows * config.n) { | |
| throw std::invalid_argument("W4 batch input, rows, or output has the wrong size"); | |
| } | |
| if (indexed) { | |
| if (row_indices.size() != rows || input_scales.empty() || | |
| quantized.size() / config.k != input_scales.size() || | |
| quantized.size() % config.k != 0) { | |
| throw std::invalid_argument("W4 indexed activation table has the wrong size"); | |
| } | |
| if (prefill_w4a4 && rows >= 2) { | |
| throw std::invalid_argument("indexed INT8 activations require W4A8 batch mode"); | |
| } | |
| for (const auto row : row_indices) { | |
| if (row >= input_scales.size() || !(input_scales[row] > 0.0F) || | |
| !std::isfinite(input_scales[row])) { | |
| throw std::invalid_argument("W4 indexed activation row or scale is invalid"); | |
| } | |
| } | |
| } | |
| (void)input; | |
| (void)rows; | |
| (void)source_weights; | |
| (void)output; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| auto & workspace = InitializeBatch(rows, indexed); | |
| BindBatchWeights(workspace, source_weights); | |
| W4RunTimings timings; | |
| const auto begin = Clock::now(); | |
| if (indexed) { | |
| for (std::size_t row = 0; row < rows; ++row) { | |
| workspace.scales[row] = input_scales[row_indices[row]]; | |
| } | |
| PackBatch(workspace, rows, quantized, row_indices); | |
| } else { | |
| StageBatch(workspace, input, rows); | |
| } | |
| const auto staged = Clock::now(); | |
| SyncBatchInputs(workspace); | |
| const auto synced = Clock::now(); | |
| RunWorkers(workspace.segments); | |
| const auto executed = Clock::now(); | |
| GatherBatch(workspace, rows, source_weights, output); | |
| const auto end = Clock::now(); | |
| timings.quantize_pack_ms = Milliseconds(begin, staged); | |
| timings.input_sync_ms = Milliseconds(staged, synced); | |
| timings.npu_ms = Milliseconds(synced, executed); | |
| timings.gather_ms = Milliseconds(executed, end); | |
| timings.total_ms = Milliseconds(begin, end); | |
| return timings; | |
| } | |
| W4RunTimings RunBatch32(std::span<const float> input, std::span<float> output) { | |
| return RunBatch(input, 32, *this, output); | |
| } | |
| void PrepareBatch(std::size_t rows, bool indexed_input) { | |
| if (rows < 1 || rows > 128) { | |
| throw std::invalid_argument("W4 batch rows must be in [1, 128]"); | |
| } | |
| (void)rows; | |
| (void)indexed_input; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| if (indexed_input && prefill_w4a4 && rows >= 2) | |
| throw std::invalid_argument("indexed INT8 activations require W4A8 batch mode"); | |
| auto & workspace = InitializeBatch(rows, indexed_input); | |
| BindBatchWeights(workspace, *this); | |
| } | |
| float PrepareInput(std::span<const float> input) { | |
| if (input.size() != static_cast<std::size_t>(config.k)) { | |
| throw std::invalid_argument("W4 input has the wrong size"); | |
| } | |
| (void)input; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| float input_scale = 1.0F; | |
| StageInput(input, input_scale); | |
| SyncInputs(); | |
| return input_scale; | |
| } | |
| W4RunTimings RunPrepared(float input_scale, std::span<float> output) { | |
| if (output.size() != static_cast<std::size_t>(config.n) || | |
| !(input_scale > 0.0F) || !std::isfinite(input_scale)) { | |
| throw std::invalid_argument("prepared W4 scale or output has the wrong value"); | |
| } | |
| (void)input_scale; | |
| (void)output; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| W4RunTimings timings; | |
| const auto begin = Clock::now(); | |
| RunWorkers(); | |
| const auto executed = Clock::now(); | |
| Gather(input_scale, output); | |
| const auto end = Clock::now(); | |
| timings.npu_ms = Milliseconds(begin, executed); | |
| timings.gather_ms = Milliseconds(executed, end); | |
| timings.total_ms = Milliseconds(begin, end); | |
| timings.input_scale = input_scale; | |
| return timings; | |
| } | |
| void ShareInputFrom(Impl & owner) { | |
| if (this == &owner) return; | |
| if (config.k != owner.config.k || config.k_splits != owner.config.k_splits) { | |
| throw std::invalid_argument("shared W4 inputs require matching K partitions"); | |
| } | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| if (segments.size() != owner.segments.size()) { | |
| throw std::invalid_argument("shared W4 inputs require matching K partitions"); | |
| } | |
| for (std::size_t index = 0; index < segments.size(); ++index) { | |
| auto & target = segments[index]; | |
| auto & source = owner.segments[index]; | |
| if (!source.owns_a || source.a == nullptr || target.a == nullptr || | |
| source.attributes.A.size != target.attributes.A.size || | |
| source.offset_k != target.offset_k || source.size_k != target.size_k) { | |
| throw std::runtime_error("shared W4 input memory is incompatible"); | |
| } | |
| if (target.owns_a) { | |
| CheckRknn( | |
| rknn_destroy_mem(target.context, target.a), | |
| "destroy private W4 input before sharing"); | |
| } | |
| target.a = source.a; | |
| target.owns_a = false; | |
| CheckRknn( | |
| rknn_matmul_set_io_mem(target.context, target.a, &target.attributes.A), | |
| "bind shared W4 input"); | |
| } | |
| } | |
| void SetSingleCore(int core) { | |
| if (config.cores.size() != 1 || core < 0 || core > 2) { | |
| throw std::invalid_argument("SetSingleCore requires one core in [0, 2]"); | |
| } | |
| if (config.cores[0] == core) return; | |
| throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); | |
| for (auto & segment : segments) { | |
| CheckRknn( | |
| rknn_matmul_set_core_mask( | |
| segment.context, static_cast<rknn_core_mask>(1U << core)), | |
| "rknn_matmul_set_core_mask W4 rebind"); | |
| segment.core = core; | |
| } | |
| for (auto & workspace : batch_workspaces) { | |
| if (!workspace) continue; | |
| for (auto & segment : workspace->segments) { | |
| CheckRknn( | |
| rknn_matmul_set_core_mask( | |
| segment.context, static_cast<rknn_core_mask>(1U << core)), | |
| "rknn_matmul_set_core_mask W4 batch rebind"); | |
| segment.core = core; | |
| } | |
| } | |
| config.cores[0] = core; | |
| } | |
| }; | |
| DynamicW4Linear::DynamicW4Linear( | |
| W4LinearConfig config, | |
| std::span<const std::byte> packed_weights, | |
| std::span<const float> weight_scales, | |
| std::span<const std::int32_t> correction) | |
| : impl_(std::make_unique<Impl>( | |
| std::move(config), packed_weights, weight_scales, correction)) {} | |
| DynamicW4Linear::~DynamicW4Linear() = default; | |
| DynamicW4Linear::DynamicW4Linear(DynamicW4Linear &&) noexcept = default; | |
| DynamicW4Linear & DynamicW4Linear::operator=(DynamicW4Linear &&) noexcept = default; | |
| W4RunTimings DynamicW4Linear::Run(std::span<const float> input, std::span<float> output) { | |
| return impl_->Run(input, output); | |
| } | |
| W4RunTimings DynamicW4Linear::RunBatch( | |
| std::span<const float> input, | |
| std::size_t rows, | |
| std::span<float> output) { | |
| return impl_->RunBatch(input, rows, *impl_, output); | |
| } | |
| W4RunTimings DynamicW4Linear::RunBatchWithWeights( | |
| std::span<const float> input, | |
| std::size_t rows, | |
| const DynamicW4Linear & weights, | |
| std::span<float> output) { | |
| return impl_->RunBatch(input, rows, *weights.impl_, output); | |
| } | |
| W4RunTimings DynamicW4Linear::RunBatch32( | |
| std::span<const float> input, std::span<float> output) { | |
| return impl_->RunBatch32(input, output); | |
| } | |
| W4RunTimings DynamicW4Linear::RunBatchQuantizedRows( | |
| std::span<const std::int8_t> input, | |
| std::span<const float> input_scales, | |
| std::span<const std::size_t> row_indices, | |
| const DynamicW4Linear & weights, | |
| std::span<float> output) { | |
| return impl_->RunBatch({}, row_indices.size(), *weights.impl_, output, | |
| input, input_scales, row_indices); | |
| } | |
| void DynamicW4Linear::PrepareBatch(std::size_t rows, bool indexed_input) { | |
| impl_->PrepareBatch(rows, indexed_input); | |
| } | |
| float DynamicW4Linear::PrepareInput(std::span<const float> input) { | |
| return impl_->PrepareInput(input); | |
| } | |
| W4RunTimings DynamicW4Linear::RunPrepared(float input_scale, std::span<float> output) { | |
| return impl_->RunPrepared(input_scale, output); | |
| } | |
| void DynamicW4Linear::ShareInputFrom(DynamicW4Linear & owner) { | |
| impl_->ShareInputFrom(*owner.impl_); | |
| } | |
| void DynamicW4Linear::SetSingleCore(int core) { impl_->SetSingleCore(core); } | |
| const W4LinearConfig & DynamicW4Linear::config() const noexcept { return impl_->config; } | |
| std::size_t DynamicW4Linear::resident_weight_bytes() const noexcept { return impl_->resident_bytes; } | |
| } // namespace ling3 | |