Ling-3.0-tiny-RKNN / src /decoder_state.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
6.66 kB
#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