File size: 6,657 Bytes
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
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
#include "ling3/decoder_state.h"
#include "ling3/model_package.h"
#include <bit>
#include <cmath>
#include <fstream>
#include <limits>
#include <stdexcept>
#include <unistd.h>

namespace ling3 {
namespace {
constexpr std::size_t kStateElements = 16 * 128 * 128;
constexpr std::size_t kConvElements = 16 * 128 * 3;
void Check(bool ok) { if (!ok) throw std::invalid_argument("invalid or incompatible decoder state"); }
template<class T> void Write(std::ostream & out, T value) {
    out.write(reinterpret_cast<const char *>(&value), sizeof(value));
    if (!out) throw std::runtime_error("state write failed");
}
template<class T> T Read(std::istream & in) {
    T value {}; in.read(reinterpret_cast<char *>(&value), sizeof(value));
    Check(bool(in)); return value;
}
template<class T> void WriteVector(std::ostream & out, const std::vector<T> & values) {
    Write<std::uint64_t>(out, values.size());
    Write(out, Crc32(reinterpret_cast<const std::byte *>(values.data()), values.size()*sizeof(T)));
    out.write(reinterpret_cast<const char *>(values.data()), values.size()*sizeof(T));
    if (!out) throw std::runtime_error("state buffer write failed");
}
template<class T> std::vector<T> ReadVector(std::istream & in, std::size_t maximum) {
    const auto count = Read<std::uint64_t>(in);
    Check(count <= maximum && count <= std::numeric_limits<std::size_t>::max()/sizeof(T));
    const auto crc = Read<std::uint32_t>(in);
    // Reject truncated files before allocating a potentially large KV buffer.
    const auto start = in.tellg(); in.seekg(0, std::ios::end); const auto end = in.tellg();
    Check(start >= 0 && end >= start && std::uint64_t(end-start) >= count*sizeof(T));
    in.seekg(start);
    std::vector<T> values(count);
    in.read(reinterpret_cast<char *>(values.data()), count*sizeof(T)); Check(bool(in));
    Check(Crc32(reinterpret_cast<const std::byte *>(values.data()), count*sizeof(T)) == crc);
    return values;
}
void HostCheck() { Check(std::endian::native == std::endian::little && sizeof(float) == 4); }
} // namespace

std::size_t AttentionState::bytes() const {
    std::size_t total = (gdn.fp16.size()+keys.size()+values.size())*2 + gdn.fp32.size()*4;
    for (const auto & v : conv) total += v.size()*4;
    return total;
}
std::size_t DecoderState::bytes() const {
    std::size_t total = signature.size();
    for (const auto & layer : layers) total += layer.bytes();
    return total;
}
void ValidateDecoderState(const DecoderState & state, std::string_view signature, std::size_t capacity) {
    Check(state.signature == signature && !signature.empty());
    Check(state.position <= capacity && state.position <= 262144 && state.layers.size() == 24);
    for (std::size_t i = 0; i < state.layers.size(); ++i) {
        const auto & l = state.layers[i];
        if ((i+1)%4 == 0) {
            Check(l.keys.size() == state.position*16*192 && l.values.size() == state.position*16*128);
            Check(l.gdn.fp16.empty() && l.gdn.fp32.empty() && !l.gdn.fp32_valid);
            for (const auto & v : l.conv) Check(v.empty());
            for (auto v : l.keys) Check((v & 0x7f80) != 0x7f80);
            for (auto v : l.values) Check((v & 0x7f80) != 0x7f80);
        } else {
            Check(l.keys.empty() && l.values.empty());
            Check(l.gdn.fp16.size() == kStateElements);
            Check(l.gdn.fp32.empty() || l.gdn.fp32.size() == kStateElements);
            Check(!l.gdn.fp32_valid || l.gdn.fp32.size() == kStateElements);
            // FP16 shadow can overflow while the authoritative FP32 state is
            // finite. Validate the representation actually used on restore.
            if (!l.gdn.fp32_valid) for (auto v : l.gdn.fp16) Check((v & 0x7c00) != 0x7c00);
            for (auto v : l.gdn.fp32) Check(std::isfinite(v));
            for (const auto & v : l.conv) {
                Check(v.size() == kConvElements);
                for (auto x : v) Check(std::isfinite(x));
            }
        }
    }
}
void WriteDecoderState(const DecoderState & state, const std::filesystem::path & path) {
    HostCheck(); ValidateDecoderState(state, state.signature, 262144);
    std::string temporary = path.string()+".tmp-XXXXXX";
    const int fd = mkstemp(temporary.data());
    if (fd < 0) throw std::runtime_error("cannot create state temporary file");
    close(fd);
    try {
        std::ofstream out(temporary, std::ios::binary | std::ios::trunc);
        Write<std::uint64_t>(out, 0x0001315453334cULL); // L3ST1, format v1
        Write<std::uint64_t>(out, state.position);
        WriteVector(out, std::vector<char>(state.signature.begin(), state.signature.end()));
        for (const auto & l : state.layers) {
            Write<std::uint8_t>(out, l.gdn.fp32_valid);
            Write<std::uint8_t>(out, !l.gdn.fp32_valid);
            WriteVector(out, l.gdn.fp16); WriteVector(out, l.gdn.fp32);
            for (const auto & v : l.conv) WriteVector(out, v);
            WriteVector(out, l.keys); WriteVector(out, l.values);
        }
        out.flush(); if (!out) throw std::runtime_error("state flush failed");
        out.close(); if (!out) throw std::runtime_error("state close failed");
        std::filesystem::rename(temporary, path);
    } catch (...) { std::filesystem::remove(temporary); throw; }
}
DecoderState ReadDecoderState(const std::filesystem::path & path, std::string_view signature,
                              std::size_t capacity) {
    HostCheck(); std::ifstream in(path, std::ios::binary);
    Check(Read<std::uint64_t>(in) == 0x0001315453334cULL);
    DecoderState state; state.position = Read<std::uint64_t>(in);
    Check(state.position <= capacity && state.position <= 262144);
    auto key = ReadVector<char>(in, signature.size());
    state.signature.assign(key.begin(), key.end()); Check(state.signature == signature);
    state.layers.resize(24);
    for (std::size_t i = 0; i < 24; ++i) {
        auto & l = state.layers[i]; const bool mla = (i+1)%4 == 0;
        const auto valid = Read<std::uint8_t>(in);
        Check(valid <= 1 && Read<std::uint8_t>(in) == 1-valid); l.gdn.fp32_valid = valid;
        l.gdn.fp16 = ReadVector<std::uint16_t>(in, mla ? 0 : kStateElements);
        l.gdn.fp32 = ReadVector<float>(in, mla ? 0 : kStateElements);
        for (auto & v : l.conv) v = ReadVector<float>(in, mla ? 0 : kConvElements);
        l.keys = ReadVector<std::uint16_t>(in, mla ? state.position*16*192 : 0);
        l.values = ReadVector<std::uint16_t>(in, mla ? state.position*16*128 : 0);
    }
    Check(in.peek() == std::char_traits<char>::eof());
    ValidateDecoderState(state, signature, capacity); return state;
}
} // namespace ling3