#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 #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #if defined(__aarch64__) #include #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(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 std::span Typed(const TensorView & tensor, DataType type) { if (tensor.entry->dtype != static_cast(type) || tensor.entry->data_bytes % sizeof(T) != 0) { throw std::runtime_error(std::string(tensor.name) + " has an incompatible dtype"); } return { reinterpret_cast(tensor.data), static_cast(tensor.entry->data_bytes / sizeof(T)), }; } std::span Blob(const TensorView & tensor, TensorRole role) { if (tensor.entry->role != static_cast(role)) { throw std::runtime_error(std::string(tensor.name) + " has an incompatible role"); } return {tensor.data, static_cast(tensor.entry->data_bytes)}; } std::vector DecodeBf16(const TensorView & tensor, std::size_t expected) { const auto input = Typed(tensor, DataType::kBFloat16); if (input.size() != expected) { throw std::runtime_error(std::string(tensor.name) + " has an incompatible shape"); } std::vector output(expected); for (std::size_t index = 0; index < expected; ++index) { output[index] = BFloat16ToFloat(input[index]); } return output; } std::vector DecodeFloat( const TensorView & tensor, std::size_t expected) { if (tensor.entry->dtype == static_cast(DataType::kFloat32)) { const auto input = Typed(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 MakeLinear( const ModelPackage & package, const std::string & base, std::vector cores = {0, 1, 2}) { const auto & weight = package.tensor(base + ".weight"); const bool mixed_bf16 = (package.header().flags & kPackageMixedW4W8) && weight.entry->dtype == static_cast(DataType::kBFloat16) && weight.entry->layout == static_cast(TensorLayout::kRowMajor); if (weight.entry->rank != 2 || (!(package.header().flags & kPackageOfficialInt4) && !mixed_bf16 && ( weight.entry->dtype != static_cast(DataType::kInt4Low) || weight.entry->layout != static_cast(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( package, base, W4LinearConfig { static_cast(weight.entry->dims[0]), static_cast(weight.entry->dims[1]), static_cast(weight.entry->flags), std::move(cores), iommu_domain_id, }); return linear; } void NormalizeHeads(std::span 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 input, std::span weight, std::span state, std::span output) { for (int channel = 0; channel < kKdaWidth; ++channel) { auto * history = state.data() + static_cast(channel) * 3; const auto * kernel = weight.data() + static_cast(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 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 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 shared_projected, shared_hidden, shared_output; std::array, 3> lane_input, lane_projected, lane_hidden, lane_output; std::vector quantized; std::vector contributions; SparseBatchScratch() : shared_projected(kMaxBatch * 2 * kExpertWidth), shared_hidden(kMaxBatch * kExpertWidth), shared_output(kMaxBatch * kHidden) {} }; struct LayerBatchScratch { std::vector 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 input, std::size_t position, std::span output) = 0; virtual void RunBatch( std::span input, std::size_t rows, std::size_t position, std::span output) = 0; virtual void PrepareBatch(std::size_t rows) = 0; }; class KdaAttention final : public Attention { public: KdaAttention( const ModelPackage & package, int layer, std::span heads6, std::span 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 input, std::size_t, std::span output) override { projection_->Run(input, projected_); CausalConvSilu( std::span(projected_).subspan(0, kKdaWidth), q_conv_, conv_state_[0], q_); CausalConvSilu( std::span(projected_).subspan(kKdaWidth, kKdaWidth), k_conv_, conv_state_[1], k_); CausalConvSilu( std::span(projected_).subspan(2 * kKdaWidth, kKdaWidth), v_conv_, conv_state_[2], v_); NormalizeHeads(q_); NormalizeHeads(k_); const auto f = std::span(projected_).subspan(3 * kKdaWidth, kKdaWidth); const auto gate = std::span(projected_).subspan(4 * kKdaWidth, kKdaWidth); const auto beta_logits = std::span(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(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 input, std::size_t rows, std::size_t, std::span 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(batch_projected_).first(rows * 10304)); const auto projected_at = Clock::now(); constexpr std::array 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(channel) * 3; const auto & weights = stream == 0 ? q_conv_ : (stream == 1 ? k_conv_ : v_conv_); const auto * kernel = weights.data() + static_cast(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(batch_q_).first(rows * kKdaWidth), std::span(batch_k_).first(rows * kKdaWidth), std::span(batch_v_).first(rows * kKdaWidth), std::span(batch_decay_).first(rows * kKdaWidth), std::span(batch_beta_).first(rows * kHeads), std::span(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(batch_q_).subspan(chunk * chunk_vectors, chunk_vectors), std::span(batch_k_).subspan(chunk * chunk_vectors, chunk_vectors), std::span(batch_v_).subspan(chunk * chunk_vectors, chunk_vectors), std::span(batch_decay_).subspan(chunk * chunk_vectors, chunk_vectors), std::span(batch_beta_).subspan(chunk * chunk_betas, chunk_betas), std::span(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(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(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 projection_; std::unique_ptr output_projection_; std::vector q_conv_, k_conv_, v_conv_, a_log_, dt_bias_, output_norm_; GdnStep gdn_; std::array, 3> conv_state_; std::vector projected_, q_, k_, v_, decay_, beta_, recurrence_, gated_; std::vector &batch_projected_, &batch_q_, &batch_k_, &batch_v_, &batch_decay_; std::vector &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 input, std::size_t position, std::span 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 rotated_key {}; RotateInterleaved( std::span(projected_).subspan( kMlaQueryRank + kMlaKvRank, kMlaRotaryWidth), position, rotated_key); for (int head = 0; head < kHeads; ++head) { std::array rotated_query {}; RotateInterleaved( std::span(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(kMlaQueryWidth)); constexpr std::array 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::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 input, std::size_t rows, std::size_t position, std::span 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(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(batch_q_rank_).first(rows * kMlaQueryRank), rows, std::span(batch_q_all_).first(rows * q_width)); kv_projection_->RunBatch( std::span(batch_kv_rank_).first(rows * kMlaKvRank), rows, std::span(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(batch_rotated_key_).subspan( row * kMlaRotaryWidth, kMlaRotaryWidth)); } constexpr std::array 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 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(batch_q_all_).first(rows * q_width), std::span(key_cache_).first((position + rows) * q_width), std::span(value_cache_).first((position + rows) * attention_width)); const bool used_npu = npu_.Run(std::span(batch_q_all_).first(rows*q_width), std::span(key_cache_).first((position+rows)*q_width), std::span(value_cache_).first((position+rows)*attention_width), rows,position+rows,std::span(batch_attention_).first(rows*attention_width)); const float scale = 1.0F / std::sqrt(static_cast(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::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;rowRunBatch( std::span(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 input, std::size_t position, std::span output) { std::array 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(frequency) / static_cast(kMlaRotaryWidth)); const float angle = static_cast(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 projection_, q_projection_, kv_projection_, output_projection_; std::vector q_norm_, kv_norm_; std::size_t max_context_; std::vector projected_, q_rank_, kv_rank_, q_all_, kv_all_, attention_, scores_; std::vector key_cache_, value_cache_; std::vector &batch_projected_, &batch_q_rank_, &batch_kv_rank_; std::vector &batch_q_all_, &batch_kv_all_, &batch_attention_, &batch_rotated_key_; std::array, 4> batch_scores_; MlaNpu & npu_; }; class FeedForward { public: virtual ~FeedForward() = default; virtual void Run(std::span input, std::span output) = 0; virtual void RunBatch( std::span input, std::size_t rows, std::span 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 input, std::span output) override { gate_up_->Run(input, projected_); SiluMultiply(projected_.data(), projected_.data() + kDenseWidth, hidden_.data(), kDenseWidth); down_->Run(hidden_, output); } void RunBatch( std::span input, std::size_t rows, std::span 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(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(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 gate_up_, down_; std::vector projected_, hidden_; std::vector 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 {0, 1, 2} : std::vector {expert % 3})), down_(MakeLinear( package, "model.layers." + std::to_string(layer) + ".mlp.experts." + std::to_string(expert) + ".down_proj", all_cores ? std::vector {0, 1, 2} : std::vector {expert % 3})), projected_(2 * kExpertWidth), hidden_(kExpertWidth), output_(kHidden), current_core_(all_cores ? -1 : expert % 3) {} std::span Run( std::span 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 input) { return gate_up_->PrepareInput(input); } void ShareGateInputFrom(Expert & owner) { gate_up_->ShareInputFrom(*owner.gate_up_); } void RunGateBatch( std::span input, std::size_t rows, const Expert & weights, std::span output) { gate_up_->RunBatchWithWeights(input, rows, *weights.gate_up_, output); } void RunDownBatch( std::span input, std::size_t rows, const Expert & weights, std::span output) { down_->RunBatchWithWeights(input, rows, *weights.down_, output); } void RunGateQuantizedRows( std::span input, std::span scales, std::span rows, const Expert & weights, std::span 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 gate_up_, down_; std::vector 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 input, std::span 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 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 input, std::size_t rows, std::span 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 router_workers; for (std::size_t worker = 0; worker < router_workers.size(); ++worker) { router_workers[worker] = std::thread([this, input, rows, worker]() { std::array 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(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 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( std::max_element(counts.begin(), counts.end()) - counts.begin()); const int target = static_cast( 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(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(batch_shared_hidden_).first(rows * kExpertWidth), rows, std::span(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 active_experts; active_experts.reserve(kExpertCount); for (int id = 0; id < static_cast(kExpertCount); ++id) { if (batch_groups_[id].empty()) continue; if (!experts_[id]) { experts_[id] = std::make_unique(package_, layer_, id, false); } active_experts.push_back(id); } for (int core = 0; core < 3; ++core) { if (!experts_[core]) { experts_[core] = std::make_unique(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 lane_load {}; std::array lane_cost {}; std::array lane_experts {}; for (int id : active_experts) { const int lane = static_cast( 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 cores {0, 1, 2}; std::array 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(indices).first(count), *experts_[id], std::span(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(kHidden) * sizeof(float)); } runner.RunGateBatch( std::span(gathered).first(count * kHidden), count, *experts_[id], std::span(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(hidden).first(count * kExpertWidth), count, *experts_[id], std::span(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(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(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(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 errors {}; std::array workers; for (int core = 0; core < 3; ++core) { workers[core] = std::thread([this, core, &errors]() { try { for (int id = core; id < static_cast(kExpertCount); id += 3) { experts_[id] = std::make_unique( 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( package_, layer_, id, all_core_experts_); } const auto value = experts_[id]->Run( {active_input_, static_cast(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 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(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( std::max_element(counts.begin(), counts.end()) - counts.begin()); const int target = static_cast( 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(package_, layer_, id, true); } const auto value = experts_[id]->Run( {active_input_, static_cast(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 gate_weight_, expert_bias_; std::unique_ptr shared_gate_up_, shared_down_; std::vector shared_projected_, shared_hidden_, shared_output_, logits_; std::vector &batch_shared_projected_, &batch_shared_hidden_, &batch_shared_output_; std::array, kExpertCount> experts_; std::array, 3> lane_output_; Route active_route_; std::array active_lanes_ {}; std::array, kExpertCount> batch_groups_; std::array batch_routes_; std::array, kMaxBatch> batch_route_lanes_ {}; std::array, 3> batch_jobs_; std::array, 3> &batch_lane_input_; std::vector &batch_quantized_; std::array batch_input_scales_ {}; std::array, 3> batch_lane_rows_ {}; std::array, 3> &batch_lane_projected_; std::array, 3> &batch_lane_hidden_; std::array, 3> &batch_lane_expert_output_; std::vector &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(package, 24, max_context, scratch.mla)), feed_forward_(std::make_unique(package, 24, scratch.sparse)) {} void Reset() { attention_->Reset(); } void Run(std::span token_embedding, std::span trunk_hidden, std::size_t position, std::span 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 input_norm_, post_norm_, embedding_norm_, hidden_norm_, final_norm_; std::vector embedding_hidden_, projected_, residual_, normalized_, attention_output_, ffn_output_; std::unique_ptr eh_projection_; std::unique_ptr attention_; std::unique_ptr feed_forward_; }; #endif class DecoderLayer { public: DecoderLayer( const ModelPackage & package, int layer, std::size_t max_context, std::span heads6, std::span 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(std::make_unique(package, layer, max_context, scratch.mla)) : std::unique_ptr(std::make_unique(package, layer, heads6, heads5, scratch.kda))), feed_forward_(layer == 0 ? std::unique_ptr(std::make_unique(package, layer)) : std::unique_ptr(std::make_unique(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 & 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 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(batch_normalized_).first(rows * kHidden), rows, position, std::span(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(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(batch_normalized_).first(rows * kHidden), rows, std::span(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(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 input_norm_, post_norm_; std::unique_ptr attention_; std::unique_ptr feed_forward_; std::vector normalized_, attention_output_, ffn_output_; std::vector &batch_normalized_, &batch_attention_output_, &batch_ffn_output_; }; } // namespace struct DecoderCheckpoint { std::weak_ptr owner; std::size_t position = 0; bool valid = true; std::vector layers; }; struct Decoder::Impl { const ModelPackage & package; std::size_t context_capacity; std::span embeddings; std::vector final_norm; DecoderScratch scratch; // Declared before layers: outlives all borrowers. std::vector> layers; #if LING3_EXPERIMENTAL_MTP std::unique_ptr mtp; std::vector mtp_hidden, mtp_embedding; std::size_t mtp_position = 0; bool mtp_trunk_ready = false; #endif std::unique_ptr lm_head; std::vector hidden; std::vector normalized; std::vector batch_hidden; std::size_t current_position = 0; std::shared_ptr checkpoint_owner = std::make_shared(0); std::vector> 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( 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(&package.header()), sizeof(PackageHeader)); for (const auto & tensor : package.tensors()) { state_signature.append(reinterpret_cast(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( 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(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 SaveCheckpoint() { std::erase_if(checkpoints, [](const auto & weak) { return weak.expired(); }); auto checkpoint = std::make_shared(); 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 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(); 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 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(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 tokens, std::span 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(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(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 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(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(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 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 logits) { return impl_->EvalMtp(token, logits); } #endif DecodeTimings Decoder::EvalBatch( std::span tokens, std::span logits) { return impl_->EvalBatch(tokens, logits, true); } DecodeTimings Decoder::EvalBatchState(std::span tokens) { return impl_->EvalBatch(tokens, {}, false); } std::shared_ptr Decoder::SaveCheckpoint() { return impl_->SaveCheckpoint(); } std::size_t Decoder::RestoreCheckpoint(const DecoderCheckpoint & checkpoint) { return impl_->RestoreCheckpoint(checkpoint); } std::shared_ptr 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 tokens, std::span 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 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 Decoder::Generate( std::span prompt, std::size_t max_new_tokens) { if (prompt.empty()) throw std::invalid_argument("prompt must contain at least one token"); Reset(); std::vector logits(impl_->package.header().vocab_size); for (std::uint32_t token : prompt) Eval(token, logits); std::vector 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(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