Ling-3.0-tiny-RKNN / src /decoder.cpp
Sariel00's picture
Keep MTP opt-in: default-off build, isolated experimental package and measurements
3706f1c verified
Raw History Blame Contribute Delete
94.3 kB
#include "ling3/decoder.h"
#include "core_workers.h"
#include "numeric_trace.h"
#include "mla_npu.h"
#include "ling3/cpu_kernels.h"
#include "ling3/gdn_step.h"
#include "ling3/quantization.h"
#include "ling3/router.h"
#include "ling3/w4_linear.h"
#include "ling3/linear.h"
#include <algorithm>
#include <array>
#include <charconv>
#include <chrono>
#include <cmath>
#include <cstddef>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <exception>
#include <cstdlib>
#include <limits>
#include <memory>
#include <span>
#include <stdexcept>
#include <string>
#include <thread>
#include <utility>
#include <vector>
#if defined(__aarch64__)
#include <arm_neon.h>
#endif
namespace ling3 {
namespace {
using Clock = std::chrono::steady_clock;
constexpr int kHidden = 1536;
constexpr int kHeads = 16;
constexpr int kHeadDimension = 128;
constexpr int kKdaWidth = kHeads * kHeadDimension;
constexpr int kDenseWidth = 4608;
constexpr int kExpertWidth = 512;
constexpr int kMlaQueryWidth = 192;
constexpr int kMlaValueWidth = 128;
constexpr int kMlaQueryRank = 256;
constexpr int kMlaKvRank = 512;
constexpr int kMlaRotaryWidth = 64;
constexpr int kMlaNopeWidth = 128;
constexpr std::size_t kMaxBatch = 128;
constexpr float kEpsilon = 1.0e-6F;
double Milliseconds(Clock::time_point begin, Clock::time_point end) {
return std::chrono::duration<double, std::milli>(end - begin).count();
}
std::size_t BatchBucketRows(std::size_t rows) {
for (const std::size_t candidate : {1, 2, 4, 8, 16, 32, 64, 128}) {
if (rows <= candidate) return candidate;
}
throw std::invalid_argument("batch supports at most 128 rows");
}
std::size_t ExpertBatchCost(std::size_t rows) {
// Both expert projections execute at the rounded RKNN shape. The small
// fixed term accounts for rebinding weights and launching two matmuls;
// the active-row term covers CPU gather, quantization and dequantization.
return BatchBucketRows(rows) + rows + 4;
}
template <typename T>
std::span<const T> Typed(const TensorView & tensor, DataType type) {
if (tensor.entry->dtype != static_cast<std::uint32_t>(type) ||
tensor.entry->data_bytes % sizeof(T) != 0) {
throw std::runtime_error(std::string(tensor.name) + " has an incompatible dtype");
}
return {
reinterpret_cast<const T *>(tensor.data),
static_cast<std::size_t>(tensor.entry->data_bytes / sizeof(T)),
};
}
std::span<const std::byte> Blob(const TensorView & tensor, TensorRole role) {
if (tensor.entry->role != static_cast<std::uint32_t>(role)) {
throw std::runtime_error(std::string(tensor.name) + " has an incompatible role");
}
return {tensor.data, static_cast<std::size_t>(tensor.entry->data_bytes)};
}
std::vector<float> DecodeBf16(const TensorView & tensor, std::size_t expected) {
const auto input = Typed<std::uint16_t>(tensor, DataType::kBFloat16);
if (input.size() != expected) {
throw std::runtime_error(std::string(tensor.name) + " has an incompatible shape");
}
std::vector<float> output(expected);
for (std::size_t index = 0; index < expected; ++index) {
output[index] = BFloat16ToFloat(input[index]);
}
return output;
}
std::vector<float> DecodeFloat(
const TensorView & tensor,
std::size_t expected) {
if (tensor.entry->dtype == static_cast<std::uint32_t>(DataType::kFloat32)) {
const auto input = Typed<float>(tensor, DataType::kFloat32);
if (input.size() != expected) throw std::runtime_error("FP32 tensor shape mismatch");
return {input.begin(), input.end()};
}
return DecodeBf16(tensor, expected);
}
std::unique_ptr<Linear> MakeLinear(
const ModelPackage & package,
const std::string & base,
std::vector<int> cores = {0, 1, 2}) {
const auto & weight = package.tensor(base + ".weight");
const bool mixed_bf16 = (package.header().flags & kPackageMixedW4W8) &&
weight.entry->dtype == static_cast<std::uint32_t>(DataType::kBFloat16) &&
weight.entry->layout == static_cast<std::uint32_t>(TensorLayout::kRowMajor);
if (weight.entry->rank != 2 || (!(package.header().flags & kPackageOfficialInt4) && !mixed_bf16 && (
weight.entry->dtype != static_cast<std::uint32_t>(DataType::kInt4Low) ||
weight.entry->layout != static_cast<std::uint32_t>(TensorLayout::kPackedInt4Low)))) {
throw std::runtime_error(base + " is not a packed W4 linear");
}
int iommu_domain_id = 0;
constexpr std::string_view layer_prefix = "model.layers.";
if (base.starts_with(layer_prefix)) {
int layer = -1;
const auto begin = base.data() + layer_prefix.size();
const auto end = base.data() + base.size();
const auto [parsed_end, error] = std::from_chars(begin, end, layer);
#if LING3_EXPERIMENTAL_MTP
constexpr int maximum_layer = 24;
#else
constexpr int maximum_layer = 23;
#endif
if (error != std::errc {} || parsed_end == begin || layer < 0 || layer > maximum_layer) {
throw std::runtime_error("cannot derive W4 IOMMU domain from " + base);
}
#if LING3_EXPERIMENTAL_MTP
iommu_domain_id = layer == 24 ? 15 : 2 + layer / 2;
#else
iommu_domain_id = 2 + layer / 2;
#endif
} else if (base == "lm_head") {
iommu_domain_id = 14;
}
auto linear = std::make_unique<Linear>(
package, base,
W4LinearConfig {
static_cast<int>(weight.entry->dims[0]),
static_cast<int>(weight.entry->dims[1]),
static_cast<int>(weight.entry->flags),
std::move(cores),
iommu_domain_id,
});
return linear;
}
void NormalizeHeads(std::span<float> values) {
for (int head = 0; head < kHeads; ++head) {
const int begin = head * kHeadDimension;
float sum = kEpsilon;
for (int index = 0; index < kHeadDimension; ++index) {
const float value = values[begin + index];
sum += value * value;
}
const float inverse = 1.0F / std::sqrt(sum);
for (int index = 0; index < kHeadDimension; ++index) values[begin + index] *= inverse;
}
}
void CausalConvSilu(
std::span<const float> input,
std::span<const float> weight,
std::span<float> state,
std::span<float> output) {
for (int channel = 0; channel < kKdaWidth; ++channel) {
auto * history = state.data() + static_cast<std::size_t>(channel) * 3;
const auto * kernel = weight.data() + static_cast<std::size_t>(channel) * 4;
const float value = history[0] * kernel[0] + history[1] * kernel[1] +
history[2] * kernel[2] + input[channel] * kernel[3];
history[0] = history[1];
history[1] = history[2];
history[2] = input[channel];
output[channel] = value / (1.0F + std::exp(-value));
}
}
// Per-decoder batch scratch. Layers execute synchronously and may borrow these
// buffers until RunBatch returns. Recurrent state and KV remain layer-owned.
struct KdaBatchScratch {
std::vector<float> projected, q, k, v, decay, beta, recurrence, gated;
KdaBatchScratch()
: projected(kMaxBatch * 10304), q(kMaxBatch * kKdaWidth),
k(kMaxBatch * kKdaWidth), v(kMaxBatch * kKdaWidth),
decay(kMaxBatch * kKdaWidth), beta(kMaxBatch * kHeads),
recurrence(kMaxBatch * kKdaWidth), gated(kMaxBatch * kKdaWidth) {}
};
struct MlaBatchScratch {
MlaNpu npu;
std::vector<float> projected, q_rank, kv_rank, q_all, kv_all, attention, rotated_key;
MlaBatchScratch()
: projected(kMaxBatch * 896), q_rank(kMaxBatch * kMlaQueryRank),
kv_rank(kMaxBatch * kMlaKvRank), q_all(kMaxBatch * kHeads * kMlaQueryWidth),
kv_all(kMaxBatch * kHeads * 256), attention(kMaxBatch * kHeads * kMlaValueWidth),
rotated_key(kMaxBatch * kMlaRotaryWidth) {}
};
struct SparseBatchScratch {
std::vector<float> shared_projected, shared_hidden, shared_output;
std::array<std::vector<float>, 3> lane_input, lane_projected, lane_hidden, lane_output;
std::vector<std::int8_t> quantized;
std::vector<float> contributions;
SparseBatchScratch()
: shared_projected(kMaxBatch * 2 * kExpertWidth),
shared_hidden(kMaxBatch * kExpertWidth), shared_output(kMaxBatch * kHidden) {}
};
struct LayerBatchScratch {
std::vector<float> normalized, attention_output, ffn_output;
LayerBatchScratch()
: normalized(kMaxBatch * kHidden), attention_output(kMaxBatch * kHidden),
ffn_output(kMaxBatch * kHidden) {}
};
struct DecoderScratch {
KdaBatchScratch kda;
MlaBatchScratch mla;
SparseBatchScratch sparse;
LayerBatchScratch layer;
};
using AttentionCheckpoint = AttentionState;
class Attention {
public:
virtual ~Attention() = default;
virtual void Reset() = 0;
virtual AttentionCheckpoint SaveCheckpoint() { return {}; }
virtual void RestoreCheckpoint(const AttentionCheckpoint &) {}
virtual AttentionState SaveState(std::size_t) { return SaveCheckpoint(); }
virtual void Run(std::span<const float> input, std::size_t position, std::span<float> output) = 0;
virtual void RunBatch(
std::span<const float> input,
std::size_t rows,
std::size_t position,
std::span<float> output) = 0;
virtual void PrepareBatch(std::size_t rows) = 0;
};
class KdaAttention final : public Attention {
public:
KdaAttention(
const ModelPackage & package,
int layer,
std::span<const std::byte> heads6,
std::span<const std::byte> heads5,
KdaBatchScratch & scratch)
: prefix_("model.layers." + std::to_string(layer) + ".attention"),
projection_(MakeLinear(package, prefix_ + ".qkvfgb")),
output_projection_(MakeLinear(package, prefix_ + ".o_proj")),
q_conv_(DecodeBf16(package.tensor(prefix_ + ".q_conv1d.weight"), kKdaWidth * 4)),
k_conv_(DecodeBf16(package.tensor(prefix_ + ".k_conv1d.weight"), kKdaWidth * 4)),
v_conv_(DecodeBf16(package.tensor(prefix_ + ".v_conv1d.weight"), kKdaWidth * 4)),
a_log_(DecodeFloat(package.tensor(prefix_ + ".A_log"), kHeads)),
dt_bias_(DecodeFloat(package.tensor(prefix_ + ".dt_bias"), kKdaWidth)),
output_norm_(DecodeBf16(package.tensor(prefix_ + ".o_norm.weight"), kHeadDimension)),
gdn_(heads6, heads5),
projected_(10304),
q_(kKdaWidth),
k_(kKdaWidth),
v_(kKdaWidth),
decay_(kKdaWidth),
beta_(kHeads),
recurrence_(kKdaWidth),
gated_(kKdaWidth),
batch_projected_(scratch.projected), batch_q_(scratch.q),
batch_k_(scratch.k), batch_v_(scratch.v), batch_decay_(scratch.decay),
batch_beta_(scratch.beta), batch_recurrence_(scratch.recurrence),
batch_gated_(scratch.gated) {
for (auto & state : conv_state_) state.assign(kKdaWidth * 3, 0.0F);
}
void Reset() override {
for (auto & state : conv_state_) std::fill(state.begin(), state.end(), 0.0F);
gdn_.Reset();
}
AttentionCheckpoint SaveCheckpoint() override { return {gdn_.SaveState(), conv_state_, {}, {}}; }
void RestoreCheckpoint(const AttentionCheckpoint & checkpoint) override {
for (const auto & values : checkpoint.conv)
if (values.size() != kKdaWidth * 3) throw std::invalid_argument("invalid convolution checkpoint");
gdn_.RestoreState(checkpoint.gdn);
conv_state_ = checkpoint.conv;
}
void Run(std::span<const float> input, std::size_t, std::span<float> output) override {
projection_->Run(input, projected_);
CausalConvSilu(
std::span<const float>(projected_).subspan(0, kKdaWidth),
q_conv_, conv_state_[0], q_);
CausalConvSilu(
std::span<const float>(projected_).subspan(kKdaWidth, kKdaWidth),
k_conv_, conv_state_[1], k_);
CausalConvSilu(
std::span<const float>(projected_).subspan(2 * kKdaWidth, kKdaWidth),
v_conv_, conv_state_[2], v_);
NormalizeHeads(q_);
NormalizeHeads(k_);
const auto f = std::span<const float>(projected_).subspan(3 * kKdaWidth, kKdaWidth);
const auto gate = std::span<const float>(projected_).subspan(4 * kKdaWidth, kKdaWidth);
const auto beta_logits = std::span<const float>(projected_).subspan(5 * kKdaWidth, kHeads);
for (int head = 0; head < kHeads; ++head) {
const float a = std::exp(a_log_[head]);
beta_[head] = 1.0F / (1.0F + std::exp(-beta_logits[head]));
for (int index = 0; index < kHeadDimension; ++index) {
const int offset = head * kHeadDimension + index;
decay_[offset] = -5.0F /
(1.0F + std::exp(-a * (f[offset] + dt_bias_[offset])));
}
}
gdn_.Run(q_, k_, v_, decay_, beta_, recurrence_);
for (int head = 0; head < kHeads; ++head) {
const int begin = head * kHeadDimension;
float sum = 0.0F;
for (int index = 0; index < kHeadDimension; ++index) {
const float value = recurrence_[begin + index];
sum += value * value;
}
const float inverse = 1.0F /
std::sqrt(sum / static_cast<float>(kHeadDimension) + kEpsilon);
for (int index = 0; index < kHeadDimension; ++index) {
const int offset = begin + index;
const float sigmoid = 1.0F / (1.0F + std::exp(-gate[offset]));
gated_[offset] = recurrence_[offset] * inverse * output_norm_[index] * sigmoid;
}
}
output_projection_->Run(gated_, output);
}
void RunBatch(
std::span<const float> input,
std::size_t rows,
std::size_t,
std::span<float> output) override {
const bool trace_batch = std::getenv("LING3_TRACE_BATCH") != nullptr;
const auto batch_begin = Clock::now();
if (rows < 1 || rows > kMaxBatch || input.size() != rows * kHidden ||
output.size() != rows * kHidden) {
throw std::invalid_argument("KDA batch has an incompatible tensor size");
}
const bool cpu_gdn = std::getenv("LING3_GDN_CPU_PREFILL") != nullptr;
if (!cpu_gdn && (!gdn_.has_batch16() || rows % 16 != 0)) {
throw std::runtime_error(
"KDA batch requires CPU GDN or a multiple of 16 with GDN prefill models");
}
const auto projection_timings = projection_->RunBatch(
input, rows, std::span<float>(batch_projected_).first(rows * 10304));
const auto projected_at = Clock::now();
constexpr std::array<int, 4> cpu_workers {0, 1, 2, 3};
CoreWorkers::Instance().Run(cpu_workers, [this, rows](int worker) {
for (int head = worker; head < kHeads; head += 4) {
const int head_begin = head * kHeadDimension;
const float a = std::exp(a_log_[head]);
for (std::size_t row = 0; row < rows; ++row) {
for (int index = 0; index < kHeadDimension; ++index) {
const int channel = head_begin + index;
for (int stream = 0; stream < 3; ++stream) {
auto * history = conv_state_[stream].data() +
static_cast<std::size_t>(channel) * 3;
const auto & weights = stream == 0 ? q_conv_ :
(stream == 1 ? k_conv_ : v_conv_);
const auto * kernel = weights.data() +
static_cast<std::size_t>(channel) * 4;
auto & destination = stream == 0 ? batch_q_ :
(stream == 1 ? batch_k_ : batch_v_);
const float input_value = batch_projected_[
row * 10304 + stream * kKdaWidth + channel];
const float value = history[0] * kernel[0] +
history[1] * kernel[1] + history[2] * kernel[2] +
input_value * kernel[3];
history[0] = history[1];
history[1] = history[2];
history[2] = input_value;
destination[row * kKdaWidth + channel] = value;
}
}
auto * q = batch_q_.data() + row * kKdaWidth + head_begin;
auto * k = batch_k_.data() + row * kKdaWidth + head_begin;
auto * v = batch_v_.data() + row * kKdaWidth + head_begin;
Silu(q, q, kHeadDimension);
Silu(k, k, kHeadDimension);
Silu(v, v, kHeadDimension);
float q_sum = kEpsilon;
float k_sum = kEpsilon;
for (int index = 0; index < kHeadDimension; ++index) {
q_sum += q[index] * q[index];
k_sum += k[index] * k[index];
}
const float q_inverse = 1.0F / std::sqrt(q_sum);
const float k_inverse = 1.0F / std::sqrt(k_sum);
for (int index = 0; index < kHeadDimension; ++index) {
q[index] *= q_inverse;
k[index] *= k_inverse;
const int channel = head_begin + index;
const float f = batch_projected_[
row * 10304 + 3 * kKdaWidth + channel];
batch_decay_[row * kKdaWidth + channel] = -5.0F /
(1.0F + std::exp(-a * (f + dt_bias_[channel])));
}
const float beta_logit = batch_projected_[
row * 10304 + 5 * kKdaWidth + head];
batch_beta_[row * kHeads + head] =
1.0F / (1.0F + std::exp(-beta_logit));
}
}
});
const auto preprocessed_at = Clock::now();
GdnRunTimings gdn_timings;
if (cpu_gdn) {
gdn_timings = gdn_.RunBatchCpu(
std::span<const float>(batch_q_).first(rows * kKdaWidth),
std::span<const float>(batch_k_).first(rows * kKdaWidth),
std::span<const float>(batch_v_).first(rows * kKdaWidth),
std::span<const float>(batch_decay_).first(rows * kKdaWidth),
std::span<const float>(batch_beta_).first(rows * kHeads),
std::span<float>(batch_recurrence_).first(rows * kKdaWidth));
} else {
constexpr std::size_t chunk_vectors = 16 * kKdaWidth;
constexpr std::size_t chunk_betas = 16 * kHeads;
for (std::size_t chunk = 0; chunk < rows / 16; ++chunk) {
const auto chunk_timings = gdn_.RunBatch16(
std::span<const float>(batch_q_).subspan(chunk * chunk_vectors, chunk_vectors),
std::span<const float>(batch_k_).subspan(chunk * chunk_vectors, chunk_vectors),
std::span<const float>(batch_v_).subspan(chunk * chunk_vectors, chunk_vectors),
std::span<const float>(batch_decay_).subspan(chunk * chunk_vectors, chunk_vectors),
std::span<const float>(batch_beta_).subspan(chunk * chunk_betas, chunk_betas),
std::span<float>(batch_recurrence_).subspan(
chunk * chunk_vectors, chunk_vectors));
gdn_timings.stage_ms += chunk_timings.stage_ms;
gdn_timings.npu_ms += chunk_timings.npu_ms;
gdn_timings.collect_ms += chunk_timings.collect_ms;
gdn_timings.total_ms += chunk_timings.total_ms;
}
}
const auto gdn_at = Clock::now();
CoreWorkers::Instance().Run(cpu_workers, [this, rows](int worker) {
for (int head = worker; head < kHeads; head += 4) {
const int begin = head * kHeadDimension;
for (std::size_t row = 0; row < rows; ++row) {
const auto * recurrence =
batch_recurrence_.data() + row * kKdaWidth;
const auto * gate =
batch_projected_.data() + row * 10304 + 4 * kKdaWidth;
auto * gated = batch_gated_.data() + row * kKdaWidth;
float sum = 0.0F;
for (int index = 0; index < kHeadDimension; ++index) {
const float value = recurrence[begin + index];
sum += value * value;
}
const float inverse = 1.0F /
std::sqrt(sum / static_cast<float>(kHeadDimension) + kEpsilon);
for (int index = 0; index < kHeadDimension; ++index) {
const int offset = begin + index;
const float sigmoid = 1.0F / (1.0F + std::exp(-gate[offset]));
gated[offset] =
recurrence[offset] * inverse * output_norm_[index] * sigmoid;
}
}
}
});
const auto gated_at = Clock::now();
const auto output_timings = output_projection_->RunBatch(
std::span<const float>(batch_gated_).first(rows * kKdaWidth), rows, output);
if (trace_batch) {
std::fprintf(
stderr,
" kda_batch rows=%zu projection=%.3f(prep=%.3f,npu=%.3f,gather=%.3f) "
"preprocess=%.3f gdn=%.3f(stage=%.3f,run=%.3f,commit=%.3f) gate=%.3f "
"output=%.3f(prep=%.3f,npu=%.3f,gather=%.3f) total=%.3f\n",
rows, projection_timings.total_ms, projection_timings.quantize_pack_ms,
projection_timings.npu_ms, projection_timings.gather_ms,
Milliseconds(projected_at, preprocessed_at), gdn_timings.total_ms,
gdn_timings.stage_ms, gdn_timings.npu_ms, gdn_timings.collect_ms,
Milliseconds(gdn_at, gated_at), output_timings.total_ms,
output_timings.quantize_pack_ms, output_timings.npu_ms,
output_timings.gather_ms, Milliseconds(batch_begin, Clock::now()));
}
}
void PrepareBatch(std::size_t rows) override {
if (rows < 1 || rows > kMaxBatch) {
throw std::invalid_argument("KDA batch rows must be in [1, 128]");
}
if (std::getenv("LING3_GDN_CPU_PREFILL") == nullptr && !gdn_.has_batch16()) {
throw std::runtime_error("KDA batch requires CPU GDN or GDN prefill models");
}
projection_->PrepareBatch(rows);
output_projection_->PrepareBatch(rows);
}
private:
std::string prefix_;
std::unique_ptr<Linear> projection_;
std::unique_ptr<Linear> output_projection_;
std::vector<float> q_conv_, k_conv_, v_conv_, a_log_, dt_bias_, output_norm_;
GdnStep gdn_;
std::array<std::vector<float>, 3> conv_state_;
std::vector<float> projected_, q_, k_, v_, decay_, beta_, recurrence_, gated_;
std::vector<float> &batch_projected_, &batch_q_, &batch_k_, &batch_v_, &batch_decay_;
std::vector<float> &batch_beta_, &batch_recurrence_, &batch_gated_;
};
float MlaDot(const float * query, const std::uint16_t * key, int count) {
#if defined(__aarch64__)
static const bool simd = std::getenv("LING3_MLA_SIMD") != nullptr;
if (simd) {
auto sum0 = vdupq_n_f32(0.0F), sum1 = vdupq_n_f32(0.0F);
for (int i = 0; i < count; i += 8) {
const auto packed = vld1q_u16(key + i);
const auto lo = vreinterpretq_f32_u32(vshll_n_u16(vget_low_u16(packed), 16));
const auto hi = vreinterpretq_f32_u32(vshll_n_u16(vget_high_u16(packed), 16));
sum0 = vfmaq_f32(sum0, vld1q_f32(query + i), lo);
sum1 = vfmaq_f32(sum1, vld1q_f32(query + i + 4), hi);
}
return vaddvq_f32(vaddq_f32(sum0, sum1));
}
#endif
float dot = 0.0F;
for (int i = 0; i < count; ++i) dot += query[i] * BFloat16ToFloat(key[i]);
return dot;
}
void MlaAccumulate(const std::uint16_t * value, float probability, float * out, int count) {
#if defined(__aarch64__)
static const bool simd = std::getenv("LING3_MLA_SIMD") != nullptr;
if (simd) {
for (int i = 0; i < count; i += 8) {
const auto packed = vld1q_u16(value + i);
const auto lo = vreinterpretq_f32_u32(vshll_n_u16(vget_low_u16(packed), 16));
const auto hi = vreinterpretq_f32_u32(vshll_n_u16(vget_high_u16(packed), 16));
vst1q_f32(out + i, vfmaq_n_f32(vld1q_f32(out + i), lo, probability));
vst1q_f32(out + i + 4, vfmaq_n_f32(vld1q_f32(out + i + 4), hi, probability));
}
return;
}
#endif
for (int i = 0; i < count; ++i) out[i] += probability * BFloat16ToFloat(value[i]);
}
class MlaAttention final : public Attention {
public:
MlaAttention(const ModelPackage & package, int layer, std::size_t max_context,
MlaBatchScratch & scratch)
: prefix_("model.layers." + std::to_string(layer) + ".attention"),
projection_(MakeLinear(package, prefix_ + ".qkv_gate_a")),
q_projection_(MakeLinear(package, prefix_ + ".q_b_proj")),
kv_projection_(MakeLinear(package, prefix_ + ".kv_b_proj")),
output_projection_(MakeLinear(package, prefix_ + ".o_proj")),
q_norm_(DecodeBf16(package.tensor(prefix_ + ".q_a_layernorm.weight"), kMlaQueryRank)),
kv_norm_(DecodeBf16(package.tensor(prefix_ + ".kv_a_layernorm.weight"), kMlaKvRank)),
max_context_(max_context),
projected_(896), q_rank_(kMlaQueryRank), kv_rank_(kMlaKvRank),
q_all_(kHeads * kMlaQueryWidth), kv_all_(kHeads * 256),
attention_(kHeads * kMlaValueWidth), scores_(max_context),
key_cache_(max_context * kHeads * kMlaQueryWidth),
value_cache_(max_context * kHeads * kMlaValueWidth),
batch_projected_(scratch.projected), batch_q_rank_(scratch.q_rank),
batch_kv_rank_(scratch.kv_rank), batch_q_all_(scratch.q_all),
batch_kv_all_(scratch.kv_all), batch_attention_(scratch.attention),
batch_rotated_key_(scratch.rotated_key), npu_(scratch.npu) {
for (auto & scores : batch_scores_) scores.resize(max_context);
}
void Reset() override {
// The decoder resets position to zero. Every read is bounded by the new
// position, and each key/value is overwritten before it can be read.
// Do not sweep a potentially multi-GB cache on every chat request.
}
AttentionState SaveState(std::size_t position) override {
AttentionState state;
state.keys.assign(key_cache_.begin(), key_cache_.begin()+position*kHeads*kMlaQueryWidth);
state.values.assign(value_cache_.begin(), value_cache_.begin()+position*kHeads*kMlaValueWidth);
return state;
}
void RestoreCheckpoint(const AttentionCheckpoint & state) override {
// Lightweight checkpoints intentionally leave MLA's live prefix in place.
if (state.keys.empty() && state.values.empty()) return;
std::copy(state.keys.begin(), state.keys.end(), key_cache_.begin());
std::copy(state.values.begin(), state.values.end(), value_cache_.begin());
}
void Run(std::span<const float> input, std::size_t position, std::span<float> output) override {
if (position >= max_context_) throw std::runtime_error("MLA cache capacity exceeded");
npu_.CpuCall();
projection_->Run(input, projected_);
RmsNorm(projected_.data(), q_norm_.data(), q_rank_.data(), kMlaQueryRank, kEpsilon);
RmsNorm(
projected_.data() + kMlaQueryRank,
kv_norm_.data(),
kv_rank_.data(),
kMlaKvRank,
kEpsilon);
q_projection_->Run(q_rank_, q_all_);
kv_projection_->Run(kv_rank_, kv_all_);
std::array<float, kMlaRotaryWidth> rotated_key {};
RotateInterleaved(
std::span<const float>(projected_).subspan(
kMlaQueryRank + kMlaKvRank, kMlaRotaryWidth),
position,
rotated_key);
for (int head = 0; head < kHeads; ++head) {
std::array<float, kMlaRotaryWidth> rotated_query {};
RotateInterleaved(
std::span<const float>(q_all_).subspan(
head * kMlaQueryWidth + kMlaNopeWidth, kMlaRotaryWidth),
position,
rotated_query);
auto * key_destination = key_cache_.data() +
(position * kHeads + head) * kMlaQueryWidth;
auto * value_destination = value_cache_.data() +
(position * kHeads + head) * kMlaValueWidth;
for (int index = 0; index < kMlaNopeWidth; ++index) {
key_destination[index] = FloatToBFloat16(kv_all_[head * 256 + index]);
}
for (int index = 0; index < kMlaRotaryWidth; ++index) {
key_destination[kMlaNopeWidth + index] = FloatToBFloat16(rotated_key[index]);
q_all_[head * kMlaQueryWidth + kMlaNopeWidth + index] = rotated_query[index];
}
for (int index = 0; index < kMlaValueWidth; ++index) {
value_destination[index] = FloatToBFloat16(kv_all_[head * 256 + kMlaNopeWidth + index]);
}
}
const float scale = 1.0F / std::sqrt(static_cast<float>(kMlaQueryWidth));
constexpr std::array<int, 4> workers {0, 1, 2, 3};
CoreWorkers::Instance().Run(workers, [&](int worker) {
auto & scores_ = batch_scores_[worker];
for (int head = worker; head < kHeads; head += 4) {
float maximum = -std::numeric_limits<float>::infinity();
const float * query = q_all_.data() + head * kMlaQueryWidth;
for (std::size_t token = 0; token <= position; ++token) {
const auto * key = key_cache_.data() +
(token * kHeads + head) * kMlaQueryWidth;
const float dot = MlaDot(query, key, kMlaQueryWidth);
scores_[token] = dot * scale;
maximum = std::max(maximum, scores_[token]);
}
float denominator = 0.0F;
for (std::size_t token = 0; token <= position; ++token) {
scores_[token] = std::exp(scores_[token] - maximum);
denominator += scores_[token];
}
auto * destination = attention_.data() + head * kMlaValueWidth;
std::fill(destination, destination + kMlaValueWidth, 0.0F);
for (std::size_t token = 0; token <= position; ++token) {
const float probability = scores_[token] / denominator;
const auto * value = value_cache_.data() +
(token * kHeads + head) * kMlaValueWidth;
MlaAccumulate(value, probability, destination, kMlaValueWidth);
}
const float gate = 1.0F /
(1.0F + std::exp(-projected_[kMlaQueryRank + kMlaKvRank + kMlaRotaryWidth + head]));
for (int index = 0; index < kMlaValueWidth; ++index) destination[index] *= gate;
}
});
output_projection_->Run(attention_, output);
}
void RunBatch(
std::span<const float> input,
std::size_t rows,
std::size_t position,
std::span<float> output) override {
constexpr std::size_t q_width = kHeads * kMlaQueryWidth;
constexpr std::size_t kv_width = kHeads * 256;
constexpr std::size_t attention_width = kHeads * kMlaValueWidth;
if (rows < 1 || rows > kMaxBatch || input.size() != rows * kHidden ||
output.size() != rows * kHidden ||
position + rows > max_context_) {
throw std::invalid_argument("MLA batch has an incompatible tensor size or position");
}
projection_->RunBatch(
input, rows, std::span<float>(batch_projected_).first(rows * 896));
for (std::size_t row = 0; row < rows; ++row) {
const auto * projected = batch_projected_.data() + row * 896;
RmsNorm(
projected, q_norm_.data(),
batch_q_rank_.data() + row * kMlaQueryRank, kMlaQueryRank, kEpsilon);
RmsNorm(
projected + kMlaQueryRank, kv_norm_.data(),
batch_kv_rank_.data() + row * kMlaKvRank, kMlaKvRank, kEpsilon);
}
q_projection_->RunBatch(
std::span<const float>(batch_q_rank_).first(rows * kMlaQueryRank), rows,
std::span<float>(batch_q_all_).first(rows * q_width));
kv_projection_->RunBatch(
std::span<const float>(batch_kv_rank_).first(rows * kMlaKvRank), rows,
std::span<float>(batch_kv_all_).first(rows * kv_width));
for (std::size_t row = 0; row < rows; ++row) {
const std::size_t token_position = position + row;
const auto * projected = batch_projected_.data() + row * 896;
RotateInterleaved(
{projected + kMlaQueryRank + kMlaKvRank, kMlaRotaryWidth},
token_position,
std::span<float>(batch_rotated_key_).subspan(
row * kMlaRotaryWidth, kMlaRotaryWidth));
}
constexpr std::array<int, 4> workers {0, 1, 2, 3};
CoreWorkers::Instance().Run(workers, [this, position, rows](int worker) {
for (int head = worker; head < kHeads; head += 4) {
for (std::size_t row = 0; row < rows; ++row) {
const std::size_t token_position = position + row;
auto * q_all = batch_q_all_.data() + row * q_width;
const auto * kv_all = batch_kv_all_.data() + row * kv_width;
const auto * rotated_key =
batch_rotated_key_.data() + row * kMlaRotaryWidth;
std::array<float, kMlaRotaryWidth> rotated_query {};
RotateInterleaved(
{q_all + head * kMlaQueryWidth + kMlaNopeWidth, kMlaRotaryWidth},
token_position,
rotated_query);
auto * key_destination = key_cache_.data() +
(token_position * kHeads + head) * kMlaQueryWidth;
auto * value_destination = value_cache_.data() +
(token_position * kHeads + head) * kMlaValueWidth;
for (int index = 0; index < kMlaNopeWidth; ++index) {
key_destination[index] = FloatToBFloat16(kv_all[head * 256 + index]);
}
for (int index = 0; index < kMlaRotaryWidth; ++index) {
key_destination[kMlaNopeWidth + index] = FloatToBFloat16(rotated_key[index]);
q_all[head * kMlaQueryWidth + kMlaNopeWidth + index] = rotated_query[index];
}
for (int index = 0; index < kMlaValueWidth; ++index) {
value_destination[index] =
FloatToBFloat16(kv_all[head * 256 + kMlaNopeWidth + index]);
}
}
}
});
NumericMlaCapture(prefix_, position, rows,
std::span<const float>(batch_q_all_).first(rows * q_width),
std::span<const std::uint16_t>(key_cache_).first((position + rows) * q_width),
std::span<const std::uint16_t>(value_cache_).first((position + rows) * attention_width));
const bool used_npu = npu_.Run(std::span<const float>(batch_q_all_).first(rows*q_width),
std::span<const std::uint16_t>(key_cache_).first((position+rows)*q_width),
std::span<const std::uint16_t>(value_cache_).first((position+rows)*attention_width),
rows,position+rows,std::span<float>(batch_attention_).first(rows*attention_width));
const float scale = 1.0F / std::sqrt(static_cast<float>(kMlaQueryWidth));
if (!used_npu) {
npu_.CpuCall();
CoreWorkers::Instance().Run(workers, [this, position, rows, scale](int worker) {
auto & scores = batch_scores_[worker];
for (int head = worker; head < kHeads; head += 4) {
for (std::size_t row = 0; row < rows; ++row) {
const std::size_t token_position = position + row;
const auto * query =
batch_q_all_.data() + row * q_width + head * kMlaQueryWidth;
float maximum = -std::numeric_limits<float>::infinity();
for (std::size_t token = 0; token <= token_position; ++token) {
const auto * cached_key = key_cache_.data() +
(token * kHeads + head) * kMlaQueryWidth;
const float dot = MlaDot(query, cached_key, kMlaQueryWidth);
scores[token] = dot * scale;
maximum = std::max(maximum, scores[token]);
}
float denominator = 0.0F;
for (std::size_t token = 0; token <= token_position; ++token) {
scores[token] = std::exp(scores[token] - maximum);
denominator += scores[token];
}
auto * destination = batch_attention_.data() +
row * attention_width + head * kMlaValueWidth;
std::fill(destination, destination + kMlaValueWidth, 0.0F);
for (std::size_t token = 0; token <= token_position; ++token) {
const float probability = scores[token] / denominator;
const auto * cached_value = value_cache_.data() +
(token * kHeads + head) * kMlaValueWidth;
MlaAccumulate(cached_value, probability, destination, kMlaValueWidth);
}
}
}
});
}
for(std::size_t row=0;row<rows;++row)for(int head=0;head<kHeads;++head){
const float gate=1.0F/(1.0F+std::exp(-batch_projected_[row*896+
kMlaQueryRank+kMlaKvRank+kMlaRotaryWidth+head]));
auto * destination=batch_attention_.data()+row*attention_width+head*kMlaValueWidth;
for(int i=0;i<kMlaValueWidth;++i)destination[i]*=gate;
}
output_projection_->RunBatch(
std::span<const float>(batch_attention_).first(rows * attention_width), rows, output);
}
void PrepareBatch(std::size_t rows) override {
projection_->PrepareBatch(rows);
q_projection_->PrepareBatch(rows);
kv_projection_->PrepareBatch(rows);
output_projection_->PrepareBatch(rows);
npu_.Prepare(rows);
}
private:
static void RotateInterleaved(
std::span<const float> input,
std::size_t position,
std::span<float> output) {
std::array<float, kMlaRotaryWidth> reordered {};
for (int index = 0; index < kMlaRotaryWidth / 2; ++index) {
reordered[index] = input[2 * index];
reordered[kMlaRotaryWidth / 2 + index] = input[2 * index + 1];
}
for (int index = 0; index < kMlaRotaryWidth; ++index) {
const int frequency = index % (kMlaRotaryWidth / 2);
const float inverse = std::pow(
6000000.0F,
-2.0F * static_cast<float>(frequency) / static_cast<float>(kMlaRotaryWidth));
const float angle = static_cast<float>(position) * inverse;
const float other = index < kMlaRotaryWidth / 2
? -reordered[index + kMlaRotaryWidth / 2]
: reordered[index - kMlaRotaryWidth / 2];
output[index] = reordered[index] * std::cos(angle) + other * std::sin(angle);
}
}
std::string prefix_;
std::unique_ptr<Linear> projection_, q_projection_, kv_projection_, output_projection_;
std::vector<float> q_norm_, kv_norm_;
std::size_t max_context_;
std::vector<float> projected_, q_rank_, kv_rank_, q_all_, kv_all_, attention_, scores_;
std::vector<std::uint16_t> key_cache_, value_cache_;
std::vector<float> &batch_projected_, &batch_q_rank_, &batch_kv_rank_;
std::vector<float> &batch_q_all_, &batch_kv_all_, &batch_attention_, &batch_rotated_key_;
std::array<std::vector<float>, 4> batch_scores_;
MlaNpu & npu_;
};
class FeedForward {
public:
virtual ~FeedForward() = default;
virtual void Run(std::span<const float> input, std::span<float> output) = 0;
virtual void RunBatch(
std::span<const float> input,
std::size_t rows,
std::span<float> output) = 0;
virtual void PrepareBatch(std::size_t rows) = 0;
};
class DenseFeedForward final : public FeedForward {
public:
DenseFeedForward(const ModelPackage & package, int layer)
: gate_up_(MakeLinear(
package, "model.layers." + std::to_string(layer) + ".mlp.gate_up")),
down_(MakeLinear(
package, "model.layers." + std::to_string(layer) + ".mlp.down_proj")),
projected_(2 * kDenseWidth), hidden_(kDenseWidth),
batch_projected_(kMaxBatch * 2 * kDenseWidth),
batch_hidden_(kMaxBatch * kDenseWidth) {}
void Run(std::span<const float> input, std::span<float> output) override {
gate_up_->Run(input, projected_);
SiluMultiply(projected_.data(), projected_.data() + kDenseWidth, hidden_.data(), kDenseWidth);
down_->Run(hidden_, output);
}
void RunBatch(
std::span<const float> input,
std::size_t rows,
std::span<float> output) override {
if (rows < 1 || rows > kMaxBatch || input.size() != rows * kHidden ||
output.size() != rows * kHidden) {
throw std::invalid_argument("dense FFN batch has an incompatible tensor size");
}
gate_up_->RunBatch(
input, rows, std::span<float>(batch_projected_).first(rows * 2 * kDenseWidth));
for (std::size_t row = 0; row < rows; ++row) {
const auto * projected = batch_projected_.data() + row * 2 * kDenseWidth;
auto * hidden = batch_hidden_.data() + row * kDenseWidth;
SiluMultiply(projected, projected + kDenseWidth, hidden, kDenseWidth);
}
down_->RunBatch(
std::span<const float>(batch_hidden_).first(rows * kDenseWidth), rows, output);
}
void PrepareBatch(std::size_t rows) override {
gate_up_->PrepareBatch(rows);
down_->PrepareBatch(rows);
}
private:
std::unique_ptr<Linear> gate_up_, down_;
std::vector<float> projected_, hidden_;
std::vector<float> batch_projected_, batch_hidden_;
};
class Expert {
public:
Expert(const ModelPackage & package, int layer, int expert, bool all_cores)
: gate_up_(MakeLinear(
package,
"model.layers." + std::to_string(layer) + ".mlp.experts." +
std::to_string(expert) + ".gate_up",
all_cores ? std::vector<int> {0, 1, 2} : std::vector<int> {expert % 3})),
down_(MakeLinear(
package,
"model.layers." + std::to_string(layer) + ".mlp.experts." +
std::to_string(expert) + ".down_proj",
all_cores ? std::vector<int> {0, 1, 2} : std::vector<int> {expert % 3})),
projected_(2 * kExpertWidth), hidden_(kExpertWidth), output_(kHidden),
current_core_(all_cores ? -1 : expert % 3) {}
std::span<const float> Run(
std::span<const float> input,
float prepared_gate_scale = 0.0F) {
if (prepared_gate_scale > 0.0F) {
gate_up_->RunPrepared(prepared_gate_scale, projected_);
} else {
gate_up_->Run(input, projected_);
}
SiluMultiply(projected_.data(), projected_.data() + kExpertWidth, hidden_.data(), kExpertWidth);
down_->Run(hidden_, output_);
return output_;
}
int current_core() const noexcept { return current_core_; }
void SetCore(int core) {
if (core == current_core_) return;
gate_up_->SetSingleCore(core);
down_->SetSingleCore(core);
current_core_ = core;
}
float PrepareGateInput(std::span<const float> input) {
return gate_up_->PrepareInput(input);
}
void ShareGateInputFrom(Expert & owner) {
gate_up_->ShareInputFrom(*owner.gate_up_);
}
void RunGateBatch(
std::span<const float> input,
std::size_t rows,
const Expert & weights,
std::span<float> output) {
gate_up_->RunBatchWithWeights(input, rows, *weights.gate_up_, output);
}
void RunDownBatch(
std::span<const float> input,
std::size_t rows,
const Expert & weights,
std::span<float> output) {
down_->RunBatchWithWeights(input, rows, *weights.down_, output);
}
void RunGateQuantizedRows(
std::span<const std::int8_t> input,
std::span<const float> scales,
std::span<const std::size_t> rows,
const Expert & weights,
std::span<float> output) {
gate_up_->RunBatchQuantizedRows(input, scales, rows, *weights.gate_up_, output);
}
void PrepareBatch(std::size_t rows, bool indexed_gate) {
gate_up_->PrepareBatch(rows, indexed_gate);
down_->PrepareBatch(rows);
}
private:
std::unique_ptr<Linear> gate_up_, down_;
std::vector<float> projected_, hidden_, output_;
int current_core_ = -1;
};
class SparseFeedForward final : public FeedForward {
public:
SparseFeedForward(const ModelPackage & package, int layer, SparseBatchScratch & scratch)
: package_(package),
layer_(layer),
gate_weight_(DecodeBf16(
package.tensor("model.layers." + std::to_string(layer) + ".mlp.gate.weight"),
kExpertCount * kHiddenSize)),
expert_bias_(DecodeFloat(
package.tensor("model.layers." + std::to_string(layer) + ".mlp.gate.expert_bias"),
kExpertCount)),
shared_gate_up_(MakeLinear(
package,
"model.layers." + std::to_string(layer) + ".mlp.shared_experts.gate_up")),
shared_down_(MakeLinear(
package,
"model.layers." + std::to_string(layer) + ".mlp.shared_experts.down_proj")),
shared_projected_(2 * kExpertWidth), shared_hidden_(kExpertWidth), shared_output_(kHidden),
logits_(kExpertCount),
batch_shared_projected_(scratch.shared_projected),
batch_shared_hidden_(scratch.shared_hidden), batch_shared_output_(scratch.shared_output),
batch_lane_input_(scratch.lane_input), batch_quantized_(scratch.quantized),
batch_lane_projected_(scratch.lane_projected), batch_lane_hidden_(scratch.lane_hidden),
batch_lane_expert_output_(scratch.lane_output), batch_contributions_(scratch.contributions),
all_core_experts_(std::getenv("LING3_EXPERT_ALL_CORES") != nullptr),
balanced_experts_(std::getenv("LING3_EXPERT_BALANCED") != nullptr),
zero_copy_expert_input_(std::getenv("LING3_EXPERT_ZERO_COPY") != nullptr) {
if (all_core_experts_ && balanced_experts_) {
throw std::runtime_error(
"LING3_EXPERT_ALL_CORES and LING3_EXPERT_BALANCED are mutually exclusive");
}
for (auto & output : lane_output_) output.assign(kHidden, 0.0F);
const bool prewarm = std::getenv("LING3_PREWARM_EXPERTS") != nullptr;
if (zero_copy_expert_input_ && (all_core_experts_ || !prewarm)) {
throw std::runtime_error(
"LING3_EXPERT_ZERO_COPY requires prewarmed single-core experts");
}
if (prewarm) PrewarmExperts();
if (zero_copy_expert_input_) ShareExpertInputs();
reuse_batch_input_ = std::getenv("LING3_DISABLE_BATCH_INPUT_REUSE") == nullptr &&
std::getenv("LING3_PREFILL_W4A4") == nullptr;
if ((package.header().flags & 0x100) && std::getenv("LING3_OFFICIAL_EXECUTION") &&
std::string_view(std::getenv("LING3_OFFICIAL_EXECUTION")) == "fp16") reuse_batch_input_ = false;
if (reuse_batch_input_) {
batch_quantized_.resize(kMaxBatch * kHidden);
} else {
for (auto & lane : batch_lane_input_) lane.resize(kMaxBatch * kHidden);
}
for (auto & lane : batch_lane_projected_) lane.resize(kMaxBatch * 2 * kExpertWidth);
for (auto & lane : batch_lane_hidden_) lane.resize(kMaxBatch * kExpertWidth);
for (auto & lane : batch_lane_expert_output_) lane.resize(kMaxBatch * kHidden);
batch_contributions_.resize(kMaxBatch * kExpertsPerToken * kHidden);
}
~SparseFeedForward() override = default;
void Run(std::span<const float> input, std::span<float> output) override {
const bool trace = std::getenv("LING3_TRACE_FFN") != nullptr;
const auto begin = Clock::now();
shared_gate_up_->Run(input, shared_projected_);
SiluMultiply(
shared_projected_.data(),
shared_projected_.data() + kExpertWidth,
shared_hidden_.data(),
kExpertWidth);
shared_down_->Run(shared_hidden_, shared_output_);
const auto shared_end = Clock::now();
RouterLogitsF32(input.data(), gate_weight_.data(), logits_.data());
active_route_ = SelectRoute(logits_.data(), expert_bias_.data());
NumericRoutes("layer" + std::to_string(layer_) + "_routes", {&active_route_, 1});
active_input_ = input.data();
if (zero_copy_expert_input_) {
active_input_scale_ = experts_[0]->PrepareGateInput(input);
}
for (auto & lane : lane_output_) std::fill(lane.begin(), lane.end(), 0.0F);
if (all_core_experts_) {
RunAllCoreExperts();
} else {
if (balanced_experts_) BalanceExpertLanes();
constexpr std::array<int, 3> cores {0, 1, 2};
CoreWorkers::Instance().Run(cores, [this](int core) { RunLane(core); });
}
const auto experts_end = Clock::now();
std::copy(shared_output_.begin(), shared_output_.end(), output.begin());
for (const auto & lane : lane_output_) {
for (int index = 0; index < kHidden; ++index) output[index] += lane[index];
}
if (trace) {
const char * mode = all_core_experts_ ? "all-core" :
(zero_copy_expert_input_ ? "balanced-zero-copy" :
(balanced_experts_ ? "balanced" : "expert-parallel"));
std::fprintf(stderr, " mode=%s shared_ms=%.3f experts_ms=%.3f total_ms=%.3f ids=",
mode,
Milliseconds(begin, shared_end), Milliseconds(shared_end, experts_end),
Milliseconds(begin, Clock::now()));
for (int id : active_route_.experts) std::fprintf(stderr, "%d,", id);
std::fprintf(stderr, "\n");
}
}
void RunBatch(
std::span<const float> input,
std::size_t rows,
std::span<float> output) override {
const bool trace_batch = std::getenv("LING3_TRACE_BATCH") != nullptr;
const auto batch_begin = Clock::now();
if (rows < 1 || rows > kMaxBatch || input.size() != rows * kHidden ||
output.size() != rows * kHidden) {
throw std::invalid_argument("sparse FFN batch has an incompatible tensor size");
}
if (all_core_experts_) {
throw std::runtime_error("sparse FFN batch requires single-core expert contexts");
}
std::array<std::thread, 4> router_workers;
for (std::size_t worker = 0; worker < router_workers.size(); ++worker) {
router_workers[worker] = std::thread([this, input, rows, worker]() {
std::array<float, kExpertCount> row_logits {};
for (std::size_t row = worker; row < rows; row += 4) {
if (reuse_batch_input_) {
batch_input_scales_[row] = QuantizeSymmetricInt8(
input.subspan(row * kHidden, kHidden),
std::span<std::int8_t>(batch_quantized_).subspan(
row * kHidden, kHidden)).scale;
}
RouterLogitsF32(
input.data() + row * kHidden, gate_weight_.data(), row_logits.data());
const auto route = SelectRoute(row_logits.data(), expert_bias_.data());
batch_routes_[row] = route;
std::array<int, 3> counts {};
for (std::size_t slot = 0; slot < kExpertsPerToken; ++slot) {
const int lane = route.experts[slot] % 3;
batch_route_lanes_[row][slot] = lane;
++counts[lane];
}
while (*std::max_element(counts.begin(), counts.end()) > 3 ||
*std::min_element(counts.begin(), counts.end()) < 2) {
const int source = static_cast<int>(
std::max_element(counts.begin(), counts.end()) - counts.begin());
const int target = static_cast<int>(
std::min_element(counts.begin(), counts.end()) - counts.begin());
const auto found = std::find(
batch_route_lanes_[row].begin(),
batch_route_lanes_[row].end(), source);
if (found == batch_route_lanes_[row].end()) break;
*found = target;
--counts[source];
++counts[target];
}
}
});
}
shared_gate_up_->RunBatch(
input, rows,
std::span<float>(batch_shared_projected_).first(rows * 2 * kExpertWidth));
const auto shared_gate_at = Clock::now();
for (std::size_t row = 0; row < rows; ++row) {
const auto * projected = batch_shared_projected_.data() + row * 2 * kExpertWidth;
SiluMultiply(
projected, projected + kExpertWidth,
batch_shared_hidden_.data() + row * kExpertWidth, kExpertWidth);
}
shared_down_->RunBatch(
std::span<const float>(batch_shared_hidden_).first(rows * kExpertWidth), rows,
std::span<float>(batch_shared_output_).first(rows * kHidden));
for (auto & worker : router_workers) worker.join();
const auto shared_done_at = Clock::now();
for (auto & group : batch_groups_) group.clear();
for (std::size_t row = 0; row < rows; ++row) {
for (std::size_t slot = 0; slot < kExpertsPerToken; ++slot) {
batch_groups_[batch_routes_[row].experts[slot]].push_back({row, slot});
}
}
std::vector<int> active_experts;
active_experts.reserve(kExpertCount);
for (int id = 0; id < static_cast<int>(kExpertCount); ++id) {
if (batch_groups_[id].empty()) continue;
if (!experts_[id]) {
experts_[id] = std::make_unique<Expert>(package_, layer_, id, false);
}
active_experts.push_back(id);
}
for (int core = 0; core < 3; ++core) {
if (!experts_[core]) {
experts_[core] = std::make_unique<Expert>(package_, layer_, core, false);
}
experts_[core]->SetCore(core);
batch_jobs_[core].clear();
}
std::sort(active_experts.begin(), active_experts.end(), [this](int left, int right) {
const auto left_size = batch_groups_[left].size();
const auto right_size = batch_groups_[right].size();
return left_size != right_size ? left_size > right_size : left < right;
});
std::array<std::size_t, 3> lane_load {};
std::array<std::size_t, 3> lane_cost {};
std::array<std::size_t, 3> lane_experts {};
for (int id : active_experts) {
const int lane = static_cast<int>(
std::min_element(lane_load.begin(), lane_load.end()) - lane_load.begin());
batch_jobs_[lane].push_back(id);
lane_load[lane] += batch_groups_[id].size();
if (trace_batch) {
lane_cost[lane] += ExpertBatchCost(batch_groups_[id].size());
++lane_experts[lane];
}
}
const auto scheduled_at = Clock::now();
constexpr std::array<int, 3> cores {0, 1, 2};
std::array<double, 3> lane_ms {};
CoreWorkers::Instance().Run(cores, [this, input, trace_batch, &lane_ms](int core) {
const auto lane_begin = trace_batch ? Clock::now() : Clock::time_point {};
auto & runner = *experts_[core];
auto & gathered = batch_lane_input_[core];
auto & projected = batch_lane_projected_[core];
auto & hidden = batch_lane_hidden_[core];
auto & expert_output = batch_lane_expert_output_[core];
for (const int id : batch_jobs_[core]) {
const auto & assignments = batch_groups_[id];
const std::size_t count = assignments.size();
if (reuse_batch_input_) {
auto & indices = batch_lane_rows_[core];
for (std::size_t index = 0; index < count; ++index) {
indices[index] = assignments[index].row;
}
runner.RunGateQuantizedRows(
batch_quantized_, batch_input_scales_,
std::span<const std::size_t>(indices).first(count), *experts_[id],
std::span<float>(projected).first(count * 2 * kExpertWidth));
} else {
for (std::size_t index = 0; index < count; ++index) {
std::memcpy(
gathered.data() + index * kHidden,
input.data() + assignments[index].row * kHidden,
static_cast<std::size_t>(kHidden) * sizeof(float));
}
runner.RunGateBatch(
std::span<const float>(gathered).first(count * kHidden), count,
*experts_[id],
std::span<float>(projected).first(count * 2 * kExpertWidth));
}
for (std::size_t index = 0; index < count; ++index) {
const auto * row = projected.data() + index * 2 * kExpertWidth;
SiluMultiply(
row, row + kExpertWidth,
hidden.data() + index * kExpertWidth, kExpertWidth);
}
runner.RunDownBatch(
std::span<const float>(hidden).first(count * kExpertWidth), count,
*experts_[id],
std::span<float>(expert_output).first(count * kHidden));
for (std::size_t index = 0; index < count; ++index) {
std::memcpy(
batch_contributions_.data() +
(assignments[index].row * kExpertsPerToken +
assignments[index].slot) * kHidden,
expert_output.data() + index * kHidden,
static_cast<std::size_t>(kHidden) * sizeof(float));
}
}
if (trace_batch) lane_ms[core] = Milliseconds(lane_begin, Clock::now());
});
const auto experts_done_at = Clock::now();
NumericRoutes("layer" + std::to_string(layer_) + "_routes",
std::span<const Route>(batch_routes_).first(rows));
std::copy_n(batch_shared_output_.begin(), rows * kHidden, output.begin());
for (std::size_t row = 0; row < rows; ++row) {
auto * destination = output.data() + row * kHidden;
for (int lane = 0; lane < 3; ++lane) {
for (std::size_t slot = 0; slot < kExpertsPerToken; ++slot) {
if (batch_route_lanes_[row][slot] != lane) continue;
WeightedAccumulate(
batch_contributions_.data() +
(row * kExpertsPerToken + slot) * kHidden,
batch_routes_[row].weights[slot], destination, kHidden);
}
}
}
if (trace_batch) {
std::fprintf(
stderr,
" sparse_batch rows=%zu active_experts=%zu lane_load=%zu,%zu,%zu "
"lane_cost=%zu,%zu,%zu lane_experts=%zu,%zu,%zu lane_ms=%.3f,%.3f,%.3f "
"shared_gate=%.3f shared_down_router=%.3f schedule=%.3f experts=%.3f "
"combine=%.3f total=%.3f\n",
rows, active_experts.size(), lane_load[0], lane_load[1], lane_load[2],
lane_cost[0], lane_cost[1], lane_cost[2],
lane_experts[0], lane_experts[1], lane_experts[2],
lane_ms[0], lane_ms[1], lane_ms[2],
Milliseconds(batch_begin, shared_gate_at),
Milliseconds(shared_gate_at, shared_done_at),
Milliseconds(shared_done_at, scheduled_at),
Milliseconds(scheduled_at, experts_done_at),
Milliseconds(experts_done_at, Clock::now()),
Milliseconds(batch_begin, Clock::now()));
}
}
void PrepareBatch(std::size_t rows) override {
if (rows < 1 || rows > kMaxBatch) {
throw std::invalid_argument("sparse FFN batch rows must be in [1, 128]");
}
if (all_core_experts_) {
throw std::runtime_error("sparse FFN batch requires single-core expert contexts");
}
shared_gate_up_->PrepareBatch(rows);
shared_down_->PrepareBatch(rows);
for (int core = 0; core < 3; ++core) {
if (!experts_[core]) {
experts_[core] = std::make_unique<Expert>(package_, layer_, core, false);
}
experts_[core]->SetCore(core);
for (std::size_t bucket = 1; bucket <= rows; bucket *= 2) {
experts_[core]->PrepareBatch(bucket, reuse_batch_input_);
}
}
}
private:
struct BatchAssignment {
std::size_t row = 0;
std::size_t slot = 0;
};
void PrewarmExperts() {
std::array<std::exception_ptr, 3> errors {};
std::array<std::thread, 3> workers;
for (int core = 0; core < 3; ++core) {
workers[core] = std::thread([this, core, &errors]() {
try {
for (int id = core; id < static_cast<int>(kExpertCount); id += 3) {
experts_[id] = std::make_unique<Expert>(
package_, layer_, id, all_core_experts_);
}
} catch (...) {
errors[core] = std::current_exception();
}
});
}
for (auto & worker : workers) if (worker.joinable()) worker.join();
for (const auto & error : errors) if (error != nullptr) std::rethrow_exception(error);
}
void RunLane(int core) {
for (std::size_t index = 0; index < kExpertsPerToken; ++index) {
const int id = active_route_.experts[index];
const int lane = balanced_experts_ ? active_lanes_[index] : id % 3;
if (lane != core) continue;
if (!experts_[id]) {
experts_[id] = std::make_unique<Expert>(
package_, layer_, id, all_core_experts_);
}
const auto value = experts_[id]->Run(
{active_input_, static_cast<std::size_t>(kHidden)},
zero_copy_expert_input_ ? active_input_scale_ : 0.0F);
WeightedAccumulate(
value.data(), active_route_.weights[index],
lane_output_[core].data(), kHidden);
}
}
void BalanceExpertLanes() {
std::array<int, 3> counts {};
for (std::size_t index = 0; index < kExpertsPerToken; ++index) {
const int id = active_route_.experts[index];
if (!experts_[id]) {
experts_[id] = std::make_unique<Expert>(package_, layer_, id, false);
}
active_lanes_[index] = id % 3;
++counts[active_lanes_[index]];
}
while (*std::max_element(counts.begin(), counts.end()) > 3 ||
*std::min_element(counts.begin(), counts.end()) < 2) {
const int source = static_cast<int>(
std::max_element(counts.begin(), counts.end()) - counts.begin());
const int target = static_cast<int>(
std::min_element(counts.begin(), counts.end()) - counts.begin());
const auto found = std::find(active_lanes_.begin(), active_lanes_.end(), source);
if (found == active_lanes_.end()) {
throw std::runtime_error("balanced expert scheduler lost its source lane");
}
*found = target;
--counts[source];
++counts[target];
}
// Core-mask changes are RKNN context mutations. Apply them serially before
// the lane workers run so dynamic scheduling stays deterministic.
for (std::size_t index = 0; index < kExpertsPerToken; ++index) {
experts_[active_route_.experts[index]]->SetCore(active_lanes_[index]);
}
}
void RunAllCoreExperts() {
for (std::size_t index = 0; index < kExpertsPerToken; ++index) {
const int id = active_route_.experts[index];
if (!experts_[id]) {
experts_[id] = std::make_unique<Expert>(package_, layer_, id, true);
}
const auto value = experts_[id]->Run(
{active_input_, static_cast<std::size_t>(kHidden)});
WeightedAccumulate(
value.data(), active_route_.weights[index],
lane_output_[0].data(), kHidden);
}
}
void ShareExpertInputs() {
if (!experts_[0]) throw std::runtime_error("expert input owner was not prewarmed");
for (std::size_t id = 1; id < experts_.size(); ++id) {
if (!experts_[id]) throw std::runtime_error("expert input target was not prewarmed");
experts_[id]->ShareGateInputFrom(*experts_[0]);
}
}
const ModelPackage & package_;
int layer_;
std::vector<float> gate_weight_, expert_bias_;
std::unique_ptr<Linear> shared_gate_up_, shared_down_;
std::vector<float> shared_projected_, shared_hidden_, shared_output_, logits_;
std::vector<float> &batch_shared_projected_, &batch_shared_hidden_, &batch_shared_output_;
std::array<std::unique_ptr<Expert>, kExpertCount> experts_;
std::array<std::vector<float>, 3> lane_output_;
Route active_route_;
std::array<int, kExpertsPerToken> active_lanes_ {};
std::array<std::vector<BatchAssignment>, kExpertCount> batch_groups_;
std::array<Route, kMaxBatch> batch_routes_;
std::array<std::array<int, kExpertsPerToken>, kMaxBatch> batch_route_lanes_ {};
std::array<std::vector<int>, 3> batch_jobs_;
std::array<std::vector<float>, 3> &batch_lane_input_;
std::vector<std::int8_t> &batch_quantized_;
std::array<float, kMaxBatch> batch_input_scales_ {};
std::array<std::array<std::size_t, kMaxBatch>, 3> batch_lane_rows_ {};
std::array<std::vector<float>, 3> &batch_lane_projected_;
std::array<std::vector<float>, 3> &batch_lane_hidden_;
std::array<std::vector<float>, 3> &batch_lane_expert_output_;
std::vector<float> &batch_contributions_;
const float * active_input_ = nullptr;
float active_input_scale_ = 1.0F;
bool all_core_experts_ = false;
bool balanced_experts_ = false;
bool zero_copy_expert_input_ = false;
bool reuse_batch_input_ = false;
};
#if LING3_EXPERIMENTAL_MTP
// Layer 24 from the optional Ling NEXTN/MTP sidecar. It is deliberately
// separate from DecoderLayer: the main trunk has KDA/MLA grouping and its
// state/checkpoints, while MTP consumes the trunk's normalized hidden state
// together with the embedding of the candidate token.
class MtpLayer {
public:
MtpLayer(const ModelPackage & package, std::size_t max_context, DecoderScratch & scratch)
: input_norm_(DecodeBf16(package.tensor("model.layers.24.input_layernorm.weight"), kHidden)),
post_norm_(DecodeBf16(package.tensor("model.layers.24.post_attention_layernorm.weight"), kHidden)),
embedding_norm_(DecodeBf16(package.tensor("model.layers.24.enorm.weight"), kHidden)),
hidden_norm_(DecodeBf16(package.tensor("model.layers.24.hnorm.weight"), kHidden)),
final_norm_(DecodeBf16(package.tensor("model.layers.24.final_layernorm.weight"), kHidden)),
embedding_hidden_(kHidden * 2), projected_(kHidden), residual_(kHidden),
normalized_(kHidden), attention_output_(kHidden), ffn_output_(kHidden),
eh_projection_(MakeLinear(package, "model.layers.24.eh_proj")),
attention_(std::make_unique<MlaAttention>(package, 24, max_context, scratch.mla)),
feed_forward_(std::make_unique<SparseFeedForward>(package, 24, scratch.sparse)) {}
void Reset() { attention_->Reset(); }
void Run(std::span<const float> token_embedding, std::span<const float> trunk_hidden,
std::size_t position, std::span<float> output) {
if (token_embedding.size() != kHidden || trunk_hidden.size() != kHidden ||
output.size() != kHidden) {
throw std::invalid_argument("MTP hidden shape mismatch");
}
RmsNorm(token_embedding.data(), embedding_norm_.data(), embedding_hidden_.data(), kHidden, kEpsilon);
RmsNorm(trunk_hidden.data(), hidden_norm_.data(), embedding_hidden_.data() + kHidden, kHidden, kEpsilon);
eh_projection_->Run(embedding_hidden_, projected_);
std::copy(projected_.begin(), projected_.end(), residual_.begin());
RmsNorm(residual_.data(), input_norm_.data(), normalized_.data(), kHidden, kEpsilon);
attention_->Run(normalized_, position, attention_output_);
for (int i = 0; i < kHidden; ++i) residual_[i] += attention_output_[i];
RmsNorm(residual_.data(), post_norm_.data(), normalized_.data(), kHidden, kEpsilon);
feed_forward_->Run(normalized_, ffn_output_);
for (int i = 0; i < kHidden; ++i) residual_[i] += ffn_output_[i];
RmsNorm(residual_.data(), final_norm_.data(), output.data(), kHidden, kEpsilon);
}
void PrepareBatch(std::size_t rows) {
eh_projection_->PrepareBatch(rows);
attention_->PrepareBatch(rows);
feed_forward_->PrepareBatch(rows);
}
private:
std::vector<float> input_norm_, post_norm_, embedding_norm_, hidden_norm_, final_norm_;
std::vector<float> embedding_hidden_, projected_, residual_, normalized_, attention_output_, ffn_output_;
std::unique_ptr<Linear> eh_projection_;
std::unique_ptr<MlaAttention> attention_;
std::unique_ptr<SparseFeedForward> feed_forward_;
};
#endif
class DecoderLayer {
public:
DecoderLayer(
const ModelPackage & package,
int layer,
std::size_t max_context,
std::span<const std::byte> heads6,
std::span<const std::byte> heads5,
DecoderScratch & scratch)
: trace_name_("layer" + std::to_string(layer)), input_norm_(DecodeBf16(
package.tensor("model.layers." + std::to_string(layer) + ".input_layernorm.weight"),
kHidden)),
post_norm_(DecodeBf16(
package.tensor("model.layers." + std::to_string(layer) + ".post_attention_layernorm.weight"),
kHidden)),
attention_((layer + 1) % 4 == 0
? std::unique_ptr<Attention>(std::make_unique<MlaAttention>(package, layer, max_context, scratch.mla))
: std::unique_ptr<Attention>(std::make_unique<KdaAttention>(package, layer, heads6, heads5, scratch.kda))),
feed_forward_(layer == 0
? std::unique_ptr<FeedForward>(std::make_unique<DenseFeedForward>(package, layer))
: std::unique_ptr<FeedForward>(std::make_unique<SparseFeedForward>(package, layer, scratch.sparse))),
normalized_(kHidden), attention_output_(kHidden), ffn_output_(kHidden),
batch_normalized_(scratch.layer.normalized),
batch_attention_output_(scratch.layer.attention_output),
batch_ffn_output_(scratch.layer.ffn_output) {}
void Reset() { attention_->Reset(); }
AttentionCheckpoint SaveCheckpoint() { return attention_->SaveCheckpoint(); }
AttentionState SaveState(std::size_t position) { return attention_->SaveState(position); }
void RestoreCheckpoint(const AttentionCheckpoint & checkpoint) { attention_->RestoreCheckpoint(checkpoint); }
void Run(std::vector<float> & hidden, std::size_t position) {
const bool trace = std::getenv("LING3_TRACE_LAYERS") != nullptr;
RmsNorm(hidden.data(), input_norm_.data(), normalized_.data(), kHidden, kEpsilon);
const auto attention_begin = Clock::now();
attention_->Run(normalized_, position, attention_output_);
const auto attention_end = Clock::now();
for (int index = 0; index < kHidden; ++index) hidden[index] += attention_output_[index];
NumericDump(trace_name_ + "_attention", attention_output_);
NumericDump(trace_name_ + "_post_attention", hidden);
RmsNorm(hidden.data(), post_norm_.data(), normalized_.data(), kHidden, kEpsilon);
const auto ffn_begin = Clock::now();
feed_forward_->Run(normalized_, ffn_output_);
const auto ffn_end = Clock::now();
for (int index = 0; index < kHidden; ++index) hidden[index] += ffn_output_[index];
NumericDump(trace_name_ + "_ffn", ffn_output_);
NumericDump(trace_name_ + "_output", hidden);
if (trace) {
std::fprintf(stderr, " attention_ms=%.3f ffn_ms=%.3f\n",
Milliseconds(attention_begin, attention_end),
Milliseconds(ffn_begin, ffn_end));
}
}
void RunBatch(
std::span<float> hidden,
std::size_t rows,
std::size_t position) {
if (rows < 1 || rows > kMaxBatch || hidden.size() != rows * kHidden) {
throw std::invalid_argument("decoder layer batch has an incompatible tensor size");
}
const bool trace = std::getenv("LING3_TRACE_LAYERS") != nullptr;
for (std::size_t row = 0; row < rows; ++row) {
RmsNorm(
hidden.data() + row * kHidden, input_norm_.data(),
batch_normalized_.data() + row * kHidden, kHidden, kEpsilon);
}
const auto attention_begin = Clock::now();
attention_->RunBatch(
std::span<const float>(batch_normalized_).first(rows * kHidden), rows, position,
std::span<float>(batch_attention_output_).first(rows * kHidden));
const auto attention_end = Clock::now();
for (std::size_t index = 0; index < rows * kHidden; ++index) {
hidden[index] += batch_attention_output_[index];
}
NumericDump(trace_name_ + "_attention", std::span<const float>(batch_attention_output_).first(rows * kHidden));
NumericDump(trace_name_ + "_post_attention", hidden);
for (std::size_t row = 0; row < rows; ++row) {
RmsNorm(
hidden.data() + row * kHidden, post_norm_.data(),
batch_normalized_.data() + row * kHidden, kHidden, kEpsilon);
}
const auto ffn_begin = Clock::now();
feed_forward_->RunBatch(
std::span<const float>(batch_normalized_).first(rows * kHidden), rows,
std::span<float>(batch_ffn_output_).first(rows * kHidden));
const auto ffn_end = Clock::now();
for (std::size_t index = 0; index < rows * kHidden; ++index) {
hidden[index] += batch_ffn_output_[index];
}
NumericDump(trace_name_ + "_ffn", std::span<const float>(batch_ffn_output_).first(rows * kHidden));
NumericDump(trace_name_ + "_output", hidden);
if (trace) {
std::fprintf(stderr, " batch_attention_ms=%.3f batch_ffn_ms=%.3f\n",
Milliseconds(attention_begin, attention_end),
Milliseconds(ffn_begin, ffn_end));
}
}
void PrepareBatch(std::size_t rows) {
attention_->PrepareBatch(rows);
feed_forward_->PrepareBatch(rows);
}
private:
std::string trace_name_;
std::vector<float> input_norm_, post_norm_;
std::unique_ptr<Attention> attention_;
std::unique_ptr<FeedForward> feed_forward_;
std::vector<float> normalized_, attention_output_, ffn_output_;
std::vector<float> &batch_normalized_, &batch_attention_output_, &batch_ffn_output_;
};
} // namespace
struct DecoderCheckpoint {
std::weak_ptr<int> owner;
std::size_t position = 0;
bool valid = true;
std::vector<AttentionCheckpoint> layers;
};
struct Decoder::Impl {
const ModelPackage & package;
std::size_t context_capacity;
std::span<const std::uint16_t> embeddings;
std::vector<float> final_norm;
DecoderScratch scratch; // Declared before layers: outlives all borrowers.
std::vector<std::unique_ptr<DecoderLayer>> layers;
#if LING3_EXPERIMENTAL_MTP
std::unique_ptr<MtpLayer> mtp;
std::vector<float> mtp_hidden, mtp_embedding;
std::size_t mtp_position = 0;
bool mtp_trunk_ready = false;
#endif
std::unique_ptr<Linear> lm_head;
std::vector<float> hidden;
std::vector<float> normalized;
std::vector<float> batch_hidden;
std::size_t current_position = 0;
std::shared_ptr<int> checkpoint_owner = std::make_shared<int>(0);
std::vector<std::weak_ptr<DecoderCheckpoint>> checkpoints;
std::string state_signature;
explicit Impl(const ModelPackage & model, std::size_t capacity)
: package(model),
context_capacity(capacity ? capacity : model.header().max_context),
embeddings(Typed<std::uint16_t>(
package.tensor("model.word_embeddings.weight"), DataType::kBFloat16)),
final_norm(DecodeBf16(package.tensor("model.norm.weight"), kHidden)),
hidden(kHidden), normalized(kHidden), batch_hidden(kMaxBatch * kHidden) {
ValidateLing3Tiny(package.header());
if (std::getenv("LING3_W8_SOURCE") || std::getenv("LING3_CALIBRATED_W4_SOURCE") ||
std::getenv("LING3_GDN_PREFILL_DIR"))
state_signature = "external-weight-overrides:";
// Persist the exact package metadata (including every tensor SHA256),
// not merely the base model revision shared by differently quantized files.
state_signature += "Ling3RKNN-numerical-state-v1";
state_signature.append(reinterpret_cast<const char *>(&package.header()), sizeof(PackageHeader));
for (const auto & tensor : package.tensors()) {
state_signature.append(reinterpret_cast<const char *>(tensor.entry), sizeof(TensorEntry));
state_signature.append(tensor.name);
}
for (const auto key : {"LING3_GDN_CPU_FP32_STATE", "LING3_GDN_CPU_DECODE", "LING3_GDN_FULL_FP32",
"LING3_GDN_CPU_PREFILL", "LING3_GDN_PREFILL_DIR", "LING3_PREFILL_W4A4", "LING3_MLA_BACKEND",
"LING3_MLA_SIMD", "LING3_VECTOR_MATH", "LING3_OFFICIAL_EXECUTION", "LING3_BRIDGE_SHARED_STAGE",
"LING3_EXPERT_BALANCED", "LING3_EXPERT_ALL_CORES", "LING3_EXPERT_ZERO_COPY",
"LING3_DISABLE_BATCH_INPUT_REUSE", "LING3_DISABLE_PARALLEL_GATHER"}) {
state_signature.append(key); state_signature.push_back('=');
if (const auto value = std::getenv(key)) state_signature.append(value);
state_signature.push_back('\0');
}
if (context_capacity < 1 || context_capacity > 262144)
throw std::invalid_argument("context capacity must be in [1, 262144]");
if ((package.header().flags & 1U) == 0) {
throw std::runtime_error("decoder requires a complete Ling3RKNN package");
}
const auto heads6 = Blob(package.tensor("rknn.gdn.heads6"), TensorRole::kRknnIsland);
const auto heads5 = Blob(package.tensor("rknn.gdn.heads5"), TensorRole::kRknnIsland);
layers.reserve(24);
for (int layer = 0; layer < 24; ++layer) {
layers.push_back(std::make_unique<DecoderLayer>(
package, layer, context_capacity, heads6, heads5, scratch));
}
lm_head = MakeLinear(package, "lm_head");
}
#if LING3_EXPERIMENTAL_MTP
bool EnableMtp() {
if (current_position != 0)
throw std::invalid_argument("MTP must be enabled before processing the prefix");
if (mtp) return true;
if (std::any_of(package.tensors().begin(), package.tensors().end(),
[](const TensorView & t) { return t.name == "model.layers.24.eh_proj.weight"; })) {
mtp = std::make_unique<MtpLayer>(package, context_capacity, scratch);
mtp_hidden.resize(kHidden);
mtp_embedding.resize(kHidden);
}
return mtp != nullptr;
}
#endif
void Reset() {
for (auto & weak : checkpoints) if (auto checkpoint = weak.lock()) checkpoint->valid = false;
checkpoints.clear();
current_position = 0;
for (auto & layer : layers) layer->Reset();
#if LING3_EXPERIMENTAL_MTP
if (mtp) mtp->Reset();
mtp_position = 0;
mtp_trunk_ready = false;
#endif
}
void InvalidateOverwrittenCheckpoints() {
std::erase_if(checkpoints, [this](const auto & weak) {
const auto checkpoint = weak.lock();
if (!checkpoint) return true;
if (checkpoint->position > current_position) checkpoint->valid = false;
return !checkpoint->valid;
});
}
std::shared_ptr<DecoderCheckpoint> SaveCheckpoint() {
std::erase_if(checkpoints, [](const auto & weak) { return weak.expired(); });
auto checkpoint = std::make_shared<DecoderCheckpoint>();
checkpoint->owner = checkpoint_owner;
checkpoint->position = current_position;
for (auto & layer : layers) checkpoint->layers.push_back(layer->SaveCheckpoint());
checkpoints.push_back(checkpoint);
return checkpoint;
}
std::shared_ptr<const DecoderState> SaveState() {
if (state_signature.starts_with("external-weight-overrides:"))
throw std::invalid_argument("full states require a self-contained model without external weights or graphs");
auto state = std::make_shared<DecoderState>();
state->signature = state_signature; state->position = current_position;
for (auto & layer : layers) state->layers.push_back(layer->SaveState(current_position));
return state;
}
std::size_t RestoreState(const DecoderState & state) {
if (state_signature.starts_with("external-weight-overrides:"))
throw std::invalid_argument("full states require a self-contained model without external weights or graphs");
ValidateDecoderState(state, state_signature, context_capacity);
// All input validation precedes mutation. Old lightweight views refer to
// overwritten KV and must not survive a full-state switch.
Reset();
try {
for (std::size_t i = 0; i < layers.size(); ++i) layers[i]->RestoreCheckpoint(state.layers[i]);
current_position = state.position;
} catch (...) { Reset(); throw; }
return current_position;
}
std::size_t RestoreCheckpoint(const DecoderCheckpoint & checkpoint) {
if (!checkpoint.valid || checkpoint.owner.lock() != checkpoint_owner ||
checkpoint.layers.size() != layers.size())
throw std::invalid_argument("stale or foreign decoder checkpoint");
for (std::size_t i = 0; i < layers.size(); ++i) layers[i]->RestoreCheckpoint(checkpoint.layers[i]);
current_position = checkpoint.position;
#if LING3_EXPERIMENTAL_MTP
mtp_trunk_ready = false;
#endif
return current_position;
}
DecodeTimings Eval(std::uint32_t token, std::span<float> logits) {
if (token >= package.header().vocab_size ||
logits.size() != package.header().vocab_size ||
current_position >= context_capacity) {
throw std::invalid_argument("decoder token, logits, or context is out of range");
}
const auto begin = Clock::now();
const auto embedding = embeddings.subspan(static_cast<std::size_t>(token) * kHidden, kHidden);
InvalidateOverwrittenCheckpoints();
for (int index = 0; index < kHidden; ++index) hidden[index] = BFloat16ToFloat(embedding[index]);
const bool trace_layers = std::getenv("LING3_TRACE_LAYERS") != nullptr;
numeric_position = current_position;
for (std::size_t index = 0; index < layers.size(); ++index) {
const auto layer_begin = Clock::now();
layers[index]->Run(hidden, current_position);
if (trace_layers) {
std::fprintf(
stderr, "layer=%zu ms=%.3f\n", index,
Milliseconds(layer_begin, Clock::now()));
}
}
const auto layers_end = Clock::now();
RmsNorm(hidden.data(), final_norm.data(), normalized.data(), kHidden, kEpsilon);
lm_head->Run(normalized, logits);
const auto end = Clock::now();
++current_position;
#if LING3_EXPERIMENTAL_MTP
mtp_trunk_ready = true;
#endif
return {
Milliseconds(begin, layers_end),
Milliseconds(layers_end, end),
Milliseconds(begin, end),
};
}
DecodeTimings EvalBatch(
std::span<const std::uint32_t> tokens,
std::span<float> logits,
bool compute_logits) {
const std::size_t rows = tokens.size();
if (rows < 1 || rows > kMaxBatch ||
(compute_logits && logits.size() != package.header().vocab_size) ||
current_position + rows > context_capacity) {
throw std::invalid_argument("decoder batch tokens, logits, or context is out of range");
}
const auto begin = Clock::now();
for (std::size_t row = 0; row < rows; ++row) {
if (tokens[row] >= package.header().vocab_size) {
throw std::invalid_argument("decoder batch token is out of range");
}
const auto embedding = embeddings.subspan(
static_cast<std::size_t>(tokens[row]) * kHidden, kHidden);
for (int index = 0; index < kHidden; ++index) {
batch_hidden[row * kHidden + index] = BFloat16ToFloat(embedding[index]);
}
}
InvalidateOverwrittenCheckpoints();
const bool trace_layers = std::getenv("LING3_TRACE_LAYERS") != nullptr;
numeric_position = current_position;
for (std::size_t index = 0; index < layers.size(); ++index) {
const auto layer_begin = Clock::now();
layers[index]->RunBatch(
std::span<float>(batch_hidden).first(rows * kHidden), rows, current_position);
if (trace_layers) {
std::fprintf(
stderr, "batch_layer=%zu ms=%.3f\n", index,
Milliseconds(layer_begin, Clock::now()));
}
}
const auto layers_end = Clock::now();
if (compute_logits) {
const auto * final_hidden = batch_hidden.data() + (rows - 1) * kHidden;
RmsNorm(final_hidden, final_norm.data(), normalized.data(), kHidden, kEpsilon);
lm_head->Run(normalized, logits);
}
const auto end = Clock::now();
current_position += rows;
#if LING3_EXPERIMENTAL_MTP
mtp_trunk_ready = compute_logits;
#endif
return {
Milliseconds(begin, layers_end),
Milliseconds(layers_end, end),
Milliseconds(begin, end),
};
}
#if LING3_EXPERIMENTAL_MTP
DecodeTimings EvalMtp(std::uint32_t token, std::span<float> logits) {
if (!mtp) throw std::runtime_error("model package has no MTP layer");
if (token >= package.header().vocab_size || logits.size() != package.header().vocab_size ||
current_position == 0 || !mtp_trunk_ready || mtp_position != current_position - 1)
throw std::invalid_argument("MTP requires a contiguous, aligned prefix and valid token/logits");
const auto begin = Clock::now();
const auto embedding = embeddings.subspan(static_cast<std::size_t>(token) * kHidden, kHidden);
for (int i = 0; i < kHidden; ++i) mtp_embedding[i] = BFloat16ToFloat(embedding[i]);
// Upstream passes the main model's final normalized hidden state.
// Keep MTP output separate so probing cannot overwrite trunk scratch.
mtp->Run(mtp_embedding, normalized, mtp_position, mtp_hidden);
const auto layers_end = Clock::now();
lm_head->Run(mtp_hidden, logits);
const auto end = Clock::now();
++mtp_position;
mtp_trunk_ready = false;
return {Milliseconds(begin, layers_end), Milliseconds(layers_end, end), Milliseconds(begin, end)};
}
#endif
void PrepareBatch(std::size_t rows) {
if (rows < 1 || rows > kMaxBatch) {
throw std::invalid_argument("decoder batch rows must be in [1, 128]");
}
if (std::getenv("LING3_GDN_CPU_PREFILL") == nullptr &&
std::getenv("LING3_GDN_PREFILL_DIR") == nullptr) {
throw std::runtime_error(
"LING3_GDN_CPU_PREFILL or LING3_GDN_PREFILL_DIR is required for batch prewarm");
}
for (auto & layer : layers) layer->PrepareBatch(rows);
}
};
Decoder::Decoder(const ModelPackage & package, std::size_t context_capacity)
: impl_(std::make_unique<Impl>(package, context_capacity)) {}
Decoder::~Decoder() = default;
void Decoder::Reset() { impl_->Reset(); }
MlaBackendStats Decoder::AttentionStats() const { return impl_->scratch.mla.npu.Stats(); }
DecodeTimings Decoder::Eval(std::uint32_t token, std::span<float> logits) {
return impl_->Eval(token, logits);
}
#if LING3_EXPERIMENTAL_MTP
bool Decoder::EnableMtp() { return impl_->EnableMtp(); }
bool Decoder::HasMtp() const noexcept { return impl_->mtp != nullptr; }
DecodeTimings Decoder::EvalMtp(std::uint32_t token, std::span<float> logits) {
return impl_->EvalMtp(token, logits);
}
#endif
DecodeTimings Decoder::EvalBatch(
std::span<const std::uint32_t> tokens,
std::span<float> logits) {
return impl_->EvalBatch(tokens, logits, true);
}
DecodeTimings Decoder::EvalBatchState(std::span<const std::uint32_t> tokens) {
return impl_->EvalBatch(tokens, {}, false);
}
std::shared_ptr<DecoderCheckpoint> Decoder::SaveCheckpoint() { return impl_->SaveCheckpoint(); }
std::size_t Decoder::RestoreCheckpoint(const DecoderCheckpoint & checkpoint) { return impl_->RestoreCheckpoint(checkpoint); }
std::shared_ptr<const DecoderState> Decoder::SaveState() { return impl_->SaveState(); }
std::size_t Decoder::RestoreState(const DecoderState & state) { return impl_->RestoreState(state); }
const std::string & Decoder::StateSignature() const { return impl_->state_signature; }
std::size_t Decoder::CheckpointBytes(const DecoderCheckpoint & checkpoint) {
std::size_t bytes = 0;
for (const auto & layer : checkpoint.layers) bytes += layer.bytes();
return bytes;
}
void Decoder::PrepareBatch(std::size_t rows) { impl_->PrepareBatch(rows); }
DecodeTimings Decoder::EvalBatch32(
std::span<const std::uint32_t> tokens,
std::span<float> logits) {
if (tokens.size() != 32) {
throw std::invalid_argument("EvalBatch32 requires exactly 32 tokens");
}
return impl_->EvalBatch(tokens, logits, true);
}
DecodeTimings Decoder::EvalBatch32State(std::span<const std::uint32_t> tokens) {
if (tokens.size() != 32) {
throw std::invalid_argument("EvalBatch32State requires exactly 32 tokens");
}
return impl_->EvalBatch(tokens, {}, false);
}
void Decoder::PrepareBatch32() { impl_->PrepareBatch(32); }
std::vector<std::uint32_t> Decoder::Generate(
std::span<const std::uint32_t> prompt,
std::size_t max_new_tokens) {
if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token");
Reset();
std::vector<float> logits(impl_->package.header().vocab_size);
for (std::uint32_t token : prompt) Eval(token, logits);
std::vector<std::uint32_t> output;
output.reserve(max_new_tokens);
for (std::size_t index = 0; index < max_new_tokens; ++index) {
const auto found = std::max_element(logits.begin(), logits.end());
const auto token = static_cast<std::uint32_t>(found - logits.begin());
if (token == impl_->package.header().eos_token) break;
output.push_back(token);
Eval(token, logits);
}
return output;
}
std::size_t Decoder::position() const noexcept { return impl_->current_position; }
bool Decoder::has_dynamic_batch() const noexcept {
return batch_granularity() != 0;
}
std::size_t Decoder::batch_granularity() const noexcept {
if (std::getenv("LING3_GDN_CPU_PREFILL") != nullptr) return 1;
if (std::getenv("LING3_GDN_PREFILL_DIR") != nullptr) return 16;
return 0;
}
bool Decoder::has_batch32() const noexcept {
return has_dynamic_batch();
}
} // namespace ling3