#![cfg(not(exclusive_memory_only))]
use std::sync::Arc;
use cubecl_ir::MemoryDeviceProperties;
use cubecl_runtime::config::memory::{MemoryPoolConfig, MemoryPoolsConfig};
use cubecl_runtime::config::size::MemorySize;
use cubecl_runtime::logging::ServerLogger;
use cubecl_runtime::memory_management::{
MemoryConfiguration, MemoryManagement, MemoryManagementOptions,
};
use cubecl_runtime::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 {
max_page_size: 128 * MIB,
alignment: 32,
}
}
#[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).unwrap();
drop(small);
let _large = memory_management.reserve(512 * 1024).unwrap();
assert_eq!(memory_management.memory_usage().bytes_reserved, MIB);
let _fill_1 = memory_management.reserve(MIB).unwrap();
let _fill_2 = memory_management.reserve(500 * 1024).unwrap();
assert!(memory_management.reserve(MIB).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).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).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).unwrap();
let bigger = MemoryConfiguration::default()
.resolve(Some(&sliced(4 * MIB, 2)), &props())
.unwrap();
assert!(!memory_management.configure(bigger.clone(), &props()));
assert!(memory_management.reserve(2 * MIB).is_err());
drop(live);
assert!(memory_management.configure(bigger, &props()));
let _large = memory_management.reserve(2 * MIB).unwrap();
}