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