use std::collections::BTreeSet;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum HostResidency {
#[default]
Pinned,
Locked,
Pageable,
}
impl HostResidency {
pub fn as_str(self) -> &'static str {
match self {
HostResidency::Pinned => "pinned",
HostResidency::Locked => "locked",
HostResidency::Pageable => "pageable",
}
}
pub fn from_label(label: &str) -> Option<Self> {
match label {
"pinned" => Some(HostResidency::Pinned),
"locked" => Some(HostResidency::Locked),
"pageable" => Some(HostResidency::Pageable),
_ => None,
}
}
pub fn is_device_addressable(self) -> bool {
matches!(self, HostResidency::Pinned)
}
}
impl std::fmt::Display for HostResidency {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
pub fn requested_labels(num_layers: usize, cpu_layers: &BTreeSet<u32>) -> Vec<HostResidency> {
(0..num_layers)
.map(|layer| {
if cpu_layers.contains(&(layer as u32)) {
HostResidency::Locked
} else {
HostResidency::Pinned
}
})
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SettleAction {
PageLockForDevice,
LockResident,
LeavePageable,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResidencyPlan {
requested: Vec<HostResidency>,
achieved: Vec<Option<HostResidency>>,
applied: bool,
lock_quota_exhausted: bool,
}
impl ResidencyPlan {
pub fn new(requested: Vec<HostResidency>) -> Self {
let achieved = vec![None; requested.len()];
Self {
requested,
achieved,
applied: false,
lock_quota_exhausted: false,
}
}
pub fn all_pinned(num_layers: usize) -> Self {
Self::new(vec![HostResidency::Pinned; num_layers])
}
pub fn num_layers(&self) -> usize {
self.requested.len()
}
pub fn has_unpinned(&self) -> bool {
self.requested.iter().any(|r| !r.is_device_addressable())
}
pub fn applied(&self) -> bool {
self.applied
}
pub fn lock_quota_exhausted(&self) -> bool {
self.lock_quota_exhausted
}
pub fn requested(&self, layer_id: usize) -> HostResidency {
self.requested[layer_id]
}
pub fn settle_action(&mut self, layer_id: usize) -> SettleAction {
self.applied = true;
match self.requested[layer_id] {
HostResidency::Pinned => SettleAction::PageLockForDevice,
HostResidency::Locked if self.lock_quota_exhausted => SettleAction::LeavePageable,
HostResidency::Locked => SettleAction::LockResident,
HostResidency::Pageable => SettleAction::LeavePageable,
}
}
pub fn record(&mut self, layer_id: usize, achieved: HostResidency) {
if self.achieved[layer_id] != Some(HostResidency::Pageable) {
self.achieved[layer_id] = Some(achieved);
}
}
pub fn record_lock(&mut self, layer_id: usize, locked: bool) {
if !locked {
self.lock_quota_exhausted = true;
}
self.record(
layer_id,
if locked {
HostResidency::Locked
} else {
HostResidency::Pageable
},
);
}
pub fn achieved(&self, layer_id: usize) -> Option<HostResidency> {
self.achieved[layer_id]
}
pub fn achieved_labels(&self) -> Vec<HostResidency> {
self.requested
.iter()
.zip(&self.achieved)
.map(|(requested, achieved)| achieved.unwrap_or(*requested))
.collect()
}
pub fn downgraded(&self) -> Vec<u32> {
self.achieved_labels()
.iter()
.zip(&self.requested)
.enumerate()
.filter(|(_, (achieved, requested))| achieved != requested)
.map(|(layer, _)| layer as u32)
.collect()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ResidencyError {
LabelCountMismatch { labels: usize, num_layers: usize },
UnpinnedLayerNotOnCpu { layers: Vec<u32> },
PrefillOverlapWithUnpinned { layers: Vec<u32> },
SlotRemapOnUnpinnedLayer { layer: u32 },
}
impl std::fmt::Display for ResidencyError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ResidencyError::LabelCountMismatch { labels, num_layers } => write!(
f,
"{labels} residency labels for a model of {num_layers} MoE layers"
),
ResidencyError::UnpinnedLayerNotOnCpu { layers } => write!(
f,
"layers {layers:?} are not page-locked for the device and are not CPU layers: \
a layer without a device address can only decode on the CPU executor, so the \
CPU-layer set must be decided before the banks are attached"
),
ResidencyError::PrefillOverlapWithUnpinned { layers } => write!(
f,
"prefill overlap DMAs from registered banks; it must be disabled when any layer \
is locked or pageable (layers {layers:?})"
),
ResidencyError::SlotRemapOnUnpinnedLayer { layer } => write!(
f,
"layer {layer} is not page-locked for the device: its only copy is the \
whole-layer materialize (position == expert id); an LRU slot remap cannot be \
honored without a device alias for the host rows"
),
}
}
}
impl std::error::Error for ResidencyError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CopyRoute {
DeviceIndexed,
WholeLayerPageable,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BankResidency {
labels: Vec<HostResidency>,
unpinned: BTreeSet<u32>,
}
impl BankResidency {
pub fn new(
labels: &[HostResidency],
num_layers: usize,
cpu_layers: &BTreeSet<u32>,
prefill_overlap: bool,
) -> Result<Self, ResidencyError> {
if labels.len() != num_layers {
return Err(ResidencyError::LabelCountMismatch {
labels: labels.len(),
num_layers,
});
}
let unpinned: BTreeSet<u32> = labels
.iter()
.enumerate()
.filter(|(_, label)| !label.is_device_addressable())
.map(|(layer, _)| layer as u32)
.collect();
if !unpinned.is_empty() {
let stranded: Vec<u32> = unpinned.difference(cpu_layers).copied().collect();
if !stranded.is_empty() {
return Err(ResidencyError::UnpinnedLayerNotOnCpu { layers: stranded });
}
if prefill_overlap {
return Err(ResidencyError::PrefillOverlapWithUnpinned {
layers: unpinned.iter().copied().collect(),
});
}
}
Ok(Self {
labels: labels.to_vec(),
unpinned,
})
}
pub fn all_pinned(num_layers: usize) -> Self {
Self {
labels: vec![HostResidency::Pinned; num_layers],
unpinned: BTreeSet::new(),
}
}
pub fn num_layers(&self) -> usize {
self.labels.len()
}
pub fn labels(&self) -> &[HostResidency] {
&self.labels
}
pub fn label(&self, layer_id: u32) -> HostResidency {
self.labels[layer_id as usize]
}
pub fn unpinned_layers(&self) -> &BTreeSet<u32> {
&self.unpinned
}
pub fn is_unpinned(&self, layer_id: u32) -> bool {
self.unpinned.contains(&layer_id)
}
pub fn has_unpinned(&self) -> bool {
!self.unpinned.is_empty()
}
pub fn copy_route(
&self,
layer_id: u32,
whole_layer: bool,
) -> Result<CopyRoute, ResidencyError> {
if !self.is_unpinned(layer_id) {
return Ok(CopyRoute::DeviceIndexed);
}
if whole_layer {
Ok(CopyRoute::WholeLayerPageable)
} else {
Err(ResidencyError::SlotRemapOnUnpinnedLayer { layer: layer_id })
}
}
}
pub const PIN_BUDGET_ENV: &str = "FERROX_PIN_BUDGET_GB";
pub const WSL_KERNEL_TAG: &str = "microsoft";
pub const WSL_PIN_FRACTION: f64 = 0.4;
const GIB: f64 = (1u64 << 30) as f64;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PinBudgetEnvError {
pub value: String,
}
impl std::fmt::Display for PinBudgetEnvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{PIN_BUDGET_ENV}={:?} is not a number of GiB",
self.value
)
}
}
impl std::error::Error for PinBudgetEnvError {}
pub fn parse_pin_budget_gb(value: &str) -> Result<Option<u64>, PinBudgetEnvError> {
let text = value.trim();
if text.is_empty() {
return Ok(None);
}
let gb: f64 = text.parse().map_err(|_| PinBudgetEnvError {
value: value.to_string(),
})?;
if !gb.is_finite() {
return Err(PinBudgetEnvError {
value: value.to_string(),
});
}
Ok(Some((gb * GIB).max(0.0) as u64))
}
pub fn is_pin_capped_host(kernel_release: &str) -> bool {
kernel_release.to_ascii_lowercase().contains(WSL_KERNEL_TAG)
}
pub fn pin_budget_bytes(kernel_release: &str, phys_ram_bytes: u64) -> Option<u64> {
if !is_pin_capped_host(kernel_release) {
return None;
}
Some((phys_ram_bytes as f64 * WSL_PIN_FRACTION) as u64)
}
pub fn resolve_pin_budget(
kernel_release: &str,
phys_ram_bytes: u64,
env_value: Option<&str>,
) -> Result<Option<u64>, PinBudgetEnvError> {
if let Some(value) = env_value {
if let Some(bytes) = parse_pin_budget_gb(value)? {
return Ok(Some(bytes));
}
}
Ok(pin_budget_bytes(kernel_release, phys_ram_bytes))
}
pub fn host_kernel_release() -> Option<String> {
std::fs::read_to_string("/proc/sys/kernel/osrelease")
.ok()
.map(|release| release.trim().to_string())
.filter(|release| !release.is_empty())
}
pub fn host_phys_ram_bytes() -> Option<u64> {
parse_mem_total_bytes(&std::fs::read_to_string("/proc/meminfo").ok()?)
}
fn parse_mem_total_bytes(meminfo: &str) -> Option<u64> {
let line = meminfo.lines().find(|line| line.starts_with("MemTotal:"))?;
let mut fields = line.split_whitespace().skip(1);
let value: u64 = fields.next()?.parse().ok()?;
let scale = match fields.next() {
None => 1,
Some("kB") | Some("KB") | Some("kb") => 1024,
Some(_) => return None,
};
value.checked_mul(scale)
}
pub fn host_pin_budget_bytes() -> Result<Option<u64>, PinBudgetEnvError> {
let env_value = std::env::var(PIN_BUDGET_ENV).ok();
resolve_pin_budget(
&host_kernel_release().unwrap_or_default(),
host_phys_ram_bytes().unwrap_or(0),
env_value.as_deref(),
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::placement::auto_cpu_layers;
fn set(items: &[u32]) -> BTreeSet<u32> {
items.iter().copied().collect()
}
#[test]
fn a_label_survives_a_round_trip_through_its_wire_form() {
for label in [
HostResidency::Pinned,
HostResidency::Locked,
HostResidency::Pageable,
] {
assert_eq!(HostResidency::from_label(label.as_str()), Some(label));
}
assert_eq!(HostResidency::from_label("registered"), None);
assert!(HostResidency::Pinned.is_device_addressable());
assert!(!HostResidency::Locked.is_device_addressable());
assert!(!HostResidency::Pageable.is_device_addressable());
}
#[test]
fn the_cpu_layers_are_the_ones_asked_to_lock() {
assert_eq!(
requested_labels(4, &set(&[0, 3])),
vec![
HostResidency::Locked,
HostResidency::Pinned,
HostResidency::Pinned,
HostResidency::Locked,
]
);
assert_eq!(
requested_labels(3, &BTreeSet::new()),
vec![HostResidency::Pinned; 3]
);
}
#[test]
fn a_plan_settles_each_layer_at_the_class_it_asked_for() {
let mut plan = ResidencyPlan::new(requested_labels(3, &set(&[1])));
assert!(!plan.applied());
assert_eq!(plan.settle_action(0), SettleAction::PageLockForDevice);
assert_eq!(plan.settle_action(1), SettleAction::LockResident);
assert!(plan.applied());
assert!(plan.has_unpinned());
assert!(!ResidencyPlan::all_pinned(3).has_unpinned());
}
#[test]
fn a_failed_lock_is_recorded_as_pageable_rather_than_assumed_locked() {
let mut plan = ResidencyPlan::new(requested_labels(3, &set(&[0, 2])));
assert_eq!(plan.settle_action(0), SettleAction::LockResident);
plan.record_lock(0, false);
assert_eq!(plan.settle_action(2), SettleAction::LeavePageable);
plan.record(2, HostResidency::Pageable);
assert_eq!(plan.achieved(0), Some(HostResidency::Pageable));
assert_eq!(
plan.achieved_labels(),
vec![
HostResidency::Pageable,
HostResidency::Pinned,
HostResidency::Pageable,
],
"the quota is spent for good, so layer 2's lock never even ran"
);
assert_eq!(plan.downgraded(), vec![0, 2]);
}
#[test]
fn an_unreported_layer_echoes_back_what_it_asked_for() {
let mut plan = ResidencyPlan::new(requested_labels(2, &set(&[1])));
plan.settle_action(1);
plan.record_lock(1, true);
assert_eq!(
plan.achieved_labels(),
vec![HostResidency::Pinned, HostResidency::Locked]
);
assert!(plan.downgraded().is_empty());
}
#[test]
fn one_pageable_bank_downgrades_the_whole_layer() {
let mut plan = ResidencyPlan::all_pinned(1);
plan.record(0, HostResidency::Pageable);
plan.record(0, HostResidency::Locked);
plan.record(0, HostResidency::Pinned);
assert_eq!(plan.achieved(0), Some(HostResidency::Pageable));
}
#[test]
fn a_spent_lock_quota_leaves_every_later_layer_pageable() {
let mut plan = ResidencyPlan::new(vec![HostResidency::Locked; 3]);
assert_eq!(plan.settle_action(0), SettleAction::LockResident);
plan.record_lock(0, false);
assert!(plan.lock_quota_exhausted());
assert_eq!(plan.settle_action(1), SettleAction::LeavePageable);
assert_eq!(plan.settle_action(2), SettleAction::LeavePageable);
}
#[test]
fn labels_that_do_not_describe_the_model_are_refused() {
assert_eq!(
BankResidency::new(&[HostResidency::Pinned; 3], 4, &BTreeSet::new(), false),
Err(ResidencyError::LabelCountMismatch {
labels: 3,
num_layers: 4
})
);
}
#[test]
fn a_non_pinned_layer_that_is_not_a_cpu_layer_is_refused() {
let labels = requested_labels(4, &set(&[0, 3]));
assert_eq!(
BankResidency::new(&labels, 4, &set(&[0]), false),
Err(ResidencyError::UnpinnedLayerNotOnCpu { layers: vec![3] })
);
let attached = BankResidency::new(&labels, 4, &set(&[0, 3]), false).unwrap();
assert_eq!(attached.unpinned_layers(), &set(&[0, 3]));
assert!(
BankResidency::new(&[HostResidency::Pinned; 4], 4, &set(&[0, 3]), true).is_ok(),
"an all-pinned load keeps prefill overlap even with CPU layers"
);
}
#[test]
fn prefill_overlap_with_any_unpinned_layer_is_refused() {
let labels = requested_labels(4, &set(&[2]));
assert_eq!(
BankResidency::new(&labels, 4, &set(&[2]), true),
Err(ResidencyError::PrefillOverlapWithUnpinned { layers: vec![2] })
);
assert!(BankResidency::new(&labels, 4, &set(&[2]), false).is_ok());
}
#[test]
fn an_unpinned_layer_accepts_only_the_whole_layer_materialize() {
let labels = requested_labels(3, &set(&[1]));
let banks = BankResidency::new(&labels, 3, &set(&[1]), false).unwrap();
assert_eq!(banks.copy_route(1, true), Ok(CopyRoute::WholeLayerPageable));
assert_eq!(
banks.copy_route(1, false),
Err(ResidencyError::SlotRemapOnUnpinnedLayer { layer: 1 })
);
}
#[test]
fn a_pinned_layer_takes_the_indexed_device_copy_either_way() {
let banks = BankResidency::all_pinned(3);
assert!(!banks.has_unpinned());
assert_eq!(banks.copy_route(0, true), Ok(CopyRoute::DeviceIndexed));
assert_eq!(banks.copy_route(0, false), Ok(CopyRoute::DeviceIndexed));
assert_eq!(banks.label(0), HostResidency::Pinned);
}
#[test]
fn the_achieved_labels_are_what_the_banks_attach_with() {
let cpu_layers = set(&[0, 5]);
let mut plan = ResidencyPlan::new(requested_labels(6, &cpu_layers));
assert_eq!(plan.settle_action(0), SettleAction::LockResident);
plan.record_lock(0, false);
assert_eq!(plan.settle_action(5), SettleAction::LeavePageable);
plan.record(5, HostResidency::Pageable);
let labels = plan.achieved_labels();
assert_eq!(labels[0], HostResidency::Pageable);
assert_eq!(labels[5], HostResidency::Pageable);
let banks = BankResidency::new(&labels, 6, &cpu_layers, false).unwrap();
assert_eq!(banks.unpinned_layers(), &cpu_layers);
assert_eq!(banks.label(0), HostResidency::Pageable);
}
#[test]
fn a_wsl_host_is_capped_at_forty_percent_of_ram() {
let release = "5.15.153.1-microsoft-standard-WSL2";
assert!(is_pin_capped_host(release));
assert_eq!(
pin_budget_bytes(release, 64 << 30),
Some((64 << 30) * 2 / 5)
);
assert!(is_pin_capped_host("5.10.16.3-Microsoft-standard-WSL2"));
}
#[test]
fn the_wsl_budget_is_what_moves_layers_off_the_gpu_path() {
let banks = 48u64 << 30;
let budget = pin_budget_bytes("5.15.153.1-microsoft-standard-WSL2", 64 << 30);
assert!(
auto_cpu_layers(48, banks, None).is_empty(),
"an uncapped host keeps everything pinned"
);
assert!(!auto_cpu_layers(48, banks, budget).is_empty());
}
#[test]
fn plain_linux_reports_no_pin_cap() {
assert!(!is_pin_capped_host("6.8.0-45-generic"));
assert_eq!(pin_budget_bytes("6.8.0-45-generic", 64 << 30), None);
assert_eq!(pin_budget_bytes("", 64 << 30), None, "no readable release");
assert_eq!(
resolve_pin_budget("6.8.0-45-generic", 64 << 30, None),
Ok(None)
);
}
#[test]
fn an_unreadable_ram_figure_on_a_capped_host_budgets_nothing() {
assert_eq!(pin_budget_bytes("microsoft-standard-WSL2", 0), Some(0));
assert_eq!(auto_cpu_layers(8, 1 << 30, Some(0)).len(), 8);
}
#[test]
fn the_environment_variable_overrides_on_any_host() {
assert_eq!(
resolve_pin_budget("6.8.0-45-generic", 64 << 30, Some("8")),
Ok(Some(8 << 30)),
"an uncapped host can still be told a budget"
);
assert_eq!(
resolve_pin_budget("microsoft-standard-WSL2", 64 << 30, Some("1.5")),
Ok(Some(1024 * 1024 * 1024 * 3 / 2)),
"and a capped host's computed budget is replaced, not clamped"
);
assert_eq!(
resolve_pin_budget("6.8.0-45-generic", 64 << 30, Some("-1")),
Ok(Some(0)),
"a negative budget means pin nothing"
);
}
#[test]
fn an_empty_pin_budget_variable_counts_as_unset() {
assert_eq!(parse_pin_budget_gb(""), Ok(None));
assert_eq!(parse_pin_budget_gb(" "), Ok(None));
assert_eq!(
resolve_pin_budget("microsoft-standard-WSL2", 64 << 30, Some("")),
Ok(Some((64 << 30) * 2 / 5)),
"an empty value in a unit file means 'use the normal rule'"
);
}
#[test]
fn an_unparsable_pin_budget_is_refused_rather_than_ignored() {
for value in ["eight", "8GiB", "inf", "NaN"] {
assert_eq!(
parse_pin_budget_gb(value),
Err(PinBudgetEnvError {
value: value.to_string()
}),
"{value:?} must not read as 'no cap'"
);
}
assert!(resolve_pin_budget("6.8.0-45-generic", 64 << 30, Some("eight")).is_err());
}
#[test]
fn mem_total_is_read_in_kilobytes() {
let meminfo = "MemTotal: 65809172 kB\nMemFree: 1234 kB\n";
assert_eq!(parse_mem_total_bytes(meminfo), Some(65809172 * 1024));
assert_eq!(parse_mem_total_bytes("MemFree: 1234 kB\n"), None);
assert_eq!(parse_mem_total_bytes("MemTotal: notanumber kB"), None);
assert_eq!(parse_mem_total_bytes("MemTotal: 12 MB"), None);
}
}