use super::output_budget::{OutputBudgetInputs, OutputCapMode, resolve_output_budget};
use crate::constants::{
DEFAULT_OLLAMA_MAX_AUTO_NUM_CTX, OLLAMA_KV_DTYPE_BYTES, OLLAMA_MEMORY_BUDGET_FRACTION,
OLLAMA_MIN_AUTO_NUM_CTX, OLLAMA_MIN_NUM_PREDICT, OLLAMA_NUM_CTX_ROUNDING,
OLLAMA_NUM_PREDICT_MARGIN,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum NumCtxSource {
Override,
GlobalConfig,
Auto,
AutoFallback,
Cloud,
}
impl NumCtxSource {
pub fn label(self) -> &'static str {
match self {
NumCtxSource::Override => "override",
NumCtxSource::GlobalConfig => "config",
NumCtxSource::Auto => "auto",
NumCtxSource::AutoFallback => "auto (fallback)",
NumCtxSource::Cloud => "cloud (full window)",
}
}
pub fn is_auto(self) -> bool {
matches!(self, NumCtxSource::Auto | NumCtxSource::AutoFallback)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct NumCtxResolution {
pub value: usize,
pub source: NumCtxSource,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct ModelDims {
pub block_count: usize,
pub head_count: usize,
pub head_count_kv: usize,
pub embedding_length: usize,
}
pub fn kv_bytes_per_token(dims: &ModelDims) -> Option<usize> {
if dims.block_count == 0
|| dims.head_count == 0
|| dims.head_count_kv == 0
|| dims.embedding_length == 0
{
return None;
}
let head_dim = dims.embedding_length / dims.head_count;
if head_dim == 0 {
return None;
}
2usize
.checked_mul(dims.block_count)?
.checked_mul(dims.head_count_kv)?
.checked_mul(head_dim)?
.checked_mul(OLLAMA_KV_DTYPE_BYTES)
}
pub fn max_tokens_for_memory(
budget_bytes: u64,
model_weight_bytes: u64,
kv_bytes_per_token: usize,
) -> Option<usize> {
if kv_bytes_per_token == 0 {
return None;
}
let usable = (budget_bytes as f64 * OLLAMA_MEMORY_BUDGET_FRACTION) as u64;
let for_kv = usable.saturating_sub(model_weight_bytes);
Some((for_kv / kv_bytes_per_token as u64) as usize)
}
pub fn converge_num_ctx(
current: usize,
size_vram_bytes: u64,
total_bytes: u64,
kv_bytes_per_token: usize,
) -> Option<usize> {
if kv_bytes_per_token == 0 || total_bytes <= size_vram_bytes {
return None; }
let overflow = total_bytes - size_vram_bytes;
let tokens_to_cut = overflow.div_ceil(kv_bytes_per_token as u64) as usize;
let after_cut = current.saturating_sub(tokens_to_cut);
if after_cut < OLLAMA_MIN_AUTO_NUM_CTX {
return None;
}
let target = round_down_to(after_cut, OLLAMA_NUM_CTX_ROUNDING).max(OLLAMA_MIN_AUTO_NUM_CTX);
(target < current).then_some(target)
}
#[derive(Debug, Clone, Copy, Default)]
pub struct NumCtxInputs {
pub model_max: Option<usize>,
pub dims: Option<ModelDims>,
pub model_weight_bytes: Option<u64>,
pub per_model_override: Option<u32>,
pub global_num_ctx: Option<i32>,
pub allow_ram_offload: bool,
pub vram_bytes: Option<u64>,
pub system_ram_bytes: Option<u64>,
pub max_auto_cap: Option<usize>,
pub is_cloud: bool,
}
pub fn resolve_ollama_num_ctx(inputs: &NumCtxInputs) -> Option<NumCtxResolution> {
if let Some(n) = inputs.per_model_override.filter(|n| *n > 0) {
return Some(NumCtxResolution {
value: n as usize,
source: NumCtxSource::Override,
});
}
if inputs.is_cloud {
return inputs
.model_max
.filter(|m| *m > 0)
.map(|model_max| NumCtxResolution {
value: model_max,
source: NumCtxSource::Cloud,
});
}
if let Some(n) = inputs.global_num_ctx.filter(|n| *n > 0) {
return Some(NumCtxResolution {
value: n as usize,
source: NumCtxSource::GlobalConfig,
});
}
let model_max = inputs.model_max.filter(|m| *m > 0)?;
let budget_bytes = if inputs.allow_ram_offload {
inputs.system_ram_bytes
} else {
inputs.vram_bytes
};
let kv = inputs.dims.as_ref().and_then(kv_bytes_per_token);
let (raw_target, source) = match (budget_bytes, kv) {
(Some(b), Some(kv)) => {
let fit = max_tokens_for_memory(b, inputs.model_weight_bytes.unwrap_or(0), kv)
.unwrap_or(DEFAULT_OLLAMA_MAX_AUTO_NUM_CTX);
(fit, NumCtxSource::Auto)
},
_ => (DEFAULT_OLLAMA_MAX_AUTO_NUM_CTX, NumCtxSource::AutoFallback),
};
let target = round_down_to(raw_target, OLLAMA_NUM_CTX_ROUNDING);
let mut value = target.min(model_max);
if let Some(cap) = inputs.max_auto_cap.filter(|c| *c > 0) {
value = value.min(cap);
}
value = value.max(OLLAMA_MIN_AUTO_NUM_CTX.min(model_max));
Some(NumCtxResolution { value, source })
}
fn round_down_to(v: usize, step: usize) -> usize {
if step == 0 {
return v;
}
(v / step) * step
}
pub fn default_ollama_num_predict(
max_tokens: usize,
num_ctx: Option<usize>,
prompt_estimate: usize,
provider_max_output: Option<usize>,
) -> Option<i32> {
resolve_output_budget(
&OutputBudgetInputs {
requested_cap: max_tokens,
window: num_ctx,
prompt_estimate,
provider_max_output,
margin: OLLAMA_NUM_PREDICT_MARGIN,
floor: OLLAMA_MIN_NUM_PREDICT,
},
OutputCapMode::NumPredict,
)
.map(|v| v.min(i32::MAX as usize) as i32)
}
#[cfg(test)]
mod tests {
use super::*;
fn gqa_dims() -> ModelDims {
ModelDims {
block_count: 28,
head_count: 28,
head_count_kv: 4,
embedding_length: 3584,
}
}
#[test]
fn kv_bytes_per_token_gqa() {
assert_eq!(kv_bytes_per_token(&gqa_dims()), Some(57_344));
}
#[test]
fn kv_bytes_per_token_missing_dims_is_none() {
assert_eq!(kv_bytes_per_token(&ModelDims::default()), None);
assert_eq!(
kv_bytes_per_token(&ModelDims {
block_count: 28,
head_count: 0, head_count_kv: 4,
embedding_length: 3584,
}),
None
);
}
#[test]
fn max_tokens_for_memory_subtracts_weights_and_headroom() {
let budget = 8 * 1024 * 1024 * 1024u64;
let weight = 5_600_000_000u64;
let kv = 57_344usize;
let got = max_tokens_for_memory(budget, weight, kv).unwrap();
assert!((28_000..=32_000).contains(&got), "expected ~30k, got {got}");
}
#[test]
fn converge_shrinks_to_clear_kv_overflow() {
let new = converge_num_ctx(10_000, 10_000_000_000, 12_000_000_000, 1_000_000)
.expect("should shrink");
assert!(new < 10_000, "must make progress, got {new}");
assert!((6_000..=8_000).contains(&new), "expected ~7-8k, got {new}");
}
#[test]
fn converge_none_when_overflow_is_weights_bound() {
let new = converge_num_ctx(8_000, 1_000_000_000, 9_000_000_000, 1_000_000);
assert_eq!(new, None);
}
#[test]
fn converge_some_only_when_shrink_actually_clears_overflow() {
let new = converge_num_ctx(16_000, 10_000_000_000, 16_000_000_000, 1_000_000)
.expect("clearable overflow should shrink");
assert!(
(OLLAMA_MIN_AUTO_NUM_CTX..16_000).contains(&new),
"got {new}"
);
assert_eq!(
converge_num_ctx(5_000, 10_000_000_000, 16_000_000_000, 1_000_000),
None
);
}
#[test]
fn converge_none_when_already_at_floor() {
let new = converge_num_ctx(
OLLAMA_MIN_AUTO_NUM_CTX,
1_000_000_000,
9_000_000_000,
1_000_000,
);
assert_eq!(new, None);
}
#[test]
fn converge_none_when_it_fits_or_kv_unknown() {
assert_eq!(
converge_num_ctx(8_000, 6_000_000_000, 6_000_000_000, 40_000),
None
);
assert_eq!(
converge_num_ctx(8_000, 6_000_000_000, 5_000_000_000, 40_000),
None
);
assert_eq!(
converge_num_ctx(8_000, 1_000_000_000, 9_000_000_000, 0),
None
);
}
#[test]
fn override_beats_everything() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
per_model_override: Some(131_072),
global_num_ctx: Some(8_192),
vram_bytes: Some(8 * 1024 * 1024 * 1024),
dims: Some(gqa_dims()),
..Default::default()
})
.unwrap();
assert_eq!(res.value, 131_072);
assert_eq!(res.source, NumCtxSource::Override);
}
#[test]
fn global_config_beats_auto() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
global_num_ctx: Some(16_384),
vram_bytes: Some(8 * 1024 * 1024 * 1024),
dims: Some(gqa_dims()),
..Default::default()
})
.unwrap();
assert_eq!(res.value, 16_384);
assert_eq!(res.source, NumCtxSource::GlobalConfig);
}
#[test]
fn cloud_model_uses_full_window_ignoring_vram_and_global() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(524_288),
is_cloud: true,
dims: Some(gqa_dims()),
vram_bytes: Some(8 * 1024 * 1024 * 1024), global_num_ctx: Some(16_384), ..Default::default()
})
.unwrap();
assert_eq!(res.value, 524_288, "cloud uses the full advertised window");
assert_eq!(res.source, NumCtxSource::Cloud);
assert!(!res.source.is_auto(), "cloud is not a VRAM auto-fit");
}
#[test]
fn cloud_model_still_honors_explicit_override() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(524_288),
is_cloud: true,
per_model_override: Some(65_536),
..Default::default()
})
.unwrap();
assert_eq!(res.value, 65_536);
assert_eq!(res.source, NumCtxSource::Override);
}
#[test]
fn cloud_model_without_known_window_omits_num_ctx() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: None,
is_cloud: true,
..Default::default()
});
assert!(res.is_none());
}
#[test]
fn auto_fits_vram_and_caps_at_model_max() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(8_192),
dims: Some(gqa_dims()),
model_weight_bytes: Some(5_600_000_000),
vram_bytes: Some(80 * 1024 * 1024 * 1024),
..Default::default()
})
.unwrap();
assert_eq!(res.value, 8_192);
assert_eq!(res.source, NumCtxSource::Auto);
}
#[test]
fn auto_fits_vram_below_model_max() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
dims: Some(gqa_dims()),
model_weight_bytes: Some(5_600_000_000),
vram_bytes: Some(8 * 1024 * 1024 * 1024),
..Default::default()
})
.unwrap();
assert_eq!(res.source, NumCtxSource::Auto);
assert!(res.value < 262_144 && res.value >= OLLAMA_MIN_AUTO_NUM_CTX);
assert_eq!(res.value % OLLAMA_NUM_CTX_ROUNDING, 0);
}
#[test]
fn offload_on_uses_system_ram() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
dims: Some(gqa_dims()),
model_weight_bytes: Some(5_600_000_000),
allow_ram_offload: true,
system_ram_bytes: Some(64 * 1024 * 1024 * 1024),
vram_bytes: None,
..Default::default()
})
.unwrap();
assert_eq!(res.source, NumCtxSource::Auto);
assert!(
res.value > DEFAULT_OLLAMA_MAX_AUTO_NUM_CTX,
"64GB RAM should fit > 32k, got {}",
res.value
);
}
#[test]
fn no_memory_detected_uses_fallback_cap() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
dims: Some(gqa_dims()),
vram_bytes: None,
system_ram_bytes: None,
..Default::default()
})
.unwrap();
assert_eq!(res.value, DEFAULT_OLLAMA_MAX_AUTO_NUM_CTX);
assert_eq!(res.source, NumCtxSource::AutoFallback);
}
#[test]
fn max_auto_cap_clamps_auto() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(262_144),
dims: Some(gqa_dims()),
model_weight_bytes: Some(5_600_000_000),
vram_bytes: Some(80 * 1024 * 1024 * 1024),
max_auto_cap: Some(16_384),
..Default::default()
})
.unwrap();
assert_eq!(res.value, 16_384);
}
#[test]
fn floor_never_exceeds_model_max() {
let res = resolve_ollama_num_ctx(&NumCtxInputs {
model_max: Some(2_048),
vram_bytes: None,
..Default::default()
})
.unwrap();
assert_eq!(res.value, 2_048);
}
#[test]
fn no_model_max_and_no_explicit_is_none() {
assert!(resolve_ollama_num_ctx(&NumCtxInputs::default()).is_none());
}
#[test]
fn num_predict_auto_uses_full_room() {
assert_eq!(
default_ollama_num_predict(0, Some(131_072), 1_000, None),
Some(131_072 - 1_000 - 256)
);
}
#[test]
fn num_predict_auto_omits_without_ctx() {
assert_eq!(default_ollama_num_predict(0, None, 1_000, None), None);
}
#[test]
fn num_predict_hard_cap_is_exact_and_room_bounded() {
assert_eq!(
default_ollama_num_predict(4_096, Some(131_072), 1_000, None),
Some(4_096)
);
assert_eq!(
default_ollama_num_predict(4_096, Some(8_192), 7_000, None),
Some(8_192 - 7_000 - 256)
);
}
#[test]
fn num_predict_floored_when_room_tiny() {
assert_eq!(
default_ollama_num_predict(4_096, Some(8_192), 8_100, None),
Some(512)
);
}
}