use crate::candle::candle_core::{Device, Result as CandleResult, Tensor};
use log::info;
pub fn gpu_mem_info(dev: &Device) -> Option<(usize, usize)> {
#[cfg(feature = "cuda")]
{
let cuda = dev.as_cuda_device().ok()?;
cuda.cuda_stream().context().mem_get_info().ok()
}
#[cfg(not(feature = "cuda"))]
{
let _ = dev;
None
}
}
pub fn auto_chunk_size(
dev: &Device,
cap: usize,
floor: usize,
target_frac: f32,
mut forward_at: impl FnMut(usize) -> CandleResult<Tensor>,
) -> Option<usize> {
let (free0, total) = gpu_mem_info(dev)?;
let cap = cap.max(1);
let floor = floor.max(1).min(cap);
let mut n = floor;
let mut verified: Option<usize> = None;
let mut per_item = 0usize;
loop {
let loss = match forward_at(n) {
Ok(t) => t,
Err(_) => break,
};
let Some((free_alive, _)) = gpu_mem_info(dev) else {
drop(loss);
verified = verified.or(Some(n));
break;
};
drop(loss);
let used = free0.saturating_sub(free_alive);
per_item = (used / n).max(1);
verified = Some(n);
if n >= cap {
break;
}
let next = (n * 2).min(cap);
if n > floor && chunk_from_measurements(free0, per_item, next, n, target_frac) == n {
break;
}
n = next;
}
let chosen = verified?;
info!(
"auto chunk: {} (cap {}, floor {}) from {:.2} GiB free of {:.2} GiB, \
~{} KiB per item, {:.0}% target",
chosen,
cap,
floor,
free0 as f64 / 1024.0 / 1024.0 / 1024.0,
total as f64 / 1024.0 / 1024.0 / 1024.0,
per_item / 1024,
f64::from(target_frac) * 100.0
);
Some(chosen)
}
pub fn chunk_from_measurements(
free: usize,
per_item: usize,
cap: usize,
floor: usize,
target_frac: f32,
) -> usize {
let floor = floor.max(1).min(cap.max(1));
let budget = (free as f64 * f64::from(target_frac.clamp(0.05, 0.95)) * 0.5) as usize;
let per_item = per_item.max(1);
let mut n = floor;
while n < cap && (n * 2).min(cap) * per_item <= budget {
n = (n * 2).min(cap);
}
n
}