#pragma once #include "ling3/model_package.h" #include "ling3/mla_stats.h" #include "ling3/decoder_state.h" #include #include #include #include #include namespace ling3 { // Owns recurrent/conv state only. MLA KV remains in its originating decoder. // Reset or overwriting any part of its prefix invalidates the checkpoint. struct DecoderCheckpoint; struct DecodeTimings { double layers_ms = 0.0; double output_head_ms = 0.0; double total_ms = 0.0; }; class Decoder { public: explicit Decoder(const ModelPackage & package, std::size_t context_capacity = 0); ~Decoder(); Decoder(const Decoder &) = delete; Decoder & operator=(const Decoder &) = delete; void Reset(); MlaBackendStats AttentionStats() const; std::shared_ptr SaveCheckpoint(); std::size_t RestoreCheckpoint(const DecoderCheckpoint & checkpoint); static std::size_t CheckpointBytes(const DecoderCheckpoint & checkpoint); std::shared_ptr SaveState(); std::size_t RestoreState(const DecoderState & state); const std::string & StateSignature() const; DecodeTimings Eval(std::uint32_t token, std::span logits); #if LING3_EXPERIMENTAL_MTP // Only the explicit experiment initializes MTP; normal chat does not. bool EnableMtp(); // Experimental NEXTN probe; requires a contiguous MTP prefix. Restored // trunk states and batch prefill alone do not populate the MTP cache. bool HasMtp() const noexcept; DecodeTimings EvalMtp(std::uint32_t next_token, std::span logits); #endif DecodeTimings EvalBatch( std::span tokens, std::span logits); DecodeTimings EvalBatchState(std::span tokens); void PrepareBatch(std::size_t rows); DecodeTimings EvalBatch32( std::span tokens, std::span logits); DecodeTimings EvalBatch32State(std::span tokens); void PrepareBatch32(); std::vector Generate( std::span prompt, std::size_t maximum_new_tokens); std::size_t position() const noexcept; bool has_dynamic_batch() const noexcept; std::size_t batch_granularity() const noexcept; bool has_batch32() const noexcept; private: struct Impl; std::unique_ptr impl_; }; } // namespace ling3