#include "ling3/w4_linear.h" #include "core_workers.h" #include "ling3/quantization.h" #include #include #include #include #include #include #include #include #include #include #if defined(__aarch64__) #include #endif #if LING3_WITH_RKNN #include #include #endif namespace ling3 { namespace { using Clock = std::chrono::steady_clock; #if LING3_WITH_RKNN double Milliseconds(Clock::time_point begin, Clock::time_point end) { return std::chrono::duration(end - begin).count(); } #endif std::vector> 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> 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; } #if LING3_WITH_RKNN 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(value) & 0x0FU; if ((index & 1U) == 0) { output[index / 2] = static_cast(nibble << 4U); } else { output[index / 2] = static_cast(output[index / 2] | nibble); } } #if defined(__aarch64__) 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)); } #endif 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; #if defined(__aarch64__) for (; column + 16 <= count; column += 16) { PackActivation16( vld1q_s8(input + column), high_output + column / 2, low_output + column / 2); } #endif 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(code) & 0x0FU; }; high_output[column / 2] = static_cast( (nibble(0, true) << 4U) | nibble(1, true)); low_output[column / 2] = static_cast( (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( (static_cast(input[column]) & 0x0FU) << 4U | (static_cast(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 segments; std::vector quantized; std::vector accumulator; std::vector 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); } }; #endif } // namespace struct DynamicW4Linear::Impl { W4LinearConfig config; std::vector scales; std::vector correction; std::vector quantized; std::vector accumulator; std::size_t resident_bytes = 0; bool prefill_w4a4 = false; #if LING3_WITH_RKNN std::vector segments; // Reverse destruction order: batch views, backing memory, source contexts. std::vector> batch_memory; std::array, 8> batch_workspaces; #endif Impl( W4LinearConfig value, std::span weights, std::span weight_scales, std::span 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); #if LING3_WITH_RKNN Initialize(weights); #else (void)weights; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #endif } ~Impl() = default; void Validate(std::span 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(config.cores.size()), 64); const auto expected = static_cast(config.k) * config.n / 2; if (weights.size() != expected || scales.size() != static_cast(config.n) || correction.size() != static_cast(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"); } } } #if LING3_WITH_RKNN void Initialize(std::span weights) { const auto k_ranges = SplitAligned(config.k, config.k_splits, 32); const auto n_ranges = SplitAligned(config.n, static_cast(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(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 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(1U << segment.core)), "rknn_matmul_set_core_mask W4"); const auto expected_a = static_cast(segment.size_k); const auto expected_b = static_cast(segment.size_k) * segment.size_n / 2; const auto expected_c = static_cast(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(segment.attributes.B.dims[2]); const int sub_k = static_cast(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(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(segment.offset_k + local_k) * config.n + segment.offset_n + local_n) : static_cast(0); const auto output_index = ((static_cast(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(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> 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(); 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(rows) * config.k); return *stored; } auto workspace = std::make_unique(); workspace->rows = rows; workspace->w4a4 = prefill_w4a4 && rows >= 2; if (!indexed_input) workspace->quantized.resize(static_cast(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(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(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(matrix_rows) * segment.size_k / 2; const auto expected_c = static_cast(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 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(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(first->a->virt_addr); for (int block = 0; block < first->size_k / sub_k; ++block) { const auto high_index = static_cast(block * 2) * sub_k; const auto low_index = static_cast(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 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(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 quantized, std::span 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(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(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( block * packed_rows + (workspace.w4a4 ? row : 2 * row)) * sub_k; if (static_cast(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 & 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 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(segment.c->virt_addr); const int sub_n = static_cast(segment.attributes.C.dims[2]); const int n_blocks = static_cast(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(block * 2) * sub_n; const auto * low = source + static_cast(block * 2 + 1) * sub_n; const int columns = std::min(sub_n, segment.size_n - block * sub_n); int column = 0; #if defined(__aarch64__) 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); } #endif for (; column < columns; ++column) { const auto initial = segment.part == 0 ? correction[global_n + column] : destination[column]; destination[column] = initial + 16 * static_cast(high[column]) + static_cast(low[column]); } } } DequantizePerChannel(accumulator, input_scale, scales, output); } void GatherBatch( BatchWorkspace & workspace, std::size_t active_rows, const Impl & source_weights, std::span 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(segment.c->virt_addr); const int sub_n = static_cast(segment.attributes.C.dims[2]); const int n_blocks = static_cast(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(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; #if defined(__aarch64__) 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); } } #endif for (; column < columns; ++column) { std::int32_t value; if (workspace.w4a4) { const auto initial = segment.part == 0 ? 0 : destination[column]; value = initial + static_cast(high[column]); } else { const auto initial = segment.part == 0 ? source_weights.correction[global_n + column] : destination[column]; value = initial + 16 * static_cast(high[column]) + static_cast(low[column]); } if (final_part) { result[column] = static_cast(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 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); } } #endif W4RunTimings Run(std::span input, std::span output) { if (input.size() != static_cast(config.k) || output.size() != static_cast(config.n)) { throw std::invalid_argument("W4 input or output has the wrong size"); } #if !LING3_WITH_RKNN (void)input; (void)output; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else 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; #endif } W4RunTimings RunBatch( std::span input, std::size_t rows, const Impl & source_weights, std::span output, std::span quantized = {}, std::span input_scales = {}, std::span 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"); } } } #if !LING3_WITH_RKNN (void)input; (void)rows; (void)source_weights; (void)output; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else 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; #endif } W4RunTimings RunBatch32(std::span input, std::span 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]"); } #if !LING3_WITH_RKNN (void)rows; (void)indexed_input; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else 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); #endif } float PrepareInput(std::span input) { if (input.size() != static_cast(config.k)) { throw std::invalid_argument("W4 input has the wrong size"); } #if !LING3_WITH_RKNN (void)input; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else float input_scale = 1.0F; StageInput(input, input_scale); SyncInputs(); return input_scale; #endif } W4RunTimings RunPrepared(float input_scale, std::span output) { if (output.size() != static_cast(config.n) || !(input_scale > 0.0F) || !std::isfinite(input_scale)) { throw std::invalid_argument("prepared W4 scale or output has the wrong value"); } #if !LING3_WITH_RKNN (void)input_scale; (void)output; throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else 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; #endif } 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"); } #if !LING3_WITH_RKNN throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else 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"); } #endif } 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; #if !LING3_WITH_RKNN throw std::runtime_error("DynamicW4Linear requires a build with RKNN support"); #else for (auto & segment : segments) { CheckRknn( rknn_matmul_set_core_mask( segment.context, static_cast(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(1U << core)), "rknn_matmul_set_core_mask W4 batch rebind"); segment.core = core; } } config.cores[0] = core; #endif } }; DynamicW4Linear::DynamicW4Linear( W4LinearConfig config, std::span packed_weights, std::span weight_scales, std::span correction) : impl_(std::make_unique( 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 input, std::span output) { return impl_->Run(input, output); } W4RunTimings DynamicW4Linear::RunBatch( std::span input, std::size_t rows, std::span output) { return impl_->RunBatch(input, rows, *impl_, output); } W4RunTimings DynamicW4Linear::RunBatchWithWeights( std::span input, std::size_t rows, const DynamicW4Linear & weights, std::span output) { return impl_->RunBatch(input, rows, *weights.impl_, output); } W4RunTimings DynamicW4Linear::RunBatch32( std::span input, std::span output) { return impl_->RunBatch32(input, output); } W4RunTimings DynamicW4Linear::RunBatchQuantizedRows( std::span input, std::span input_scales, std::span row_indices, const DynamicW4Linear & weights, std::span 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 input) { return impl_->PrepareInput(input); } W4RunTimings DynamicW4Linear::RunPrepared(float input_scale, std::span 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