#pragma once #include "common/file_source.h" #include "common/thread_pool.h" #include "executor/allocator.h" #include #include namespace mtk { // Shared weights used by a single model chunk. struct SharedWeights { std::vector files; std::vector buffers; // Helper functions explicit operator bool() const; bool empty() const; size_t size() const; bool isPreloaded(const size_t swIndex) const; }; // A global shared weights handle that can exist outside of LLM Runtime class SharedWeightsHandle { public: explicit SharedWeightsHandle(const std::vector& sharedWeightsFiles, const size_t numDlaChunks = 1); ~SharedWeightsHandle(); void setPreloadSubset(const std::unordered_set& subsetIndexes, const bool repeatAllChunks = false); void preload(const bool async = false); static void preloadUnion(const std::vector& preloadHandles, const std::vector& refHandles = {}); static void preloadUnion(const std::vector>& preloadHandles, const std::vector>& refHandles = {}); bool isPreloaded(const size_t swIndex) const; void unload(const size_t swIndex); void wait() const; SharedWeights getSharedWeights(const size_t dlaChunkIndex) const; private: const size_t kNumDlaChunks; std::shared_ptr mAllocator; std::vector> mSharedWeightsBuffers; std::vector kSharedWeightsFiles; std::unordered_set mPreloadIndexes; std::vector mPreloadedMap; mutable BasicThreadPool mThreadPool; }; } // namespace mtk