File size: 2,457 Bytes
3fd1a35 3706f1c 3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 | #pragma once
#include "ling3/model_package.h"
#include "ling3/mla_stats.h"
#include "ling3/decoder_state.h"
#include <cstddef>
#include <cstdint>
#include <memory>
#include <span>
#include <vector>
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<DecoderCheckpoint> SaveCheckpoint();
std::size_t RestoreCheckpoint(const DecoderCheckpoint & checkpoint);
static std::size_t CheckpointBytes(const DecoderCheckpoint & checkpoint);
std::shared_ptr<const DecoderState> SaveState();
std::size_t RestoreState(const DecoderState & state);
const std::string & StateSignature() const;
DecodeTimings Eval(std::uint32_t token, std::span<float> 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<float> logits);
#endif
DecodeTimings EvalBatch(
std::span<const std::uint32_t> tokens,
std::span<float> logits);
DecodeTimings EvalBatchState(std::span<const std::uint32_t> tokens);
void PrepareBatch(std::size_t rows);
DecodeTimings EvalBatch32(
std::span<const std::uint32_t> tokens,
std::span<float> logits);
DecodeTimings EvalBatch32State(std::span<const std::uint32_t> tokens);
void PrepareBatch32();
std::vector<std::uint32_t> Generate(
std::span<const std::uint32_t> 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> impl_;
};
} // namespace ling3
|