use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ModelTier {
Fast,
Capable,
Premium,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum TaskType {
Summarize,
Classify,
Code,
Plan,
Reason,
Research,
MultiStep,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Complexity {
Simple,
Medium,
Complex,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub struct TaskProfile {
pub task_type: TaskType,
pub complexity: Complexity,
}
impl TaskProfile {
#[must_use]
#[inline]
pub fn new(task_type: TaskType, complexity: Complexity) -> Self {
Self {
task_type,
complexity,
}
}
}
#[must_use]
pub fn route(profile: &TaskProfile) -> ModelTier {
use Complexity::*;
use TaskType::*;
let tier = match (&profile.task_type, &profile.complexity) {
(Summarize | Classify, Simple | Medium) => ModelTier::Fast,
(Summarize | Classify, Complex) => ModelTier::Capable,
(Code | Plan | Reason, Simple | Medium) => ModelTier::Capable,
(Code | Plan | Reason, Complex) => ModelTier::Premium,
(Research | MultiStep, Simple) => ModelTier::Capable,
(Research | MultiStep, Medium | Complex) => ModelTier::Premium,
};
tracing::debug!(
task_type = ?profile.task_type,
complexity = ?profile.complexity,
tier = ?tier,
"model tier selected"
);
tier
}
#[must_use]
pub fn default_model(tier: ModelTier) -> &'static str {
match tier {
ModelTier::Fast => "llama3",
ModelTier::Capable => "llama3:70b",
ModelTier::Premium => "llama3:405b",
}
}
#[must_use]
pub fn parse_complexity(s: &str) -> Complexity {
if s.eq_ignore_ascii_case("low") || s.eq_ignore_ascii_case("simple") {
Complexity::Simple
} else if s.eq_ignore_ascii_case("high") || s.eq_ignore_ascii_case("complex") {
Complexity::Complex
} else {
Complexity::Medium
}
}
#[cfg(feature = "hwaccel")]
pub fn suggest_quantization(
model_params: u64,
registry: &ai_hwaccel::AcceleratorRegistry,
) -> ai_hwaccel::QuantizationLevel {
registry.suggest_quantization(model_params)
}
#[cfg(feature = "hwaccel")]
pub fn estimate_model_memory(model_params: u64, quant: &ai_hwaccel::QuantizationLevel) -> u64 {
ai_hwaccel::AcceleratorRegistry::estimate_memory(model_params, quant)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn summarize_simple_is_fast() {
let p = TaskProfile {
task_type: TaskType::Summarize,
complexity: Complexity::Simple,
};
assert_eq!(route(&p), ModelTier::Fast);
}
#[test]
fn classify_medium_is_fast() {
let p = TaskProfile {
task_type: TaskType::Classify,
complexity: Complexity::Medium,
};
assert_eq!(route(&p), ModelTier::Fast);
}
#[test]
fn summarize_complex_is_capable() {
let p = TaskProfile {
task_type: TaskType::Summarize,
complexity: Complexity::Complex,
};
assert_eq!(route(&p), ModelTier::Capable);
}
#[test]
fn code_simple_is_capable() {
let p = TaskProfile {
task_type: TaskType::Code,
complexity: Complexity::Simple,
};
assert_eq!(route(&p), ModelTier::Capable);
}
#[test]
fn code_complex_is_premium() {
let p = TaskProfile {
task_type: TaskType::Code,
complexity: Complexity::Complex,
};
assert_eq!(route(&p), ModelTier::Premium);
}
#[test]
fn reason_medium_is_capable() {
let p = TaskProfile {
task_type: TaskType::Reason,
complexity: Complexity::Medium,
};
assert_eq!(route(&p), ModelTier::Capable);
}
#[test]
fn research_simple_is_capable() {
let p = TaskProfile {
task_type: TaskType::Research,
complexity: Complexity::Simple,
};
assert_eq!(route(&p), ModelTier::Capable);
}
#[test]
fn research_medium_is_premium() {
let p = TaskProfile {
task_type: TaskType::Research,
complexity: Complexity::Medium,
};
assert_eq!(route(&p), ModelTier::Premium);
}
#[test]
fn multistep_complex_is_premium() {
let p = TaskProfile {
task_type: TaskType::MultiStep,
complexity: Complexity::Complex,
};
assert_eq!(route(&p), ModelTier::Premium);
}
#[test]
fn plan_complex_is_premium() {
let p = TaskProfile {
task_type: TaskType::Plan,
complexity: Complexity::Complex,
};
assert_eq!(route(&p), ModelTier::Premium);
}
#[test]
fn default_model_fast() {
assert_eq!(default_model(ModelTier::Fast), "llama3");
}
#[test]
fn default_model_capable() {
assert_eq!(default_model(ModelTier::Capable), "llama3:70b");
}
#[test]
fn default_model_premium() {
assert_eq!(default_model(ModelTier::Premium), "llama3:405b");
}
#[test]
fn parse_complexity_variants() {
assert_eq!(parse_complexity("low"), Complexity::Simple);
assert_eq!(parse_complexity("simple"), Complexity::Simple);
assert_eq!(parse_complexity("medium"), Complexity::Medium);
assert_eq!(parse_complexity("high"), Complexity::Complex);
assert_eq!(parse_complexity("complex"), Complexity::Complex);
assert_eq!(parse_complexity("HIGH"), Complexity::Complex);
assert_eq!(parse_complexity("unknown"), Complexity::Medium);
}
#[cfg(feature = "hwaccel")]
mod hwaccel_tests {
use super::super::*;
use ai_hwaccel::{AcceleratorProfile, AcceleratorRegistry, QuantizationLevel};
#[test]
fn suggest_quantization_small_model_high_vram() {
let registry = AcceleratorRegistry::from_profiles(vec![
AcceleratorProfile::cpu(64 * 1024 * 1024 * 1024),
AcceleratorProfile::cuda(0, 80 * 1024 * 1024 * 1024),
]);
let quant = suggest_quantization(7_000_000_000, ®istry);
assert!(
quant.bits_per_param() >= 16,
"7B model with 80GB should get at least FP16, got {:?}",
quant
);
}
#[test]
fn suggest_quantization_large_model_small_vram() {
let registry = AcceleratorRegistry::from_profiles(vec![
AcceleratorProfile::cpu(32 * 1024 * 1024 * 1024),
AcceleratorProfile::cuda(0, 24 * 1024 * 1024 * 1024),
]);
let quant = suggest_quantization(70_000_000_000, ®istry);
assert!(
quant.bits_per_param() < 16,
"70B model with 24GB should be quantized below FP16, got {:?}",
quant
);
}
#[test]
fn estimate_model_memory_scales_with_quantization() {
let fp32 = estimate_model_memory(7_000_000_000, &QuantizationLevel::None);
let fp16 = estimate_model_memory(7_000_000_000, &QuantizationLevel::Float16);
let int4 = estimate_model_memory(7_000_000_000, &QuantizationLevel::Int4);
assert!(fp32 > fp16, "FP32 should use more memory than FP16");
assert!(fp16 > int4, "FP16 should use more memory than INT4");
}
}
}