pub(crate) mod nvml;
pub mod topology;
pub mod worker_pool;
use cudarc::driver::{result::device as cuda_device, sys as cuda_sys};
use nix::libc;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::{Mutex, OnceLock};
use std::{fs, mem, process::Command};
static NUMA_NODE_CACHE: OnceLock<Mutex<HashMap<String, Option<NumaNode>>>> = OnceLock::new();
pub fn is_numa_enabled() -> bool {
!crate::env_is_truthy("DYN_MEMORY_DISABLE_NUMA")
}
pub fn is_numa_disabled() -> bool {
!is_numa_enabled()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct NumaNode(pub u32);
impl NumaNode {
pub const UNKNOWN: NumaNode = NumaNode(u32::MAX);
pub fn is_unknown(&self) -> bool {
self.0 == u32::MAX
}
}
impl std::fmt::Display for NumaNode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
if self.is_unknown() {
write!(f, "UNKNOWN")
} else {
write!(f, "NumaNode({})", self.0)
}
}
}
pub fn get_current_cpu_numa_node() -> NumaNode {
unsafe {
let mut cpu: libc::c_uint = 0;
let mut node: libc::c_uint = 0;
let result = libc::syscall(
libc::SYS_getcpu,
&mut cpu,
&mut node,
std::ptr::null_mut::<libc::c_void>(),
);
if result == 0 {
NumaNode(node)
} else {
NumaNode::UNKNOWN
}
}
}
fn read_numa_node_from_sysfs(pci_address: &str) -> Option<NumaNode> {
let path = format!("/sys/bus/pci/devices/{}/numa_node", pci_address);
let content = fs::read_to_string(&path).ok()?;
let node: i32 = content.trim().parse().ok()?;
if node < 0 {
None
} else {
Some(NumaNode(node as u32))
}
}
fn get_numa_node_from_nvidia_smi(pci_address: &str) -> Option<NumaNode> {
let output = Command::new("nvidia-smi")
.args(["topo", "--get-numa-id-of-nearby-cpu", "-i", pci_address])
.output()
.ok()?;
if !output.status.success() {
return None;
}
let stdout = std::str::from_utf8(&output.stdout).ok()?;
let line = stdout.lines().next()?;
let numa_str = line.split(':').nth(1)?;
let node: u32 = numa_str.trim().parse().ok()?;
Some(NumaNode(node))
}
pub fn get_device_numa_node(device_id: u32) -> Option<NumaNode> {
let pci_address = match get_pci_bus_address_from_cuda(device_id) {
Some(addr) => addr,
None => {
tracing::warn!(
"Failed to get PCI address from CUDA for device {}, skipping NUMA optimization",
device_id
);
return None;
}
};
let cache = NUMA_NODE_CACHE.get_or_init(|| Mutex::new(HashMap::new()));
{
let guard = cache.lock().unwrap();
if let Some(cached) = guard.get(&pci_address) {
return *cached;
}
}
let result = read_numa_node_from_sysfs(&pci_address)
.or_else(|| get_numa_node_from_nvidia_smi(&pci_address));
match result {
Some(node) => {
tracing::trace!(
"GPU {} (PCI {}) on NUMA node {}",
device_id,
pci_address,
node.0
);
}
None => {
tracing::warn!(
"Could not determine NUMA node for GPU {} (PCI {}), skipping NUMA optimization",
device_id,
pci_address
);
}
}
cache.lock().unwrap().insert(pci_address, result);
result
}
pub fn pin_thread_to_numa_node(node: NumaNode) -> Result<(), String> {
let topology =
topology::get_numa_topology().map_err(|e| format!("Can not get NUMA topology: {}", e))?;
let cpus = topology
.cpus_for_node(node.0)
.ok_or_else(|| format!("No CPUs found for NUMA node {}", node.0))?;
if cpus.is_empty() {
return Err(format!("No CPUs found for NUMA node {}", node.0));
}
unsafe {
let mut cpu_set: libc::cpu_set_t = mem::zeroed();
for cpu in cpus {
libc::CPU_SET(*cpu, &mut cpu_set);
}
let result = libc::sched_setaffinity(
0, mem::size_of::<libc::cpu_set_t>(),
&cpu_set,
);
if result != 0 {
let err = std::io::Error::last_os_error();
return Err(format!("Failed to set CPU affinity: {}", err));
}
}
Ok(())
}
fn get_pci_bus_address_from_cuda(device_id: u32) -> Option<String> {
unsafe {
let mut dev = std::mem::MaybeUninit::uninit();
if cuda_sys::cuDeviceGet(dev.as_mut_ptr(), device_id as i32)
.result()
.is_err()
{
return None;
}
let dev = dev.assume_init();
let domain = cuda_device::get_attribute(
dev,
cuda_sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_PCI_DOMAIN_ID,
)
.ok()?;
let bus = cuda_device::get_attribute(
dev,
cuda_sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_PCI_BUS_ID,
)
.ok()?;
let device = cuda_device::get_attribute(
dev,
cuda_sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_PCI_DEVICE_ID,
)
.ok()?;
Some(format!("{:04x}:{:02x}:{:02x}.0", domain, bus, device))
}
}
#[derive(Debug, Clone)]
struct GpuTopoInfo {
pci_address: String,
numa_node: Option<u32>,
}
fn enumerate_cuda_gpus() -> Vec<GpuTopoInfo> {
let count = match cuda_device::get_count() {
Ok(c) => c,
Err(_) => return Vec::new(),
};
(0..count as u32)
.filter_map(|i| {
let pci = get_pci_bus_address_from_cuda(i)?;
let numa = read_numa_node_from_sysfs(&pci).map(|n| n.0);
Some(GpuTopoInfo {
pci_address: pci,
numa_node: numa,
})
})
.collect()
}
fn enumerate_all_gpus() -> Vec<GpuTopoInfo> {
if let Some(nvml) = nvml::try_nvml() {
let nvml_gpus = nvml.enumerate_gpus();
if !nvml_gpus.is_empty() {
tracing::debug!(
"NVML enumerated {} GPUs (ignoring CUDA_VISIBLE_DEVICES)",
nvml_gpus.len()
);
return nvml_gpus
.into_iter()
.map(|g| {
let numa = read_numa_node_from_sysfs(&g.pci_address).map(|n| n.0);
GpuTopoInfo {
pci_address: g.pci_address,
numa_node: numa,
}
})
.collect();
}
}
tracing::debug!("Falling back to CUDA driver GPU enumeration");
enumerate_cuda_gpus()
}
static DEVICE_CPU_SETS: OnceLock<HashMap<u32, Option<Vec<usize>>>> = OnceLock::new();
pub fn get_device_cpu_set(device_id: u32) -> Option<Vec<usize>> {
DEVICE_CPU_SETS
.get_or_init(compute_all_device_cpu_sets)
.get(&device_id)
.cloned()
.flatten()
}
fn compute_all_device_cpu_sets() -> HashMap<u32, Option<Vec<usize>>> {
let topology = match topology::get_numa_topology() {
Ok(t) => t,
Err(e) => {
tracing::warn!("Cannot subdivide CPU sets: {e}");
return HashMap::new();
}
};
let cuda_count = cuda_device::get_count().unwrap_or(0);
if cuda_count == 0 {
return HashMap::new();
}
let mut cuda_devices: Vec<(u32, String, Option<u32>)> = Vec::new();
for i in 0..cuda_count as u32 {
if let Some(pci) = get_pci_bus_address_from_cuda(i) {
let numa = read_numa_node_from_sysfs(&pci).map(|n| n.0);
cuda_devices.push((i, pci, numa));
}
}
let all_gpus = enumerate_all_gpus();
let mut node_groups: HashMap<u32, Vec<String>> = HashMap::new();
for gpu in &all_gpus {
if let Some(node) = gpu.numa_node {
node_groups
.entry(node)
.or_default()
.push(gpu.pci_address.clone());
}
}
for group in node_groups.values_mut() {
group.sort();
}
let mut results = HashMap::new();
for (device_id, pci_addr, numa_node) in &cuda_devices {
let cpu_set = numa_node.and_then(|node| {
let group = node_groups.get(&node)?;
let position = group.iter().position(|addr| addr == pci_addr)?;
let all_cpus = topology.cpus_for_node(node)?;
if all_cpus.is_empty() || group.is_empty() {
return None;
}
let n = group.len();
let chunk_size = all_cpus.len() / n;
if chunk_size == 0 {
return Some(all_cpus.to_vec());
}
let start = position * chunk_size;
let end = if position == n - 1 {
all_cpus.len() } else {
start + chunk_size
};
Some(all_cpus[start..end].to_vec())
});
results.insert(*device_id, cpu_set);
}
results
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_numa_node_equality() {
let node0a = NumaNode(0);
let node0b = NumaNode(0);
let node1 = NumaNode(1);
assert_eq!(node0a, node0b);
assert_ne!(node0a, node1);
}
#[test]
fn test_numa_node_unknown() {
let unknown = NumaNode::UNKNOWN;
assert!(unknown.is_unknown());
assert_eq!(unknown.0, u32::MAX);
let valid = NumaNode(0);
assert!(!valid.is_unknown());
}
#[test]
fn test_numa_node_display() {
assert_eq!(format!("{}", NumaNode(0)), "NumaNode(0)");
assert_eq!(format!("{}", NumaNode(7)), "NumaNode(7)");
assert_eq!(format!("{}", NumaNode::UNKNOWN), "UNKNOWN");
}
#[test]
fn test_numa_node_serialization() {
let node = NumaNode(1);
let json = serde_json::to_string(&node).unwrap();
let deserialized: NumaNode = serde_json::from_str(&json).unwrap();
assert_eq!(node, deserialized);
}
#[test]
fn test_get_current_cpu_numa_node() {
let node = get_current_cpu_numa_node();
if !node.is_unknown() {
assert!(node.0 < 8, "NUMA node {} seems unreasonably high", node.0);
}
}
#[test]
fn test_numa_node_hash() {
use std::collections::HashMap;
let mut map = HashMap::new();
map.insert(NumaNode(0), "node0");
map.insert(NumaNode(1), "node1");
assert_eq!(map.get(&NumaNode(0)), Some(&"node0"));
assert_eq!(map.get(&NumaNode(1)), Some(&"node1"));
assert_eq!(map.get(&NumaNode(2)), None);
}
#[test]
fn test_numa_node_copy_clone() {
let node1 = NumaNode(5);
let node2 = node1;
let node3 = node1;
assert_eq!(node1, node2);
assert_eq!(node1, node3);
assert_eq!(node2, node3);
}
#[test]
fn test_read_numa_node_from_sysfs_nonexistent() {
assert!(read_numa_node_from_sysfs("ffff:ff:ff.0").is_none());
}
}
#[cfg(all(test, feature = "testing-cuda"))]
mod cuda_tests {
use super::*;
#[test]
fn test_get_pci_bus_address_from_cuda() {
let addr = get_pci_bus_address_from_cuda(0).expect("should get PCI address for GPU 0");
let parts: Vec<&str> = addr.split(':').collect();
assert_eq!(
parts.len(),
3,
"PCI address should have 3 colon-separated parts: {}",
addr
);
assert_eq!(parts[0].len(), 4, "domain should be 4 hex chars: {}", addr);
assert!(parts[2].ends_with(".0"), "should end with .0: {}", addr);
println!("GPU 0 PCI address: {}", addr);
}
#[test]
fn test_read_numa_node_from_sysfs_real_gpu() {
let addr = get_pci_bus_address_from_cuda(0).expect("should get PCI address for GPU 0");
if let Some(node) = read_numa_node_from_sysfs(&addr) {
assert!(node.0 < 16, "NUMA node {} seems unreasonably high", node.0);
println!("GPU 0 (PCI {}) sysfs NUMA node: {}", addr, node.0);
} else {
println!(
"GPU 0 (PCI {}) has no sysfs NUMA info (single-socket?)",
addr
);
}
}
#[test]
fn test_get_device_numa_node_returns_some_or_none() {
let result = get_device_numa_node(0);
match result {
Some(node) => {
assert!(node.0 < 16, "NUMA node {} seems unreasonably high", node.0);
assert!(
!node.is_unknown(),
"should never return UNKNOWN inside Some"
);
println!("GPU 0 detected on NUMA node: {}", node.0);
}
None => {
println!("GPU 0 has no determinable NUMA node (single-socket or no sysfs info)");
}
}
}
}