#include "ling3/decoder_state.h" #include "ling3/model_package.h" #include #include #include #include #include #include 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 void Write(std::ostream & out, T value) { out.write(reinterpret_cast(&value), sizeof(value)); if (!out) throw std::runtime_error("state write failed"); } template T Read(std::istream & in) { T value {}; in.read(reinterpret_cast(&value), sizeof(value)); Check(bool(in)); return value; } template void WriteVector(std::ostream & out, const std::vector & values) { Write(out, values.size()); Write(out, Crc32(reinterpret_cast(values.data()), values.size()*sizeof(T))); out.write(reinterpret_cast(values.data()), values.size()*sizeof(T)); if (!out) throw std::runtime_error("state buffer write failed"); } template std::vector ReadVector(std::istream & in, std::size_t maximum) { const auto count = Read(in); Check(count <= maximum && count <= std::numeric_limits::max()/sizeof(T)); const auto crc = Read(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 values(count); in.read(reinterpret_cast(values.data()), count*sizeof(T)); Check(bool(in)); Check(Crc32(reinterpret_cast(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(out, 0x0001315453334cULL); // L3ST1, format v1 Write(out, state.position); WriteVector(out, std::vector(state.signature.begin(), state.signature.end())); for (const auto & l : state.layers) { Write(out, l.gdn.fp32_valid); Write(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(in) == 0x0001315453334cULL); DecoderState state; state.position = Read(in); Check(state.position <= capacity && state.position <= 262144); auto key = ReadVector(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(in); Check(valid <= 1 && Read(in) == 1-valid); l.gdn.fp32_valid = valid; l.gdn.fp16 = ReadVector(in, mla ? 0 : kStateElements); l.gdn.fp32 = ReadVector(in, mla ? 0 : kStateElements); for (auto & v : l.conv) v = ReadVector(in, mla ? 0 : kConvElements); l.keys = ReadVector(in, mla ? state.position*16*192 : 0); l.values = ReadVector(in, mla ? state.position*16*128 : 0); } Check(in.peek() == std::char_traits::eof()); ValidateDecoderState(state, signature, capacity); return state; } } // namespace ling3