#![cfg(not(exclusive_memory_only))]
use std::sync::Arc;
use cubecl_ir::MemoryDeviceProperties;
use cubecl_server::config::memory::{MemoryPoolConfig, MemoryPoolsConfig};
use cubecl_server::config::size::MemorySize;
use cubecl_server::dry_run::{DryRun, RealRun};
use cubecl_server::logging::ServerLogger;
use cubecl_server::memory_management::{
ErrorGraph, InstallMemoryPoolsError, MemoryAllocationMode, MemoryConfiguration,
MemoryManagement, MemoryManagementOptions, MemoryPoolKind,
};
use cubecl_server::storage::BytesStorage;
const MIB: u64 = 1024 * 1024;
fn sliced(page_size: u64, pages: u64) -> MemoryPoolsConfig {
MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(page_size),
max_slice_size: None,
max_pool_size: Some(MemorySize(page_size * pages)),
dealloc_period: None,
}])
}
fn props() -> MemoryDeviceProperties {
MemoryDeviceProperties::new(128 * MIB, 32)
}
fn manage(pools: &MemoryPoolsConfig) -> MemoryManagement<BytesStorage> {
let resolved = MemoryConfiguration::default()
.resolve(Some(pools), &props())
.unwrap();
MemoryManagement::from_configuration(
BytesStorage::default(),
&props(),
resolved,
Arc::new(ServerLogger::default()),
MemoryManagementOptions::new("Main GPU Memory"),
)
}
#[test]
fn programmatic_pools_override_runtime_default() {
let pools = sliced(MIB, 2);
let resolved = MemoryConfiguration::default()
.resolve(Some(&pools), &props())
.unwrap();
let mut memory_management = MemoryManagement::from_configuration(
BytesStorage::default(),
&props(),
resolved,
Arc::new(ServerLogger::default()),
MemoryManagementOptions::new("Main GPU Memory"),
);
let small = memory_management
.reserve(4096, &mut ErrorGraph::default())
.unwrap();
drop(small);
let _large = memory_management
.reserve(512 * 1024, &mut ErrorGraph::default())
.unwrap();
assert_eq!(memory_management.memory_usage().bytes_reserved, MIB);
let _fill_1 = memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap();
let _fill_2 = memory_management
.reserve(500 * 1024, &mut ErrorGraph::default())
.unwrap();
assert!(
memory_management
.reserve(MIB, &mut ErrorGraph::default())
.is_err()
);
assert_eq!(memory_management.memory_usage().bytes_reserved, 2 * MIB);
}
#[test]
fn capped_pool_respects_max_slice_size() {
let pools = MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: Some(MemorySize(64 * 1024)),
max_pool_size: Some(MemorySize(2 * MIB)),
dealloc_period: None,
},
MemoryPoolConfig::Sliced {
page_size: MemorySize(16 * MIB),
max_slice_size: None,
max_pool_size: Some(MemorySize(32 * MIB)),
dealloc_period: None,
},
]);
let resolved = MemoryConfiguration::default()
.resolve(Some(&pools), &props())
.unwrap();
let mut memory_management = MemoryManagement::from_configuration(
BytesStorage::default(),
&props(),
resolved,
Arc::new(ServerLogger::default()),
MemoryManagementOptions::new("Main GPU Memory"),
);
let strays: Vec<_> = (0..4)
.map(|_| {
memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap()
})
.collect();
assert_eq!(
memory_management.memory_usage().bytes_reserved,
16 * MIB,
"strays must land in the arena, not the capped small pool"
);
let _tiny = memory_management
.reserve(4096, &mut ErrorGraph::default())
.unwrap();
assert_eq!(memory_management.memory_usage().bytes_reserved, 17 * MIB);
drop(strays);
}
#[test]
fn configure_rebuilds_pools_in_place() {
let resolved = MemoryConfiguration::default()
.resolve(Some(&sliced(MIB, 2)), &props())
.unwrap();
let mut memory_management = MemoryManagement::from_configuration(
BytesStorage::default(),
&props(),
resolved,
Arc::new(ServerLogger::default()),
MemoryManagementOptions::new("Main GPU Memory"),
);
let live = memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap();
let bigger = MemoryConfiguration::default()
.resolve(Some(&sliced(4 * MIB, 2)), &props())
.unwrap();
assert!(
matches!(
memory_management.install_pools(bigger.clone(), &props(), &mut ErrorGraph::default()),
Err(InstallMemoryPoolsError::PoolsInUse { bytes_in_use }) if bytes_in_use > 0
),
"the refusal names the live bytes that caused it"
);
assert!(
memory_management
.reserve(2 * MIB, &mut ErrorGraph::default())
.is_err()
);
drop(live);
memory_management
.install_pools(bigger, &props(), &mut ErrorGraph::default())
.unwrap();
let _large = memory_management
.reserve(2 * MIB, &mut ErrorGraph::default())
.unwrap();
}
#[test]
fn measured_plan_cycle() {
let growable = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
}]);
let mut memory_management = manage(&growable);
let workload = |memory_management: &mut MemoryManagement<BytesStorage>| {
let a = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
let b = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
drop(a);
let c = memory_management
.reserve(900 * 1024, &mut ErrorGraph::default())
.unwrap();
drop(b);
drop(c);
};
workload(&mut memory_management);
let report = memory_management.memory_report();
let arena = &report.dynamic[0];
let MemoryPoolKind::Sliced { page_size, .. } = arena.kind else {
panic!("the arena is a sliced pool");
};
assert_eq!(arena.largest_alloc, 900 * 1024);
assert_eq!(arena.pages_peak, 2, "two pages while a, b overlap");
let capped = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(page_size),
max_slice_size: None,
max_pool_size: Some(MemorySize(page_size * arena.pages_peak)),
dealloc_period: None,
}]);
let resolved = MemoryConfiguration::default()
.resolve(Some(&capped), &props())
.unwrap();
memory_management.cleanup(true, &mut ErrorGraph::default());
memory_management
.install_pools(resolved, &props(), &mut ErrorGraph::default())
.unwrap();
workload(&mut memory_management);
let replayed = memory_management.memory_report();
assert_eq!(replayed.dynamic[0].pages_peak, arena.pages_peak);
}
#[test]
fn full_capped_pool_spills_to_tail() {
let pools = MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: Some(MemorySize(MIB)),
dealloc_period: None,
},
MemoryPoolConfig::Sliced {
page_size: MemorySize(4 * MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
},
]);
let mut memory_management = manage(&pools);
let _planned = memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap();
let _off_plan = memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap();
let report = memory_management.memory_report();
assert_eq!(report.dynamic[0].pages_peak, 1, "the arena stayed capped");
assert_eq!(report.dynamic[1].pages_peak, 1, "the tail caught the spill");
}
#[test]
#[serial_test::serial]
fn a_measurement_maps_only_what_it_resolves() {
let growable = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
}]);
let mut memory_management = manage(&growable);
let dry_run = DryRun::new();
let workload = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
let scratch = {
let _measurement = RealRun::new();
memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap()
};
let report = memory_management.memory_report();
assert_eq!(report.dynamic[0].pages, 2, "{report:?}");
assert_eq!(
report.dynamic[0].pages_unmapped, 2,
"neither is backed until something resolves it: {report:?}"
);
memory_management
.get_storage(scratch.clone().binding())
.unwrap();
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages_unmapped, 1,
"the measurement's page is backed and the workload's is not: {report:?}"
);
drop(workload);
drop(scratch);
drop(dry_run);
}
#[test]
#[serial_test::serial]
fn dry_run_reservations_stay_unmapped_until_resolved() {
let growable = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
}]);
let mut memory_management = manage(&growable);
let dry_run = DryRun::new();
let reserved = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
let report = memory_management.memory_report();
assert_eq!(report.dynamic[0].pages, 1, "the page was carved");
assert_eq!(
report.dynamic[0].pages_peak, 1,
"and counts toward the plan"
);
assert_eq!(
report.dynamic[0].pages_unmapped, 1,
"but has no device backing: {report:?}"
);
let storage = memory_management
.get_storage(reserved.clone().binding())
.unwrap();
assert_eq!(storage.size(), 600 * 1024);
assert_eq!(
memory_management.memory_report().dynamic[0].pages_unmapped,
0,
"resolution installed the backing"
);
drop(reserved);
drop(dry_run);
memory_management.cleanup(true, &mut ErrorGraph::default());
assert_eq!(memory_management.memory_report().dynamic[0].pages, 0);
}
#[test]
#[serial_test::serial]
fn dry_run_persistent_reservations_stay_unmapped() {
let growable = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
}]);
let mut memory_management = manage(&growable);
memory_management.mode(MemoryAllocationMode::Persistent);
let dry_run = DryRun::new();
let kv = memory_management
.reserve(512 * 1024, &mut ErrorGraph::default())
.unwrap();
let report = memory_management.memory_report();
assert_eq!(report.persistent.pages, 1);
assert_eq!(
report.persistent.pages_unmapped, 1,
"a dry-run persistent slice has no backing yet: {report:?}"
);
let storage = memory_management.get_storage(kv.clone().binding()).unwrap();
assert_eq!(storage.size(), 512 * 1024);
assert_eq!(
memory_management.memory_report().persistent.pages_unmapped,
0
);
drop(kv);
drop(dry_run);
memory_management.mode(MemoryAllocationMode::Auto);
}
#[test]
#[serial_test::serial]
fn a_warm_second_pass_measures_the_workload_alone() {
let growable = MemoryPoolsConfig::Explicit(vec![MemoryPoolConfig::Sliced {
page_size: MemorySize(MIB),
max_slice_size: None,
max_pool_size: None,
dealloc_period: None,
}]);
let mut memory_management = manage(&growable);
let dry_run = DryRun::new();
let live = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
{
let _measurement = RealRun::new();
let scratch = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
drop(scratch);
}
drop(live);
assert_eq!(
memory_management.memory_report().dynamic[0].pages_peak,
2,
"the scratch is in the marks: it is an allocation like any other"
);
let resolved = MemoryConfiguration::default()
.resolve(Some(&growable), &props())
.unwrap();
memory_management
.install_pools(resolved, &props(), &mut ErrorGraph::default())
.unwrap();
let live = memory_management
.reserve(600 * 1024, &mut ErrorGraph::default())
.unwrap();
drop(live);
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages_peak, 1,
"the plan is the workload's own high-water: {report:?}"
);
drop(dry_run);
}
#[test]
fn direct_pool_pads_only_to_alignment() {
let mut memory_management = manage(&MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Direct { reclaim_at: None },
]));
let odd = 700 * 1024 + 17;
let _live = memory_management
.reserve(odd, &mut ErrorGraph::default())
.unwrap();
let report = memory_management.memory_report();
assert!(matches!(report.dynamic[0].kind, MemoryPoolKind::Direct));
assert_eq!(report.dynamic[0].usage.bytes_in_use, odd);
assert_eq!(
report.dynamic[0].usage.bytes_reserved,
odd.next_multiple_of(props().alignment),
"nothing is reserved beyond what alignment demands: {report:?}"
);
}
#[test]
fn direct_pool_reuses_below_the_ceiling() {
let mut memory_management = manage(&MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Direct {
reclaim_at: Some(MemorySize(8 * MIB)),
},
]));
for _ in 0..4 {
let scratch = memory_management
.reserve(300 * 1024, &mut ErrorGraph::default())
.unwrap();
drop(scratch);
}
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages_peak, 1,
"four iterations of one shape allocated once: {report:?}"
);
}
#[test]
fn direct_pool_reclaims_at_the_ceiling() {
let mut memory_management = manage(&MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Direct {
reclaim_at: Some(MemorySize(7 * MIB)),
},
]));
for size in [MIB, 2 * MIB, 3 * MIB] {
let held = memory_management
.reserve(size, &mut ErrorGraph::default())
.unwrap();
drop(held);
}
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages, 3,
"held below the ceiling: {report:?}"
);
let _crossed = memory_management
.reserve(4 * MIB, &mut ErrorGraph::default())
.unwrap();
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages, 2,
"just enough was released, not everything free: {report:?}"
);
assert_eq!(
report.dynamic[0].usage.bytes_reserved,
7 * MIB,
"the kept 3 MiB slice plus the new 4 MiB one: {report:?}"
);
}
#[test]
fn direct_pool_cleanup_releases_everything_free() {
let mut memory_management = manage(&MemoryPoolsConfig::Explicit(vec![
MemoryPoolConfig::Direct {
reclaim_at: Some(MemorySize(64 * MIB)),
},
]));
let live = memory_management
.reserve(MIB, &mut ErrorGraph::default())
.unwrap();
let freed = memory_management
.reserve(2 * MIB, &mut ErrorGraph::default())
.unwrap();
drop(freed);
memory_management.cleanup(true, &mut ErrorGraph::default());
let report = memory_management.memory_report();
assert_eq!(
report.dynamic[0].pages, 1,
"only the live slice survives: {report:?}"
);
drop(live);
}