Download src/decoder.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 94.3 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/decoder.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/decoder.cpp
-
curl -L -o decoder.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/decoder.cpp
94.3 kB
| 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); | |
| constexpr int maximum_layer = 24; | |
| constexpr int maximum_layer = 23; | |
| if (error != std::errc {} || parsed_end == begin || layer < 0 || layer > maximum_layer) { | |
| throw std::runtime_error("cannot derive W4 IOMMU domain from " + base); | |
| } | |
| iommu_domain_id = layer == 24 ? 15 : 2 + layer / 2; | |
| iommu_domain_id = 2 + layer / 2; | |
| } 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) { | |
| 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)); | |
| } | |
| 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) { | |
| 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; | |
| } | |
| 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; | |
| }; | |
| // 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_; | |
| }; | |
| 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; | |
| std::unique_ptr<MtpLayer> mtp; | |
| std::vector<float> mtp_hidden, mtp_embedding; | |
| std::size_t mtp_position = 0; | |
| bool mtp_trunk_ready = false; | |
| 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"); | |
| } | |
| 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; | |
| } | |
| 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 (mtp) mtp->Reset(); | |
| mtp_position = 0; | |
| mtp_trunk_ready = false; | |
| } | |
| 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; | |
| mtp_trunk_ready = false; | |
| 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; | |
| mtp_trunk_ready = true; | |
| 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; | |
| mtp_trunk_ready = compute_logits; | |
| return { | |
| Milliseconds(begin, layers_end), | |
| Milliseconds(layers_end, end), | |
| Milliseconds(begin, end), | |
| }; | |
| } | |
| 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)}; | |
| } | |
| 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); | |
| } | |
| 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); | |
| } | |
| 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 | |