use std::ffi::c_void;
use std::os::raw::{c_int, c_uint};
use std::time::Instant;
use kime_tensor::{Error, Result};
use libloading::Library;
type Device = *mut c_void;
type Init = unsafe extern "C" fn() -> c_int;
type Handle = unsafe extern "C" fn(c_uint, *mut Device) -> c_int;
type Energy = unsafe extern "C" fn(Device, *mut u64) -> c_int;
type Power = unsafe extern "C" fn(Device, *mut c_uint) -> c_int;
pub struct Meter {
device: Device,
energy: Energy,
power: Power,
_lib: Library,
}
impl std::fmt::Debug for Meter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Meter").finish_non_exhaustive()
}
}
fn nvml(what: &str, code: c_int) -> Result<()> {
if code == 0 { Ok(()) } else { Err(Error::Device(format!("NVML {what} returned {code}"))) }
}
impl Meter {
pub fn open(ordinal: u32) -> Result<Self> {
let name = if cfg!(windows) { "nvml.dll" } else { "libnvidia-ml.so.1" };
let lib =
unsafe { Library::new(name) }.map_err(|e| Error::Device(format!("{name}: {e}")))?;
let sym = |s: &str| Error::Device(format!("NVML has no {s}"));
let (init, handle, energy, power) = unsafe {
(
*lib.get::<Init>(b"nvmlInit_v2\0").map_err(|_| sym("nvmlInit_v2"))?,
*lib.get::<Handle>(b"nvmlDeviceGetHandleByIndex_v2\0")
.map_err(|_| sym("nvmlDeviceGetHandleByIndex_v2"))?,
*lib.get::<Energy>(b"nvmlDeviceGetTotalEnergyConsumption\0")
.map_err(|_| sym("nvmlDeviceGetTotalEnergyConsumption"))?,
*lib.get::<Power>(b"nvmlDeviceGetPowerUsage\0")
.map_err(|_| sym("nvmlDeviceGetPowerUsage"))?,
)
};
let mut device: Device = std::ptr::null_mut();
unsafe {
nvml("init", init())?;
nvml("device handle", handle(ordinal, &mut device))?;
}
Ok(Self { device, energy, power, _lib: lib })
}
pub fn millijoules(&self) -> Result<u64> {
let mut mj = 0u64;
nvml("total energy", unsafe { (self.energy)(self.device, &mut mj) })?;
Ok(mj)
}
pub fn watts(&self) -> Result<f64> {
let mut mw: c_uint = 0;
nvml("power", unsafe { (self.power)(self.device, &mut mw) })?;
Ok(f64::from(mw) / 1e3)
}
pub fn measure<T>(&self, f: impl FnOnce() -> T) -> Result<(T, f64, f64)> {
let (e, t) = (self.millijoules()?, Instant::now());
let out = f();
let s = t.elapsed().as_secs_f64();
Ok((out, s, (self.millijoules()? - e) as f64 / 1e3))
}
pub fn idle_watts(&self, secs: f64) -> Result<f64> {
let (_, s, j) =
self.measure(|| std::thread::sleep(std::time::Duration::from_secs_f64(secs)))?;
Ok(j / s)
}
}