#include <gtest/gtest.h>
#include "cpu/memory_manager.h"
#include <cstring>
namespace cpu_inference {
class MemoryManagerTest : public ::testing::Test {
protected:
void SetUp() override {
MemoryManager::get();
}
void TearDown() override {
MemoryManager::get();
}
};
template<typename T>
std::vector<T> read_from_tensor(const Tensor& tensor, int numa_node, size_t count) {
const void* ptr = tensor.at(numa_node, {0});
if (!ptr) return {};
std::vector<T> out(count);
std::memcpy(out.data(), ptr, count * sizeof(T));
return out;
}
template<typename T>
bool write_to_tensor(Tensor& tensor, int numa_node, const std::vector<T>& data) {
void* ptr = tensor.at(numa_node, {0});
if (!ptr) return false;
std::memcpy(ptr, data.data(), data.size() * sizeof(T));
return true;
}
TEST_F(MemoryManagerTest, AllocSingleNuma_Basic) {
auto& manager = MemoryManager::get();
std::vector<float> input = {1.0f, 2.0f, 3.0f, 4.0f};
Tensor& tensor = manager.alloc_single_numa("single", {4}, 0, "float32");
ASSERT_TRUE(write_to_tensor(tensor, 0, input));
auto output = read_from_tensor<float>(tensor, 0, input.size());
EXPECT_EQ(input, output);
}
TEST_F(MemoryManagerTest, AllocMultiNuma_TwoNodes) {
auto& manager = MemoryManager::get();
int tot_numa_cnt = manager.get_tot_numa_cnt();
printf("Total NUMA nodes available: %d\n", tot_numa_cnt);
if (tot_numa_cnt < 2) {
GTEST_SKIP() << "Only " << tot_numa_cnt << " NUMA node(s) available.";
}
Tensor& tensor = manager.alloc_multi_numa("multi", {5}, {0, 1}, "float32");
std::vector<float> data0 = {10.0f, 20.0f, 30.0f, 40.0f, 50.0f};
std::vector<float> data1 = {40.0f, 50.0f, 10.0f, 20.0f, 30.0f};
ASSERT_TRUE(write_to_tensor(tensor, 0, data0));
ASSERT_TRUE(write_to_tensor(tensor, 1, data1));
auto out0 = read_from_tensor<float>(tensor, 0, data0.size());
auto out1 = read_from_tensor<float>(tensor, 1, data1.size());
EXPECT_EQ(data0, out0);
EXPECT_EQ(data1, out1);
}
TEST_F(MemoryManagerTest, AllocAllNuma_Basic) {
auto& manager = MemoryManager::get();
Tensor& tensor = manager.alloc_all_numa("all", {8}, "float32");
size_t total = 8;
for (int node = 0; node < 4; ++node) {
std::vector<float> block(total, static_cast<float>(100 + node));
ASSERT_TRUE(write_to_tensor(tensor, node, block));
}
auto first = read_from_tensor<float>(tensor, 0, 1);
EXPECT_FLOAT_EQ(first[0], 100.0f);
}
TEST_F(MemoryManagerTest, GetMemory_Exists) {
auto& manager = MemoryManager::get();
manager.alloc_single_numa("lookup", {3}, 0, "float32");
Tensor& tensor = manager.get_memory("lookup");
std::vector<float> in = {1.5f, 2.5f, 3.5f};
ASSERT_TRUE(write_to_tensor(tensor, 0, in));
auto out = read_from_tensor<float>(tensor, 0, in.size());
EXPECT_EQ(in, out);
}
TEST_F(MemoryManagerTest, AllocSingleNuma_NameConflict) {
auto& manager = MemoryManager::get();
manager.alloc_single_numa("dup", {1}, 0, "float32");
EXPECT_THROW(
manager.alloc_single_numa("dup", {1}, 0, "float32"),
std::runtime_error
);
}
TEST_F(MemoryManagerTest, AllocSingleNuma_InvalidNuma) {
auto& manager = MemoryManager::get();
EXPECT_THROW(
manager.alloc_single_numa("bad", {1}, -1, "float32"),
std::out_of_range
);
EXPECT_THROW(
manager.alloc_single_numa("bad", {1}, manager.get_tot_numa_cnt(), "float32"),
std::out_of_range
);
}
}