Download src/decoder_state.cpp from Sariel00/Ling-3.0-tiny-RKNN: direct link, hf CLI and curl.
- Browser
- Download file 6.66 kB
-
https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/decoder_state.cpp
- Command line
-
hf download hf://Sariel00/Ling-3.0-tiny-RKNN/src/decoder_state.cpp
-
curl -L -o decoder_state.cpp https://huggingface.co/Sariel00/Ling-3.0-tiny-RKNN/resolve/main/src/decoder_state.cpp
6.66 kB
| 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 | |