#pragma once
#include <cstddef>
#include <cstdint>
#include <list>
#include <unordered_map>
inline constexpr size_t kAlignBytes = 256; inline constexpr size_t kMinSegment = size_t(1) << 20; inline constexpr size_t kMinSplit = 256;
inline size_t align_up(size_t v, size_t a) { return (v + a - 1) & ~(a - 1); }
struct Stats {
size_t allocated_bytes = 0; size_t reserved_bytes = 0; size_t peak_bytes = 0; size_t live_blocks = 0; size_t free_blocks = 0; size_t segments = 0; size_t cuda_malloc_calls = 0; size_t cuda_free_calls = 0; };
struct Segment;
struct Block {
Segment* seg; size_t offset; size_t size; bool free; size_t id; Block* prev; Block* next; };
struct Segment {
void* ptr; size_t size; size_t live_bytes; Block* first; Segment* next; };
class CachingAllocator {
public:
CachingAllocator() = default;
~CachingAllocator();
CachingAllocator(const CachingAllocator&) = delete;
CachingAllocator& operator=(const CachingAllocator&) = delete;
void* alloc(size_t bytes); void free(void* p); void empty_cache(); Stats stats() const;
void assert_consistency() const;
private:
Block* best_fit(size_t bytes);
void split_block(Block* b, size_t take); void insert_free(Block* b);
void remove_free(Block* b);
Block* add_segment(size_t bytes); void coalesce(Block* b);
std::list<Block*> free_; std::unordered_map<void*, Block*> by_ptr_; Segment* segments_ = nullptr;
Segment* tail_ = nullptr;
size_t id_counter_ = 0;
size_t allocated_ = 0;
size_t reserved_ = 0;
size_t peak_ = 0;
size_t live_blocks_ = 0;
size_t free_blocks_ = 0;
size_t segments_n_ = 0;
size_t n_malloc_ = 0;
size_t n_free_ = 0;
};