File size: 1,882 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
#pragma once

#include "ling3/model_format.h"

#include <cstddef>
#include <atomic>
#include <filesystem>
#include <string_view>
#include <unordered_map>
#include <vector>

namespace ling3 {

struct TensorView {
    std::string_view name;
    const TensorEntry * entry = nullptr;
    const std::byte * data = nullptr;
    const std::byte * aux = nullptr;
};

class ModelPackage {
public:
    explicit ModelPackage(const std::filesystem::path & path);
    ~ModelPackage();

    ModelPackage(const ModelPackage &) = delete;
    ModelPackage & operator=(const ModelPackage &) = delete;

    const PackageHeader & header() const noexcept { return *header_; }
    const std::vector<TensorView> & tensors() const noexcept { return tensors_; }
    const TensorView & tensor(std::string_view name) const;
    std::size_t mapped_bytes() const noexcept { return mapped_bytes_; }

    // After a consumer has synchronously copied a linear weight, discard only
    // whole source pages inside that tensor. The read-only file remains valid.
    void DiscardCopiedLinearWeight(const TensorView & weight) const;
    std::size_t discarded_weight_bytes() const noexcept { return discarded_weight_bytes_; }
    std::size_t cache_advice_failures() const noexcept { return cache_advice_failures_; }
    std::size_t resident_linear_weight_bytes() const;

private:
    void Reset() noexcept;

    int fd_ = -1;
    const std::byte * mapping_ = nullptr;
    std::size_t mapped_bytes_ = 0;
    const PackageHeader * header_ = nullptr;
    std::vector<TensorView> tensors_;
    std::unordered_map<std::string_view, std::size_t> tensor_index_;
    mutable std::atomic_size_t discarded_weight_bytes_ {0};
    mutable std::atomic_size_t cache_advice_failures_ {0};
};

void ValidateLing3Tiny(const PackageHeader & header);
std::uint32_t Crc32(const std::byte * data, std::size_t bytes);

} // namespace ling3