Ling-3.0-tiny-RKNN / src /w4_linear.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
48.8 kB
#include "ling3/w4_linear.h"
#include "core_workers.h"
#include "ling3/quantization.h"
#include <algorithm>
#include <array>
#include <chrono>
#include <cmath>
#include <cstdlib>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <string>
#include <utility>
#if defined(__aarch64__)
#include <arm_neon.h>
#endif
#if LING3_WITH_RKNN
#include <rknn_api.h>
#include <rknn_matmul_api.h>
#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<double, std::milli>(end - begin).count();
}
#endif
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;
}
#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<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);
}
}
#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<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);
}
};
#endif
} // 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;
#if LING3_WITH_RKNN
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;
#endif
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);
#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<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");
}
}
}
#if LING3_WITH_RKNN
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;
#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<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;
#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<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);
}
}
#endif
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");
}
#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<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");
}
}
}
#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<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]");
}
#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<const float> input) {
if (input.size() != static_cast<std::size_t>(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<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");
}
#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<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;
#endif
}
};
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