diff --git a/benchmarks/run_all.py b/benchmarks/run_all.py index a9931c8..bf4b8ac 100755 --- a/benchmarks/run_all.py +++ b/benchmarks/run_all.py @@ -220,43 +220,6 @@ def main(): ], "passive-test", ) - tools_dev_dir = run_step( - "tools-lane-dev", - [ - "cargo", - "run", - "-q", - "-p", - "klbr-bench", - "--", - "tools-lane", - args.active_dataset, - args.config_dev, - args.router_model, - run_dir / "tools-lane-dev", - args.llm_url, - ], - "tools-lane-dev", - ) - tools_test_dir = run_step( - "tools-lane-test", - [ - "cargo", - "run", - "-q", - "-p", - "klbr-bench", - "--", - "tools-lane", - args.active_dataset, - args.config_test, - args.router_model, - run_dir / "tools-lane-test", - args.llm_url, - ], - "tools-lane-test", - ) - (run_dir / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n", "utf-8") sections = [] @@ -265,8 +228,6 @@ def main(): ("Active Recall (Retrieval) — Test", retrieval_test_dir), ("Passive Recall — Dev", passive_dev_dir), ("Passive Recall — Test", passive_test_dir), - ("Tools Lane — Dev", tools_dev_dir), - ("Tools Lane — Test", tools_test_dir), ]: report = out_dir / "report.md" metrics = parse_metrics(report) if report.exists() else {} diff --git a/klbr-bench/src/main.rs b/klbr-bench/src/main.rs index c976970..9801f3b 100644 --- a/klbr-bench/src/main.rs +++ b/klbr-bench/src/main.rs @@ -10,8 +10,7 @@ use std::{ use anyhow::{bail, Context, Result}; use klbr_core::{ - config::Config, - context::{format_recalled_memories, Context as AgentContext, RecalledMemory}, + context::{format_recalled_memories, RecalledMemory}, memory::MemoryStore, models::{LlmClient, Message, ModelsConfig}, mvp::{ @@ -25,8 +24,7 @@ use klbr_core::{ retrieval::{self, RetrievalConfig}, router::{ LinearDecisionParams as RuntimeLinearDecisionParams, - LinearRouterModel as RuntimeLinearRouterModel, RouteDecision as RuntimeRouteDecision, - Router as RuntimeRouter, RouterScores as RuntimeRouterScores, + LinearRouterModel as RuntimeLinearRouterModel, }, tools as core_tools, }; @@ -64,7 +62,6 @@ struct PreparedBenchmark { } const PASSIVE_RECALL_TOP_K: usize = 3; -const MAX_TOOLS_LANE_TOOL_ITERATIONS: usize = 20; fn truncate(s: &str, max_chars: usize) -> String { if s.chars().count() <= max_chars { @@ -75,84 +72,7 @@ fn truncate(s: &str, max_chars: usize) -> String { out } -const TOOLS_LANE_ANCHOR: &str = r#" -you are running in tool lane. -- do not guess. if the answer depends on runtime state or source code, call tools. -- prefer `read_file` for inspecting repo files; prefer `shell` for runtime state. -- if a tool result indicates failure (starts with `error:` or non-zero exit), do not repeat the same approach. try a different tool once, then answer with what you learned. -- if the question's premise seems wrong (e.g., mentions postgres), quickly verify by searching the repo; if you find no evidence, answer that it's not used here. -- for GPU questions: `nvidia-smi` may be unavailable. prefer sysfs (`/sys/class/drm/*/device/uevent`, `vendor`, `device`, `driver`) to identify the GPU; use `rocminfo` if present. -- for vector-db / extension questions: check for sqlite + sqlite-vec usage (e.g. `sqlite-vec`, `sqlite3_vec_init`, `sqlite3_auto_extension`) before assuming postgres/pgvector. -- stop calling tools once you have enough evidence to answer. avoid infinite tool loops. -- keep the response short and factual. -"#; - -#[derive(Debug, Clone, Serialize)] -struct ToolsLaneToolCallLog { - name: String, - args: String, - status: String, // ok|blocked|error - result: String, -} - -#[derive(Debug, Clone, Serialize)] -struct ToolsLaneTrace { - query_id: String, - split: DatasetSplit, - text: String, - gold: EvalRouteLabel, - required_tools: Vec, - router_decision: EvalRouteLabel, - #[serde(skip_serializing_if = "Option::is_none")] - router_scores: Option, - router_correct: bool, - tool_iterations: usize, - tool_calls: Vec, - final_response: String, - called_any_tool: bool, - called_any_required_tool: Option, -} - -#[derive(Debug, Clone, Serialize)] -struct ToolsLaneRouterScores { - tools: f32, - memory: f32, - abstain: f32, - kind: String, -} - -#[derive(Debug, Clone, Serialize, Default)] -struct ToolsLaneMetrics { - query_count: usize, - gold_tools: usize, - gold_memory: usize, - gold_abstain: usize, - router_accuracy: f32, - router_tools_recall: f32, - router_tools_precision: f32, - router_memory_false_tools_rate: f32, - router_abstain_false_tools_rate: f32, - router_routed_tools_count: usize, - tools_lane_gold_tools_called_any_rate: f32, - tools_lane_gold_memory_called_any_rate: f32, - tools_lane_any_tool_calls_count: usize, - tools_lane_required_tool_hit_rate: Option, - tools_lane_tool_exec_error_rate: f32, -} - -#[derive(Debug, Clone, Serialize)] -struct ToolsLaneReport { - dataset_id: String, - experiment_id: String, - evaluated_split: Option, - llm_url: String, - embed_url: String, - embed_model: String, - router_model_path: String, - write_root: String, - metrics: ToolsLaneMetrics, -} #[derive(Debug, Clone, Copy, Serialize)] struct SweepRange { @@ -243,7 +163,7 @@ async fn main() -> Result<()> { let args: Vec = env::args().collect(); if args.len() < 2 { bail!( - "usage:\n cargo run -p klbr-bench -- run --suite longmemeval-s --data --out [--retrieval-only]\n cargo run -p klbr-bench -- retrieval \n cargo run -p klbr-bench -- passive-recall \n cargo run -p klbr-bench -- router \n cargo run -p klbr-bench -- router-multi [dataset2.json ...]\n cargo run -p klbr-bench -- router-multi-linear [dataset2.json ...]\n cargo run -p klbr-bench -- tools-lane [llm_url]\n cargo run -p klbr-bench -- longmem [llm_url]\n cargo run -p klbr-bench -- longmem-retrieval [split.json subset]\n cargo run -p klbr-bench -- dump-tools \n cargo run -p klbr-bench -- sweep [score_start score_end score_step margin_start margin_end margin_step [support_start support_end support_step]]" + "usage:\n cargo run -p klbr-bench -- run --suite longmemeval-s --data --out [--retrieval-only]\n cargo run -p klbr-bench -- retrieval \n cargo run -p klbr-bench -- passive-recall \n cargo run -p klbr-bench -- router \n cargo run -p klbr-bench -- router-multi [dataset2.json ...]\n cargo run -p klbr-bench -- router-multi-linear [dataset2.json ...]\n cargo run -p klbr-bench -- longmem [llm_url]\n cargo run -p klbr-bench -- longmem-retrieval [split.json subset]\n cargo run -p klbr-bench -- dump-tools \n cargo run -p klbr-bench -- sweep [score_start score_end score_step margin_start margin_end margin_step [support_start support_end support_step]]" ); } @@ -289,15 +209,6 @@ async fn main() -> Result<()> { } run_router_multi_linear_command(&args[2], &args[3], &args[4..]).await } - "tools-lane" => { - if args.len() != 6 && args.len() != 7 { - bail!( - "usage: cargo run -p klbr-bench -- tools-lane [llm_url]" - ); - } - let llm_url = args.get(6).cloned(); - run_tools_lane_command(&args[2], &args[3], &args[4], &args[5], llm_url).await - } "dump-tools" => { if args.len() != 3 { bail!("usage: cargo run -p klbr-bench -- dump-tools "); @@ -388,7 +299,7 @@ async fn main() -> Result<()> { run_sweep_command(&args[2], &args[3], &args[4], grid).await } other => bail!( - "unknown subcommand '{}'; expected 'retrieval', 'passive-recall', 'router', 'router-multi', 'router-multi-linear', 'tools-lane', 'dump-tools', or 'sweep'", + "unknown subcommand '{}'; expected 'retrieval', 'passive-recall', 'router', 'router-multi', 'router-multi-linear', 'dump-tools', or 'sweep'", other ), } @@ -407,110 +318,6 @@ async fn run_dump_tools_command(output_path: &str) -> Result<()> { Ok(()) } -async fn run_tools_lane_command( - dataset_path: &str, - config_path: &str, - router_model_path: &str, - output_dir: &str, - llm_url_override: Option, -) -> Result<()> { - let dataset_path = Path::new(dataset_path); - let config_path = Path::new(config_path); - let output_dir = PathBuf::from(output_dir); - fs::create_dir_all(&output_dir)?; - - let dataset: InternalEvalDataset = read_json(dataset_path) - .with_context(|| format!("failed to read dataset from {}", dataset_path.display()))?; - let experiment: RetrievalExperimentConfig = read_json(config_path) - .with_context(|| format!("failed to read config from {}", config_path.display()))?; - - let mut runtime = ModelsConfig::default(); - runtime.embedder.url = normalize_url(&experiment.embed_url); - runtime.embedder.model = experiment.embed_model.clone(); - runtime.embed_dim = experiment.embed_dim; - if let Some(url) = &experiment.rerank_url { - runtime.reranker.url = normalize_url(url); - } - if let Some(url) = llm_url_override { - runtime.llm.url = normalize_url(&url); - } else if let Ok(url) = env::var("KLBR_BENCH_LLM_URL") { - if !url.trim().is_empty() { - runtime.llm.url = normalize_url(&url); - } - } - let llm = LlmClient::new(runtime.clone()); - - let router = RuntimeRouter::load(Path::new(router_model_path)) - .with_context(|| format!("failed to load router model from {router_model_path}"))?; - - let write_root = output_dir.join("tool_lane_writes"); - fs::create_dir_all(&write_root)?; - - let db_path = output_dir.join("tool_lane_agent.db"); - let memory = MemoryStore::open(db_path.to_string_lossy().as_ref(), experiment.embed_dim) - .context("failed to create tool-lane memory store")?; - let tool_ctx = core_tools::ToolContext::new( - memory, - llm.clone(), - llm.clone(), - Config::default().memory.sim_threshold, - Config::default().memory.verbatim_count, - ); - let registry = core_tools::all_tools(); - let tool_defs = registry.definitions(); - - let mut traces = Vec::new(); - for query in &dataset.queries { - if query.objective != EvalObjective::AnswerQuery { - continue; - } - if let Some(split) = &experiment.dataset_split { - if &query.split != split { - continue; - } - } - let Some(gold) = gold_route_label(query) else { - continue; - }; - - traces.push( - run_tools_lane_query( - query, - gold, - &router, - &llm, - &tool_ctx, - ®istry, - &tool_defs, - &write_root, - ) - .await?, - ); - } - - let report = summarize_tools_lane( - &dataset, - &experiment, - router_model_path, - &llm, - &write_root, - &traces, - ); - let report_md = render_tools_lane_report(&report, &traces); - - fs::write( - output_dir.join("report.json"), - serde_json::to_vec_pretty(&report)?, - )?; - fs::write(output_dir.join("report.md"), report_md)?; - fs::write( - output_dir.join("traces.json"), - serde_json::to_vec_pretty(&traces)?, - )?; - - println!("wrote tools-lane outputs to {}", output_dir.display()); - Ok(()) -} async fn run_retrieval_command( dataset_path: &str, @@ -571,7 +378,6 @@ struct RouterModel { trained_on: DatasetSplit, threshold: f32, memory_centroids: Vec>, - tools_centroids: Vec>, abstain_centroids: Vec>, } @@ -579,11 +385,8 @@ struct RouterModel { struct RouterMetrics { total: usize, accuracy: f32, - tool_precision: f32, - tool_recall: f32, - memory_false_tool_rate: f32, - abstain_false_tool_rate: f32, - tool_count: usize, + memory_false_abstain_rate: f32, + abstain_false_memory_rate: f32, memory_count: usize, abstain_count: usize, misclassified: Vec, @@ -595,7 +398,6 @@ struct RouterTrace { text: String, gold: EvalRouteLabel, pred: EvalRouteLabel, - sim_tools: f32, sim_memory: f32, sim_abstain: f32, } @@ -614,13 +416,9 @@ struct RouterReport { struct LinearRouterMetrics { total: usize, accuracy: f32, - confusion: [[usize; 3]; 3], // [gold][pred] - tool_precision: f32, - tool_recall: f32, - memory_false_tool_rate: f32, - abstain_false_tool_rate: f32, + confusion: [[usize; 2]; 2], // [gold][pred] memory_false_abstain_rate: f32, - tool_count: usize, + abstain_false_memory_rate: f32, memory_count: usize, abstain_count: usize, misclassified: Vec, @@ -632,7 +430,6 @@ struct LinearRouterTrace { text: String, gold: EvalRouteLabel, pred: EvalRouteLabel, - p_tools: f32, p_memory: f32, p_abstain: f32, } @@ -824,9 +621,7 @@ async fn run_router_multi_linear_command( train_softmax_router(&examples, train_split.clone(), experiment.embed_dim)?; eprintln!("router-linear: tuning thresholds on dev split"); - let tuned = tune_linear_thresholds(&examples, DatasetSplit::Dev, &weights, &bias)?; - let decision = select_linear_decision_with_holdouts(&examples, &weights, &bias, tuned.clone()) - .unwrap_or(tuned); + let decision = tune_linear_thresholds(&examples, DatasetSplit::Dev, &weights, &bias)?; let model = RuntimeLinearRouterModel { kind: "linear_softmax_v1".into(), @@ -835,7 +630,7 @@ async fn run_router_multi_linear_command( embed_dim: experiment.embed_dim, trained_on: format!("{:?}", train_split).to_lowercase(), dataset_id: Some(dataset.dataset_id.clone()), - labels: vec!["tools".into(), "memory".into(), "abstain".into()], + labels: vec!["memory".into(), "abstain".into()], weights: weights.clone(), bias: bias.clone(), decision: decision.clone(), @@ -862,211 +657,6 @@ async fn run_router_multi_linear_command( Ok(()) } -fn select_linear_decision_with_holdouts( - examples: &[RouterExample], - weights: &[Vec], - bias: &[f32], - tuned: RuntimeLinearDecisionParams, -) -> Result { - // Optional post-tune selection: - // Use holdout buckets (query_id starts with holdout_) as a safety/behavior constraint. - // This helps avoid "dev-only" tuning picks that crater tools recall on tools_recall holdouts. - // - // NOTE: this is intentionally conservative: if constraints can't be met, we fall back to `tuned`. - const MAX_DEV_ABSTAIN_FALSE_TOOL_RATE: f32 = 0.10; - const MAX_DEV_MEMORY_FALSE_ABSTAIN_RATE: f32 = 0.03; - const MAX_HOLDOUT_ABSTAIN_FALSE_TOOL_RATE: f32 = 0.10; - - let dev: Vec<&RouterExample> = examples - .iter() - .filter(|e| e.split == DatasetSplit::Dev) - .collect(); - let test: Vec<&RouterExample> = examples - .iter() - .filter(|e| e.split == DatasetSplit::Test) - .collect(); - if dev.is_empty() || test.is_empty() { - return Ok(tuned); - } - let dev_points = precompute_linear_points(&dev, weights, bias)?; - let test_points = precompute_linear_points(&test, weights, bias)?; - let holdout_points: Vec = test_points - .iter() - .filter(|p| p.query_id.starts_with("holdout_")) - .cloned() - .collect(); - if holdout_points.is_empty() { - return Ok(tuned); - } - - let mut holdout_by_bucket: BTreeMap> = BTreeMap::new(); - for pt in &holdout_points { - if let Some(bucket) = parse_holdout_bucket(&pt.query_id) { - holdout_by_bucket - .entry(bucket) - .or_default() - .push(pt.clone()); - } - } - - let tools_prob_values: Vec = (2..=18).map(|i| i as f32 * 0.05).collect(); // 0.10..0.90 - let tools_margin_values: Vec = (0..=25).map(|i| i as f32 * 0.02).collect(); // 0.00..0.50 - - let mut best: Option<(RuntimeLinearDecisionParams, (f32, f32, f32))> = None; - let mut best_must: Option<(RuntimeLinearDecisionParams, LinearRouterMetrics, f32, f32)> = None; - for &t_tools_prob in &tools_prob_values { - for &t_tools_margin in &tools_margin_values { - let decision = RuntimeLinearDecisionParams { - tools_prob_threshold: t_tools_prob, - tools_margin_threshold: t_tools_margin, - // keep abstain thresholds from the dev-tuned point (we're selecting tools thresholds). - abstain_prob_threshold: tuned.abstain_prob_threshold, - abstain_margin_threshold: tuned.abstain_margin_threshold, - }; - - // Check must-tools routing first (fast), so we can optionally emit diagnostics - // if constraints prevent satisfying these key examples. - let mut routes_must_tools = true; - for must_tools in ["test_q9", "test_q15"] { - if let Some(pt) = test_points.iter().find(|p| p.query_id == must_tools) { - let pred = predict_linear_route(pt.probs, &decision); - if pred != EvalRouteLabel::Tools { - routes_must_tools = false; - break; - } - } - } - - let dev_m = score_linear_thresholds_precomputed(&dev_points, &decision)?; - if routes_must_tools { - let tools_recall_holdout = holdout_by_bucket - .get("tools_recall") - .map(|pts| score_linear_thresholds_precomputed(pts, &decision).ok()) - .flatten() - .map(|m| m.tool_recall) - .unwrap_or(0.0); - let abstain_false_tool_holdout = holdout_by_bucket - .get("abstain_toolish") - .map(|pts| score_linear_thresholds_precomputed(pts, &decision).ok()) - .flatten() - .map(|m| m.abstain_false_tool_rate) - .unwrap_or(0.0); - best_must = match best_must.take() { - None => Some(( - decision.clone(), - dev_m.clone(), - tools_recall_holdout, - abstain_false_tool_holdout, - )), - Some((bd, bm, btr, baft)) => { - // Prefer higher tools_recall on tools_recall holdout, then lower abstain->tools. - let cur = (tools_recall_holdout, -abstain_false_tool_holdout); - let bestt = (btr, -baft); - if cur > bestt { - Some(( - decision.clone(), - dev_m.clone(), - tools_recall_holdout, - abstain_false_tool_holdout, - )) - } else { - Some((bd, bm, btr, baft)) - } - } - }; - } - if dev_m.memory_false_tool_rate != 0.0 { - continue; - } - if dev_m.memory_false_abstain_rate > MAX_DEV_MEMORY_FALSE_ABSTAIN_RATE { - continue; - } - if dev_m.abstain_false_tool_rate > MAX_DEV_ABSTAIN_FALSE_TOOL_RATE { - continue; - } - - // Ensure a couple of known tool-needed internal_eval_starter queries route to Tools. - // These are exactly the kind of failures we want the router to prevent in practice. - if !routes_must_tools { - continue; - } - - let tools_recall_holdout = holdout_by_bucket - .get("tools_recall") - .map(|pts| score_linear_thresholds_precomputed(pts, &decision).ok()) - .flatten() - .map(|m| m.tool_recall) - .unwrap_or(0.0); - - let abstain_false_tool_holdout = holdout_by_bucket - .get("abstain_toolish") - .map(|pts| score_linear_thresholds_precomputed(pts, &decision).ok()) - .flatten() - .map(|m| m.abstain_false_tool_rate) - .unwrap_or(0.0); - if abstain_false_tool_holdout > MAX_HOLDOUT_ABSTAIN_FALSE_TOOL_RATE { - continue; - } - - let overall_holdout_acc = - score_linear_thresholds_precomputed(&holdout_points, &decision)?.accuracy; - - // Objective: maximize tools_recall on tools_recall holdout bucket. - // Tie-break: minimize abstain->tools leakage on abstain_toolish holdout. - // Then: maximize overall holdout accuracy. - let score_tuple = ( - tools_recall_holdout, - -abstain_false_tool_holdout, - overall_holdout_acc, - ); - best = match best.take() { - None => Some((decision, score_tuple)), - Some((best_d, best_tuple)) => { - if score_tuple > best_tuple { - Some((decision, score_tuple)) - } else { - Some((best_d, best_tuple)) - } - } - }; - } - } - - if let Some((picked, (tr, neg_aft, acc))) = best { - let aft = -neg_aft; - let changed = picked.tools_prob_threshold != tuned.tools_prob_threshold - || picked.tools_margin_threshold != tuned.tools_margin_threshold - || picked.abstain_prob_threshold != tuned.abstain_prob_threshold - || picked.abstain_margin_threshold != tuned.abstain_margin_threshold; - if changed { - eprintln!( - "router-linear: picked tools thresholds via holdouts tools_prob={:.2} tools_margin={:.2} (holdout tools_recall={:.3}, abstain→tools={:.3}, acc={:.3})", - picked.tools_prob_threshold, - picked.tools_margin_threshold, - tr, - aft, - acc - ); - } - return Ok(picked); - } - if let Some((d, dm, tr, aft)) = best_must { - eprintln!( - "router-linear: note: best decision that routes test_q9/test_q15 is tools_prob={:.2} tools_margin={:.2} (dev: tools_recall={:.3} precision={:.3} mem→tools={:.3} abstain→tools={:.3} mem→abstain={:.3}; holdout tools_recall={:.3} abstain→tools={:.3})", - d.tools_prob_threshold, - d.tools_margin_threshold, - dm.tool_recall, - dm.tool_precision, - dm.memory_false_tool_rate, - dm.abstain_false_tool_rate, - dm.memory_false_abstain_rate, - tr, - aft, - ); - } - Ok(tuned) -} - fn rebalance_router_train_dev_if_needed(examples: &[RouterExample]) -> Vec { // Tuned for our current dataset sizes; keep it simple and deterministic. const MIN_TRAIN_TOTAL: usize = 120; @@ -1147,562 +737,9 @@ fn validate_router_inputs(experiment: &RetrievalExperimentConfig) -> Result<()> Ok(()) } -fn route_decision_to_label(d: RuntimeRouteDecision) -> EvalRouteLabel { - match d { - RuntimeRouteDecision::Tools => EvalRouteLabel::Tools, - RuntimeRouteDecision::Memory => EvalRouteLabel::Memory, - RuntimeRouteDecision::Abstain => EvalRouteLabel::Abstain, - } -} - -fn scores_for_trace(scores: RuntimeRouterScores) -> ToolsLaneRouterScores { - ToolsLaneRouterScores { - tools: scores.tools, - memory: scores.memory, - abstain: scores.abstain, - kind: scores.kind.to_string(), - } -} - -fn is_shell_blocked(cmd: &str) -> bool { - let s = cmd.to_lowercase(); - let blocked_substrings = [ - "rm -rf", - " sudo ", - "\nsudo ", - "mkfs", - " dd ", - "\ndd ", - "mount ", - "umount ", - "chmod ", - "chown ", - "shred ", - "truncate ", - ">>", - " | tee ", - ]; - if blocked_substrings.iter().any(|pat| s.contains(pat)) { - return true; - } - // block redirects that write to files. allow common safe redirections to /dev/null and fd wiring. - if s.contains('>') { - let mut t = s.as_str().to_string(); - for pat in [ - "2>/dev/null", - "1>/dev/null", - ">/dev/null", - "&>/dev/null", - "2>&1", - "1>&2", - ] { - t = t.replace(pat, ""); - } - if t.contains('>') { - return true; - } - } - false -} - -fn normalize_path_lexical(path: &std::path::Path) -> std::path::PathBuf { - use std::path::Component; - let mut out = std::path::PathBuf::new(); - for c in path.components() { - match c { - Component::CurDir => {} - Component::ParentDir => { - out.pop(); - } - other => out.push(other.as_os_str()), - } - } - out -} - -fn resolve_abs_path(path: &str) -> Result { - let p = PathBuf::from(path); - if p.is_absolute() { - Ok(normalize_path_lexical(&p)) - } else { - let cwd = env::current_dir().context("failed to get cwd")?; - Ok(normalize_path_lexical(&cwd.join(p))) - } -} - -async fn execute_tool_restricted( - registry: &core_tools::Subroutines, - call: &klbr_core::models::ToolCall, - tool_ctx: &core_tools::ToolContext, - write_root: &Path, -) -> ToolsLaneToolCallLog { - let name = call.function.name.clone(); - let args_str = call.function.arguments.clone(); - let args_json: serde_json::Value = serde_json::from_str(&args_str).unwrap_or_default(); - - let mut status = "ok".to_string(); - let mut blocked_reason = None::; - - match name.as_str() { - "write_file" => { - let path = args_json - .get("path") - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - match resolve_abs_path(&path) { - Ok(abs) => { - let allowed = normalize_path_lexical(write_root); - if !abs.starts_with(&allowed) { - status = "blocked".into(); - blocked_reason = Some(format!( - "error: write_file blocked by bench sandbox (allowed root: {})", - allowed.display() - )); - } - } - Err(e) => { - status = "blocked".into(); - blocked_reason = Some(format!("error: write_file path resolution failed: {e}")); - } - } - } - "shell" => { - let cmd = args_json.get("cmd").and_then(|v| v.as_str()).unwrap_or(""); - if is_shell_blocked(cmd) { - status = "blocked".into(); - blocked_reason = Some("error: shell blocked by bench sandbox".into()); - } - } - _ => {} - } - - let mut result = if let Some(msg) = blocked_reason { - msg - } else { - registry.execute(call, tool_ctx).await - }; - - // Normalize shell failures into `error:` so the model sees them clearly and our metrics count them. - if status == "ok" && name == "shell" { - let cmd = args_json.get("cmd").and_then(|v| v.as_str()).unwrap_or(""); - if let Some(code) = result - .strip_prefix("exit ") - .and_then(|rest| rest.lines().next()) - .and_then(|n| n.trim().parse::().ok()) - { - if code != 0 { - let is_grep_no_match = (cmd.contains("grep") || cmd.contains("rg")) && code == 1; - if !is_grep_no_match { - result = format!("error: {result}"); - } - } - } - } - - let is_error = result.starts_with("error:") || result.starts_with("unknown tool:"); - if status == "ok" && is_error { - status = "error".into(); - } - - ToolsLaneToolCallLog { - name, - args: truncate(&args_str, 4000), - status, - result: truncate(&result, 10_000), - } -} - -async fn run_tools_lane_query( - query: &EvalQuery, - gold: EvalRouteLabel, - router: &RuntimeRouter, - llm: &LlmClient, - tool_ctx: &core_tools::ToolContext, - registry: &core_tools::Subroutines, - tool_defs: &[klbr_core::models::ToolDef], - write_root: &Path, -) -> Result { - let emb = llm - .embed(&query.text) - .await - .with_context(|| format!("embedding failed for tools-lane query {}", query.query_id))?; - let (route, scores) = router.predict_raw_with_scores(&emb); - let router_decision = route_decision_to_label(route); - let router_correct = router_decision == gold; - - let mut tool_iterations = 0usize; - let mut tool_calls_log: Vec = Vec::new(); - let mut final_response: Option = None; - - // Run the tool loop for: - // - gold Tools queries (even if the router predicts Memory), since in the real agent the model - // can still decide to call tools from the normal lane. - // - queries the router predicts Tools (to measure tool-call behavior under router positives). - // - // For gold Memory queries routed to Memory, we skip the tool loop to avoid incentivizing tool - // use on queries that would normally be answered from recalled context in the full system. - if gold == EvalRouteLabel::Tools || route == RuntimeRouteDecision::Tools { - let mut ctx = AgentContext::new(TOOLS_LANE_ANCHOR, &[]); - if query.required_tools.is_empty() { - ctx.push_input(&format!("[source:bench] {}", query.text)); - } else { - ctx.push_input(&format!( - "[source:bench]\n[required_tools: {}]\n{}", - query.required_tools.join(","), - query.text - )); - } - - loop { - let (tok_tx, mut tok_rx) = tokio::sync::mpsc::channel(256); - let llm2 = llm.clone(); - let msgs = ctx.as_messages(); - let defs = tool_defs.to_vec(); - let stream_task = tokio::spawn(async move { llm2.stream(&msgs, &defs, tok_tx).await }); - - let mut response = String::new(); - let mut reasoning = String::new(); - let mut tool_calls: Vec = Vec::new(); - - while let Some(ev) = tok_rx.recv().await { - match ev { - klbr_core::models::LlmEvent::Token(tok) => response.push_str(&tok), - klbr_core::models::LlmEvent::ThinkToken(tok) => reasoning.push_str(&tok), - klbr_core::models::LlmEvent::Usage(_) => {} - klbr_core::models::LlmEvent::ToolCalls(calls) => tool_calls = calls, - klbr_core::models::LlmEvent::ToolCallStarted => {} - } - } - - // Ensure the stream task didn't fail (surface error in trace, don't crash the run). - match stream_task.await { - Ok(Ok(())) => {} - Ok(Err(e)) => { - final_response = Some(format!("error: llm stream failed: {e}")); - break; - } - Err(e) => { - final_response = Some(format!("error: llm stream task failed: {e}")); - break; - } - } - - if !tool_calls.is_empty() && tool_iterations < MAX_TOOLS_LANE_TOOL_ITERATIONS { - tool_iterations += 1; - let text_content = (!response.is_empty()).then_some(response.clone()); - let reasoning_content = (!reasoning.is_empty()).then_some(reasoning.clone()); - ctx.push_assistant_tool_calls(tool_calls.clone(), text_content, reasoning_content); - for call in &tool_calls { - let log = execute_tool_restricted(registry, call, tool_ctx, write_root).await; - ctx.push_tool_result(&call.id, &log.result); - tool_calls_log.push(log); - } - continue; - } - - if !tool_calls.is_empty() && response.trim().is_empty() { - final_response = - Some("error: tool iteration limit reached without a final response".into()); - break; - } - - if response.trim().is_empty() { - if !reasoning.trim().is_empty() { - final_response = Some(reasoning); - } else { - final_response = Some("error: model returned an empty response".into()); - } - break; - } - - final_response = Some(response); - let final_reasoning = (!reasoning.is_empty()).then_some(reasoning.as_str()); - ctx.push_assistant( - final_response.as_deref().unwrap_or_default(), - final_reasoning, - ); - break; - } - } - - let called_any_tool = !tool_calls_log.is_empty(); - let required_tools = query.required_tools.clone(); - let called_any_required_tool = if required_tools.is_empty() { - None - } else { - Some( - tool_calls_log - .iter() - .any(|call| required_tools.iter().any(|t| t == &call.name)), - ) - }; - - Ok(ToolsLaneTrace { - query_id: query.query_id.clone(), - split: query.split.clone(), - text: query.text.clone(), - gold, - required_tools, - router_decision, - router_scores: Some(scores_for_trace(scores)), - router_correct, - tool_iterations, - tool_calls: tool_calls_log, - final_response: truncate(final_response.as_deref().unwrap_or(""), 20_000), - called_any_tool, - called_any_required_tool, - }) -} - -fn summarize_tools_lane( - dataset: &InternalEvalDataset, - experiment: &RetrievalExperimentConfig, - router_model_path: &str, - llm: &LlmClient, - write_root: &Path, - traces: &[ToolsLaneTrace], -) -> ToolsLaneReport { - let mut total = 0usize; - let mut correct = 0usize; - - let mut gold_tools = 0usize; - let mut gold_memory = 0usize; - let mut gold_abstain = 0usize; - - let mut tp_tools = 0usize; - let mut fp_tools_total = 0usize; - let mut fp_mem_to_tools = 0usize; - let mut fp_abs_to_tools = 0usize; - - let mut routed_tools = 0usize; - let mut any_tool_calls = 0usize; - let mut any_tool_call_errors = 0usize; - let mut gold_tools_called_any = 0usize; - let mut gold_memory_called_any = 0usize; - - let mut required_defined = 0usize; - let mut required_hit = 0usize; - - for t in traces { - total += 1; - if t.router_correct { - correct += 1; - } - - match t.gold { - EvalRouteLabel::Tools => gold_tools += 1, - EvalRouteLabel::Memory => gold_memory += 1, - EvalRouteLabel::Abstain => gold_abstain += 1, - } - - if t.router_decision == EvalRouteLabel::Tools { - fp_tools_total += 1; - match t.gold { - EvalRouteLabel::Tools => tp_tools += 1, - EvalRouteLabel::Memory => fp_mem_to_tools += 1, - EvalRouteLabel::Abstain => fp_abs_to_tools += 1, - } - } - if t.router_decision == EvalRouteLabel::Tools { - routed_tools += 1; - } - - if t.called_any_tool { - any_tool_calls += 1; - if t.tool_calls.iter().any(|c| c.status == "error") { - any_tool_call_errors += 1; - } - } - - match t.gold { - EvalRouteLabel::Tools => { - if t.called_any_tool { - gold_tools_called_any += 1; - } - if let Some(hit) = t.called_any_required_tool { - required_defined += 1; - if hit { - required_hit += 1; - } - } - } - EvalRouteLabel::Memory => { - // Only count as a memory-query tool-call if the router predicted Tools (since - // we skip the tool loop for gold Memory routed to Memory). - if t.router_decision == EvalRouteLabel::Tools && t.called_any_tool { - gold_memory_called_any += 1; - } - } - EvalRouteLabel::Abstain => {} - } - } - - let router_accuracy = if total == 0 { - 0.0 - } else { - correct as f32 / total as f32 - }; - let tools_precision = if fp_tools_total == 0 { - 0.0 - } else { - tp_tools as f32 / fp_tools_total as f32 - }; - let tools_recall = if gold_tools == 0 { - 0.0 - } else { - tp_tools as f32 / gold_tools as f32 - }; - let memory_false_tools_rate = if gold_memory == 0 { - 0.0 - } else { - fp_mem_to_tools as f32 / gold_memory as f32 - }; - let abstain_false_tools_rate = if gold_abstain == 0 { - 0.0 - } else { - fp_abs_to_tools as f32 / gold_abstain as f32 - }; - - let tools_lane_gold_tools_called_any_rate = if gold_tools == 0 { - 0.0 - } else { - gold_tools_called_any as f32 / gold_tools as f32 - }; - let tools_lane_gold_memory_called_any_rate = if gold_memory == 0 { - 0.0 - } else { - gold_memory_called_any as f32 / gold_memory as f32 - }; - let tools_lane_required_tool_hit_rate = if required_defined == 0 { - None - } else { - Some(required_hit as f32 / required_defined as f32) - }; - let tools_lane_tool_exec_error_rate = if any_tool_calls == 0 { - 0.0 - } else { - any_tool_call_errors as f32 / any_tool_calls as f32 - }; - - ToolsLaneReport { - dataset_id: dataset.dataset_id.clone(), - experiment_id: experiment.experiment_id.clone(), - evaluated_split: experiment.dataset_split.clone(), - llm_url: llm.config.llm.url.clone(), - embed_url: llm.config.embedder.url.clone(), - embed_model: llm.config.embedder.model.clone(), - router_model_path: router_model_path.to_string(), - write_root: write_root.display().to_string(), - metrics: ToolsLaneMetrics { - query_count: total, - gold_tools, - gold_memory, - gold_abstain, - router_accuracy, - router_tools_recall: tools_recall, - router_tools_precision: tools_precision, - router_memory_false_tools_rate: memory_false_tools_rate, - router_abstain_false_tools_rate: abstain_false_tools_rate, - router_routed_tools_count: routed_tools, - tools_lane_gold_tools_called_any_rate, - tools_lane_gold_memory_called_any_rate, - tools_lane_any_tool_calls_count: any_tool_calls, - tools_lane_required_tool_hit_rate, - tools_lane_tool_exec_error_rate, - }, - } -} -fn render_tools_lane_report(report: &ToolsLaneReport, traces: &[ToolsLaneTrace]) -> String { - let m = &report.metrics; - let mut out = format!( - "# Tools Lane Benchmark Report\n\n\ -dataset: `{}`\n\ -experiment: `{}`\n\ -split: `{:?}`\n\ -llm_url: `{}`\n\ -embed_url: `{}`\n\ -embed_model: `{}`\n\ -router_model: `{}`\n\ -write_root: `{}`\n\n\ -## Router Metrics (Gold = route_label)\n\n\ -- queries: `{}` (tools `{}` / memory `{}` / abstain `{}`)\n\ -- router accuracy: `{:.4}`\n\ -- tools precision: `{:.4}`\n\ -- tools recall: `{:.4}`\n\ -- memory→tools rate: `{:.4}`\n\ -- abstain→tools rate: `{:.4}`\n\n\ -## End-to-End Tool Lane\n\n\ -- router routed-to-tools count: `{}`\n\ -- gold-tools called any tool rate: `{:.4}`\n\ -- gold-memory called any tool rate: `{:.4}`\n\ -- queries that called any tool: `{}`\n\ -- required tool hit rate (if defined): `{}`\n\ -- tool exec error rate (given tool calls): `{:.4}`\n", - report.dataset_id, - report.experiment_id, - report.evaluated_split, - report.llm_url, - report.embed_url, - report.embed_model, - report.router_model_path, - report.write_root, - m.query_count, - m.gold_tools, - m.gold_memory, - m.gold_abstain, - m.router_accuracy, - m.router_tools_precision, - m.router_tools_recall, - m.router_memory_false_tools_rate, - m.router_abstain_false_tools_rate, - m.router_routed_tools_count, - m.tools_lane_gold_tools_called_any_rate, - m.tools_lane_gold_memory_called_any_rate, - m.tools_lane_any_tool_calls_count, - m.tools_lane_required_tool_hit_rate - .map(|v| format!("{v:.4}")) - .unwrap_or_else(|| "n/a".into()), - m.tools_lane_tool_exec_error_rate, - ); - let mut interesting: Vec<&ToolsLaneTrace> = traces - .iter() - .filter(|t| t.gold == EvalRouteLabel::Tools || t.router_decision == EvalRouteLabel::Tools) - .collect(); - interesting.sort_by(|a, b| a.query_id.cmp(&b.query_id)); - if !interesting.is_empty() { - out.push_str("\n## Traces (Tools Gold or Tools Pred)\n\n"); - out.push_str("| id | gold | pred | called_tool | required_hit | tools_score | memory_score | abstain_score | kind |\n"); - out.push_str("|---|---|---|---:|---:|---:|---:|---:|---|\n"); - for t in interesting { - let (st, sm, sa, kind) = t - .router_scores - .as_ref() - .map(|s| (s.tools, s.memory, s.abstain, s.kind.as_str())) - .unwrap_or((0.0, 0.0, 0.0, "unknown")); - let req = t - .called_any_required_tool - .map(|v| if v { "1" } else { "0" }) - .unwrap_or("n/a"); - out.push_str(&format!( - "| `{}` | `{:?}` | `{:?}` | {} | {} | {:.2} | {:.2} | {:.2} | `{}` |\n", - t.query_id, - t.gold, - t.router_decision, - if t.called_any_tool { 1 } else { 0 }, - req, - st, - sm, - sa, - kind - )); - } - } - out -} fn build_embed_only_client(experiment: &RetrievalExperimentConfig) -> LlmClient { let mut runtime = ModelsConfig::default(); @@ -1764,9 +801,8 @@ async fn embed_router_examples( fn label_index(label: &EvalRouteLabel) -> usize { match label { - EvalRouteLabel::Tools => 0, - EvalRouteLabel::Memory => 1, - EvalRouteLabel::Abstain => 2, + EvalRouteLabel::Memory => 0, + EvalRouteLabel::Abstain => 1, } } @@ -1780,19 +816,19 @@ fn train_softmax_router( bail!("router-linear: no training examples for split {:?}", split); } - let mut counts = [0usize; 3]; + let mut counts = [0usize; 2]; for ex in &train { counts[label_index(&ex.gold)] += 1; } if counts[0] == 0 || counts[1] == 0 { bail!( - "router-linear training requires tools+memory labels on split {:?}; got tools={} memory={} abstain={}", - split, counts[0], counts[1], counts[2] + "router-linear training requires memory+abstain labels on split {:?}; got memory={} abstain={}", + split, counts[0], counts[1] ); } - // Softmax regression (3-way). Full-batch gradient descent; dataset is tiny. - let c = 3usize; + // Softmax regression (2-way). Full-batch gradient descent; dataset is tiny. + let c = 2usize; let d = embed_dim; let mut w = vec![vec![0.0f32; d]; c]; let mut b = vec![0.0f32; c]; @@ -1816,8 +852,8 @@ fn train_softmax_router( ); } let y = label_index(&ex.gold); - let logits = logits3(&w, &b, &ex.emb_unit); - let probs = softmax3(logits); + let logits = logits2(&w, &b, &ex.emb_unit); + let probs = softmax2(logits); let py = probs[y].max(1e-9); loss += -py.ln(); @@ -1868,95 +904,66 @@ fn tune_linear_thresholds( let tuning_points = precompute_linear_points(&tuning, weights, bias)?; const MAX_MEMORY_TO_ABSTAIN_RATE: f32 = 0.03; - // For safety, be conservative about routing abstain-lane queries to tools. - const MAX_ABSTAIN_FALSE_TOOL_RATE: f32 = 0.05; - // Keep this sweep cheap but wide enough to find a reasonable operating point. - // (Too-high tools_prob_threshold was a common failure mode: tools queries routed to memory.) let prob_values: Vec = (2..=18).map(|i| i as f32 * 0.05).collect(); // 0.10..0.90 let margin_values: Vec = (0..=25).map(|i| i as f32 * 0.02).collect(); // 0.00..0.50 let mut best: Option<(RuntimeLinearDecisionParams, LinearRouterMetrics)> = None; - for &t_tools_prob in &prob_values { - for &t_tools_margin in &margin_values { - for &t_abs_prob in &prob_values { - for &t_abs_margin in &margin_values { - let decision = RuntimeLinearDecisionParams { - tools_prob_threshold: t_tools_prob, - tools_margin_threshold: t_tools_margin, - abstain_prob_threshold: t_abs_prob, - abstain_margin_threshold: t_abs_margin, - }; - let m = score_linear_thresholds_precomputed(&tuning_points, &decision)?; - - let ok = m.memory_false_tool_rate == 0.0 - && m.memory_false_abstain_rate <= MAX_MEMORY_TO_ABSTAIN_RATE - && m.abstain_false_tool_rate <= MAX_ABSTAIN_FALSE_TOOL_RATE; - - best = match best.take() { - None => Some((decision, m)), - Some((best_d, best_m)) => { - let best_ok = best_m.memory_false_tool_rate == 0.0 - && best_m.memory_false_abstain_rate <= MAX_MEMORY_TO_ABSTAIN_RATE - && best_m.abstain_false_tool_rate <= MAX_ABSTAIN_FALSE_TOOL_RATE; - - let better = if ok && best_ok { - // primary: maximize tools recall (keep tool-lane capability). - // secondary: minimize abstain→tools leakage (avoid spurious tool runs). - // then: maximize overall accuracy and tools precision. - ( - -m.tool_recall, - m.abstain_false_tool_rate, - -m.accuracy, - -m.tool_precision, - ) < ( - -best_m.tool_recall, - best_m.abstain_false_tool_rate, - -best_m.accuracy, - -best_m.tool_precision, - ) - } else if ok && !best_ok { - true - } else if !ok && best_ok { - false - } else { - // fallback: still prefer fewer memory→tools then higher precision - ( - m.memory_false_tool_rate, - -m.tool_precision, - -m.tool_recall, - -m.accuracy, - ) < ( - best_m.memory_false_tool_rate, - -best_m.tool_precision, - -best_m.tool_recall, - -best_m.accuracy, - ) - }; - - if better { - Some((decision, m)) - } else { - Some((best_d, best_m)) - } - } + for &t_abs_prob in &prob_values { + for &t_abs_margin in &margin_values { + let decision = RuntimeLinearDecisionParams { + abstain_prob_threshold: t_abs_prob, + abstain_margin_threshold: t_abs_margin, + }; + let m = score_linear_thresholds_precomputed(&tuning_points, &decision)?; + + let ok = m.memory_false_abstain_rate <= MAX_MEMORY_TO_ABSTAIN_RATE; + + best = match best.take() { + None => Some((decision, m)), + Some((best_d, best_m)) => { + let best_ok = best_m.memory_false_abstain_rate <= MAX_MEMORY_TO_ABSTAIN_RATE; + + let better = if ok && best_ok { + ( + -m.accuracy, + m.memory_false_abstain_rate, + ) < ( + -best_m.accuracy, + best_m.memory_false_abstain_rate, + ) + } else if ok && !best_ok { + true + } else if !ok && best_ok { + false + } else { + ( + m.memory_false_abstain_rate, + -m.accuracy, + ) < ( + best_m.memory_false_abstain_rate, + -best_m.accuracy, + ) }; + + if better { + Some((decision, m)) + } else { + Some((best_d, best_m)) + } } - } + }; } } let (decision, metrics) = best.context("router-linear threshold sweep produced no points")?; eprintln!( - "router-linear: tuned thresholds tools_prob={:.2} tools_margin={:.2} abs_prob={:.2} abs_margin={:.2} (dev: tools_recall={:.3}, tools_precision={:.3}, mem→tools={:.3})", - decision.tools_prob_threshold, - decision.tools_margin_threshold, + "router-linear: tuned thresholds abs_prob={:.2} abs_margin={:.2} (dev: accuracy={:.3}, mem→abstain={:.3})", decision.abstain_prob_threshold, decision.abstain_margin_threshold, - metrics.tool_recall, - metrics.tool_precision, - metrics.memory_false_tool_rate + metrics.accuracy, + metrics.memory_false_abstain_rate ); Ok(decision) } @@ -1966,7 +973,7 @@ struct LinearPoint { query_id: String, text: String, gold: EvalRouteLabel, - probs: [f32; 3], // tools/memory/abstain + probs: [f32; 2], // memory/abstain } fn precompute_linear_points( @@ -1976,7 +983,7 @@ fn precompute_linear_points( ) -> Result> { let mut out = Vec::with_capacity(examples.len()); for ex in examples { - let probs = softmax3(logits3(weights, bias, &ex.emb_unit)); + let probs = softmax2(logits2(weights, bias, &ex.emb_unit)); out.push(LinearPoint { query_id: ex.query_id.clone(), text: ex.text.clone(), @@ -1993,16 +1000,11 @@ fn score_linear_thresholds_precomputed( ) -> Result { let mut total = 0usize; let mut correct = 0usize; - let mut confusion = [[0usize; 3]; 3]; + let mut confusion = [[0usize; 2]; 2]; - let mut tool_tp = 0usize; - let mut tool_fp_total = 0usize; - let mut memory_to_tools_fp = 0usize; - let mut abstain_to_tools_fp = 0usize; - let mut tool_fn = 0usize; let mut memory_to_abstain_fp = 0usize; + let mut abstain_to_memory_fp = 0usize; - let mut tool_count = 0usize; let mut memory_count = 0usize; let mut abstain_count = 0usize; @@ -2026,64 +1028,36 @@ fn score_linear_thresholds_precomputed( text: ex.text.clone(), gold: gold.clone(), pred: pred.clone(), - p_tools: probs[0], - p_memory: probs[1], - p_abstain: probs[2], + p_memory: probs[0], + p_abstain: probs[1], }); } match gold { - EvalRouteLabel::Tools => tool_count += 1, EvalRouteLabel::Memory => memory_count += 1, EvalRouteLabel::Abstain => abstain_count += 1, } match (gold, &pred) { - (EvalRouteLabel::Tools, EvalRouteLabel::Tools) => tool_tp += 1, - (EvalRouteLabel::Memory, EvalRouteLabel::Tools) => { - tool_fp_total += 1; - memory_to_tools_fp += 1; + (EvalRouteLabel::Memory, EvalRouteLabel::Abstain) => { + memory_to_abstain_fp += 1; } - (EvalRouteLabel::Abstain, EvalRouteLabel::Tools) => { - tool_fp_total += 1; - abstain_to_tools_fp += 1; + (EvalRouteLabel::Abstain, EvalRouteLabel::Memory) => { + abstain_to_memory_fp += 1; } - (EvalRouteLabel::Tools, _) => tool_fn += 1, _ => {} } - - if matches!( - (gold, &pred), - (EvalRouteLabel::Memory, EvalRouteLabel::Abstain) - ) { - memory_to_abstain_fp += 1; - } } let denom = total.max(1) as f32; let acc = correct as f32 / denom; - let tool_precision = if tool_tp + tool_fp_total > 0 { - tool_tp as f32 / (tool_tp + tool_fp_total) as f32 - } else { - 0.0 - }; - let tool_recall = if tool_tp + tool_fn > 0 { - tool_tp as f32 / (tool_tp + tool_fn) as f32 - } else { - 0.0 - }; - let memory_false_tool_rate = if memory_count > 0 { - memory_to_tools_fp as f32 / (memory_count as f32) + let memory_false_abstain_rate = if memory_count > 0 { + memory_to_abstain_fp as f32 / (memory_count as f32) } else { 0.0 }; - let abstain_false_tool_rate = if abstain_count > 0 { - abstain_to_tools_fp as f32 / (abstain_count as f32) - } else { - 0.0 - }; - let memory_false_abstain_rate = if memory_count > 0 { - memory_to_abstain_fp as f32 / (memory_count as f32) + let abstain_false_memory_rate = if abstain_count > 0 { + abstain_to_memory_fp as f32 / (abstain_count as f32) } else { 0.0 }; @@ -2092,44 +1066,33 @@ fn score_linear_thresholds_precomputed( total, accuracy: acc, confusion, - tool_precision, - tool_recall, - memory_false_tool_rate, - abstain_false_tool_rate, memory_false_abstain_rate, - tool_count, + abstain_false_memory_rate, memory_count, abstain_count, misclassified, }) } -fn logits3(weights: &[Vec], bias: &[f32], x_unit: &[f32]) -> [f32; 3] { - let mut out = [0.0f32; 3]; - for cls in 0..3 { +fn logits2(weights: &[Vec], bias: &[f32], x_unit: &[f32]) -> [f32; 2] { + let mut out = [0.0f32; 2]; + for cls in 0..2 { out[cls] = dot(x_unit, &weights[cls]) + bias[cls]; } out } -fn softmax3(logits: [f32; 3]) -> [f32; 3] { - let m = logits[0].max(logits[1]).max(logits[2]); +fn softmax2(logits: [f32; 2]) -> [f32; 2] { + let m = logits[0].max(logits[1]); let e0 = (logits[0] - m).exp(); let e1 = (logits[1] - m).exp(); - let e2 = (logits[2] - m).exp(); - let denom = (e0 + e1 + e2).max(1e-12); - [e0 / denom, e1 / denom, e2 / denom] + let denom = (e0 + e1).max(1e-12); + [e0 / denom, e1 / denom] } -fn predict_linear_route(probs: [f32; 3], decision: &RuntimeLinearDecisionParams) -> EvalRouteLabel { - let p_tools = probs[0]; - let p_memory = probs[1]; - let p_abstain = probs[2]; - - let tools_margin = p_tools - p_memory.max(p_abstain); - if p_tools >= decision.tools_prob_threshold && tools_margin >= decision.tools_margin_threshold { - return EvalRouteLabel::Tools; - } +fn predict_linear_route(probs: [f32; 2], decision: &RuntimeLinearDecisionParams) -> EvalRouteLabel { + let p_memory = probs[0]; + let p_abstain = probs[1]; let abstain_margin = p_abstain - p_memory; if p_abstain >= decision.abstain_prob_threshold @@ -2208,70 +1171,56 @@ embed_model: `{}`\n\ trained_on: `{:?}`\n\ evaluated_on: `{:?}`\n\n\ ## Decision Params\n\n\ -- tools_prob_threshold: `{:.3}`\n\ -- tools_margin_threshold: `{:.3}`\n\ - abstain_prob_threshold: `{:.3}`\n\ - abstain_margin_threshold: `{:.3}`\n\n\ ## Metrics\n\n\ - labeled queries evaluated: `{}`\n\ -- tools: `{}` memory: `{}` abstain: `{}`\n\ +- memory: `{}` abstain: `{}`\n\ - accuracy: `{:.4}`\n\ -- tools precision: `{:.4}`\n\ -- tools recall: `{:.4}`\n\ -- memory false-tools rate (memory→tools): `{:.4}`\n\ -- abstain false-tools rate (abstain→tools): `{:.4}`\n\ -- memory false-abstain rate (memory→abstain): `{:.4}`\n", +- memory false-abstain rate (memory→abstain): `{:.4}`\n\ +- abstain false-memory rate (abstain→memory): `{:.4}`\n", dataset.dataset_id, report.embed_model, report.trained_on, report.evaluated_on, - d.tools_prob_threshold, - d.tools_margin_threshold, d.abstain_prob_threshold, d.abstain_margin_threshold, m.total, - m.tool_count, m.memory_count, m.abstain_count, m.accuracy, - m.tool_precision, - m.tool_recall, - m.memory_false_tool_rate, - m.abstain_false_tool_rate, m.memory_false_abstain_rate, + m.abstain_false_memory_rate, ); out.push_str("\n## Confusion Matrix (gold x pred)\n\n"); - out.push_str("| gold \\ pred | tools | memory | abstain |\n"); - out.push_str("|---|---:|---:|---:|\n"); + out.push_str("| gold \\ pred | memory | abstain |\n"); + out.push_str("|---|---:|---:|\n"); let rows = [ - EvalRouteLabel::Tools, EvalRouteLabel::Memory, EvalRouteLabel::Abstain, ]; - let names = ["tools", "memory", "abstain"]; + let names = ["memory", "abstain"]; for (i, gold) in rows.iter().enumerate() { out.push_str(&format!( - "| {} | {} | {} | {} |\n", + "| {} | {} | {} |\n", names[label_index(gold)], m.confusion[i][0], - m.confusion[i][1], - m.confusion[i][2] + m.confusion[i][1] )); } if !m.misclassified.is_empty() { out.push_str("\n## Misclassified Queries\n\n"); - out.push_str("| id | text | gold | pred | p_tools | p_memory | p_abstain |\n"); - out.push_str("|---|---|---|---|---:|---:|---:|\n"); + out.push_str("| id | text | gold | pred | p_memory | p_abstain |\n"); + out.push_str("|---|---|---|---|---:|---:|\n"); for trace in &m.misclassified { out.push_str(&format!( - "| `{}` | {} | `{:?}` | `{:?}` | {:.2} | {:.2} | {:.2} |\n", + "| `{}` | {} | `{:?}` | `{:?}` | {:.2} | {:.2} |\n", trace.query_id, trace.text, trace.gold, trace.pred, - trace.p_tools, trace.p_memory, trace.p_abstain )); @@ -2282,33 +1231,27 @@ evaluated_on: `{:?}`\n\n\ out.push_str("\n## Holdout Metrics (query_id starts with `holdout_`)\n\n"); out.push_str(&format!( "- labeled queries evaluated: `{}`\n\ -- tools precision: `{:.4}`\n\ -- tools recall: `{:.4}`\n\ -- memory false-tools rate (memory→tools): `{:.4}`\n\ -- abstain false-tools rate (abstain→tools): `{:.4}`\n\ +- memory false-abstain rate (memory→abstain): `{:.4}`\n\ +- abstain false-memory rate (abstain→memory): `{:.4}`\n\ - accuracy: `{:.4}`\n", h.total, - h.tool_precision, - h.tool_recall, - h.memory_false_tool_rate, - h.abstain_false_tool_rate, + h.memory_false_abstain_rate, + h.abstain_false_memory_rate, h.accuracy, )); } if !report.holdout_metrics_by_bucket.is_empty() { out.push_str("\n## Holdout Metrics By Bucket\n\n"); - out.push_str("| bucket | total | tools precision | tools recall | memory→tools | abstain→tools | accuracy |\n"); - out.push_str("|---|---:|---:|---:|---:|---:|---:|\n"); + out.push_str("| bucket | total | memory→abstain | abstain→memory | accuracy |\n"); + out.push_str("|---|---:|---:|---:|---:|\n"); for (bucket, h) in &report.holdout_metrics_by_bucket { out.push_str(&format!( - "| `{}` | {} | {:.4} | {:.4} | {:.4} | {:.4} | {:.4} |\n", + "| `{}` | {} | {:.4} | {:.4} | {:.4} |\n", bucket, h.total, - h.tool_precision, - h.tool_recall, - h.memory_false_tool_rate, - h.abstain_false_tool_rate, + h.memory_false_abstain_rate, + h.abstain_false_memory_rate, h.accuracy )); } @@ -2352,7 +1295,6 @@ async fn train_router_centroids( split: DatasetSplit, ) -> Result<(RouterModel, f32)> { let mut memory_embs: Vec> = Vec::new(); - let mut tools_embs: Vec> = Vec::new(); let mut abstain_embs: Vec> = Vec::new(); for query in dataset @@ -2372,17 +1314,14 @@ async fn train_router_centroids( })?; match label { EvalRouteLabel::Memory => memory_embs.push(emb), - EvalRouteLabel::Tools => tools_embs.push(emb), EvalRouteLabel::Abstain => abstain_embs.push(emb), } } - if memory_embs.is_empty() || tools_embs.is_empty() { + if memory_embs.is_empty() { bail!( - "router training requires both memory and tools labels on split {:?}; got memory={} tools={}", - split, - memory_embs.len(), - tools_embs.len() + "router training requires memory labels on split {:?}", + split ); } let embed_dim = memory_embs[0].len(); @@ -2392,7 +1331,6 @@ async fn train_router_centroids( experiment.embed_dim, embed_dim ); } - // abstain is optional for training, but better if present if abstain_embs.is_empty() { eprintln!( "warning: no abstain examples for training, abstain centroid will be zero (dim={})", @@ -2401,26 +1339,16 @@ async fn train_router_centroids( abstain_embs.push(vec![0.0; embed_dim]); } - // Multiple prototypes per class. This usually improves tools recall because the "tools lane" - // is a union of different intents (repo introspection, runtime state, file IO, etc). let memory_centroids = kmeans_prototypes_unit(&memory_embs, choose_k(memory_embs.len(), 3), 8); - let tools_centroids = kmeans_prototypes_unit(&tools_embs, choose_k(tools_embs.len(), 4), 8); let abstain_centroids = kmeans_prototypes_unit(&abstain_embs, choose_k(abstain_embs.len(), 2), 8); - // Tune a conservative threshold on the training split: prioritize high precision for tools. let mut best = None::<(f32, RouterMetrics)>; - // If we insist on *zero* memory→tools, tools recall tends to collapse because many "tools-needed" - // queries are semantically close to "memory questions about the same topic" (config/reranker/etc). - // Allowing a tiny amount of memory→tools (usually ~1 query on dev) can substantially improve recall. - const MAX_MEMORY_FALSE_TOOL_RATE: f32 = 0.03; - const MIN_TOOL_PRECISION: f32 = 0.98; for threshold in (-40..=40).map(|i| i as f32 * 0.05) { let metrics = evaluate_router_on_split( dataset, llm, &memory_centroids, - &tools_centroids, &abstain_centroids, threshold, split.clone(), @@ -2429,45 +1357,13 @@ async fn train_router_centroids( best = match best.take() { None => Some((threshold, metrics)), Some((best_t, best_m)) => { - // Constrained objective: - // - keep memory→tools mistakes near-zero - // - keep tool precision high - // - then maximize tool recall - // - // If constraints can't be met, fall back to the old conservative lexicographic compare. - let ok = metrics.memory_false_tool_rate <= MAX_MEMORY_FALSE_TOOL_RATE - && metrics.tool_precision >= MIN_TOOL_PRECISION; - let best_ok = best_m.memory_false_tool_rate <= MAX_MEMORY_FALSE_TOOL_RATE - && best_m.tool_precision >= MIN_TOOL_PRECISION; - - let better = if ok && best_ok { - ( - -metrics.tool_recall, - -metrics.tool_precision, - -metrics.accuracy, - ) < ( - -best_m.tool_recall, - -best_m.tool_precision, - -best_m.accuracy, - ) - } else if ok && !best_ok { - true - } else if !ok && best_ok { - false - } else { - // Conservative fallback. - ( - metrics.memory_false_tool_rate, - -metrics.tool_precision, - -metrics.tool_recall, - -metrics.accuracy, - ) < ( - best_m.memory_false_tool_rate, - -best_m.tool_precision, - -best_m.tool_recall, - -best_m.accuracy, - ) - }; + let better = ( + -metrics.accuracy, + metrics.memory_false_abstain_rate, + ) < ( + -best_m.accuracy, + best_m.memory_false_abstain_rate, + ); if better { Some((threshold, metrics)) } else { @@ -2487,7 +1383,6 @@ async fn train_router_centroids( trained_on: split.clone(), threshold, memory_centroids, - tools_centroids, abstain_centroids, }, threshold, @@ -2505,7 +1400,6 @@ async fn evaluate_router( dataset, llm, &model.memory_centroids, - &model.tools_centroids, &model.abstain_centroids, model.threshold, split, @@ -2525,19 +1419,14 @@ async fn evaluate_router_on_split( dataset: &InternalEvalDataset, llm: &LlmClient, memory_centroids: &[Vec], - tools_centroids: &[Vec], abstain_centroids: &[Vec], threshold: f32, split: DatasetSplit, ) -> Result { let mut total = 0usize; let mut correct = 0usize; - let mut tool_tp = 0usize; - let mut tool_fp_total = 0usize; - let mut memory_to_tools_fp = 0usize; - let mut abstain_to_tools_fp = 0usize; - let mut tool_fn = 0usize; - let mut tool_count = 0usize; + let mut memory_to_abstain_fp = 0usize; + let mut abstain_to_memory_fp = 0usize; let mut memory_count = 0usize; let mut abstain_count = 0usize; let mut misclassified = Vec::new(); @@ -2557,14 +1446,12 @@ async fn evaluate_router_on_split( format!("embedding failed for router eval query {}", query.query_id) })?; let emb_unit = normalize_vec(&emb); - let sim_tools = max_sim(&emb_unit, tools_centroids); let sim_memory = max_sim(&emb_unit, memory_centroids); let sim_abstain = max_sim(&emb_unit, abstain_centroids); let pred = predict_route( &emb_unit, memory_centroids, - tools_centroids, abstain_centroids, threshold, ); @@ -2577,52 +1464,36 @@ async fn evaluate_router_on_split( text: query.text.clone(), gold: gold.clone(), pred: pred.clone(), - sim_tools, sim_memory, sim_abstain, }); } match gold { - EvalRouteLabel::Tools => tool_count += 1, EvalRouteLabel::Memory => memory_count += 1, EvalRouteLabel::Abstain => abstain_count += 1, } match (gold, pred) { - (EvalRouteLabel::Tools, EvalRouteLabel::Tools) => tool_tp += 1, - (EvalRouteLabel::Memory, EvalRouteLabel::Tools) => { - tool_fp_total += 1; - memory_to_tools_fp += 1; + (EvalRouteLabel::Memory, EvalRouteLabel::Abstain) => { + memory_to_abstain_fp += 1; } - (EvalRouteLabel::Abstain, EvalRouteLabel::Tools) => { - tool_fp_total += 1; - abstain_to_tools_fp += 1; + (EvalRouteLabel::Abstain, EvalRouteLabel::Memory) => { + abstain_to_memory_fp += 1; } - (EvalRouteLabel::Tools, _) => tool_fn += 1, _ => {} } } let denom = total.max(1) as f32; let acc = correct as f32 / denom; - let tool_precision = if tool_tp + tool_fp_total > 0 { - tool_tp as f32 / (tool_tp + tool_fp_total) as f32 - } else { - 0.0 - }; - let tool_recall = if tool_tp + tool_fn > 0 { - tool_tp as f32 / (tool_tp + tool_fn) as f32 - } else { - 0.0 - }; - let memory_false_tool_rate = if memory_count > 0 { - memory_to_tools_fp as f32 / (memory_count as f32) + let memory_false_abstain_rate = if memory_count > 0 { + memory_to_abstain_fp as f32 / (memory_count as f32) } else { 0.0 }; - let abstain_false_tool_rate = if abstain_count > 0 { - abstain_to_tools_fp as f32 / (abstain_count as f32) + let abstain_false_memory_rate = if abstain_count > 0 { + abstain_to_memory_fp as f32 / (abstain_count as f32) } else { 0.0 }; @@ -2630,11 +1501,8 @@ async fn evaluate_router_on_split( Ok(RouterMetrics { total, accuracy: acc, - tool_precision, - tool_recall, - memory_false_tool_rate, - abstain_false_tool_rate, - tool_count, + memory_false_abstain_rate, + abstain_false_memory_rate, memory_count, abstain_count, misclassified, @@ -2644,20 +1512,16 @@ async fn evaluate_router_on_split( fn predict_route( query_embedding_unit: &[f32], memory_centroids_unit: &[Vec], - tools_centroids_unit: &[Vec], abstain_centroids_unit: &[Vec], threshold: f32, ) -> EvalRouteLabel { - let sim_tools = max_sim(query_embedding_unit, tools_centroids_unit); let sim_memory = max_sim(query_embedding_unit, memory_centroids_unit); let sim_abstain = max_sim(query_embedding_unit, abstain_centroids_unit); - if (sim_tools - sim_memory) >= threshold && (sim_tools - sim_abstain) >= threshold { - EvalRouteLabel::Tools - } else if sim_memory >= sim_abstain { - EvalRouteLabel::Memory - } else { + if (sim_abstain - sim_memory) >= threshold { EvalRouteLabel::Abstain + } else { + EvalRouteLabel::Memory } } @@ -2790,45 +1654,37 @@ evaluated_on: `{:?}`\n\ threshold: `{:.3}`\n\n\ ## Metrics\n\n\ - labeled queries evaluated: `{}`\n\ -- tools: `{}` memory: `{}` abstain: `{}`\n\ +- memory: `{}` abstain: `{}`\n\ - accuracy: `{:.4}`\n\ -- tools precision: `{:.4}`\n\ -- tools recall: `{:.4}`\n\ -- memory false-tools rate: `{:.4}`\n", +- memory false-abstain rate: `{:.4}`\n", dataset.dataset_id, report.embed_model, report.trained_on, report.evaluated_on, report.threshold, m.total, - m.tool_count, m.memory_count, m.abstain_count, m.accuracy, - m.tool_precision, - m.tool_recall, - m.memory_false_tool_rate, + m.memory_false_abstain_rate, ); - // Include this in the report so we can see when the threshold is mostly protecting against - // abstain→tools mistakes vs actually skipping memory incorrectly. out.push_str(&format!( - "- abstain false-tools rate: `{:.4}`\n", - m.abstain_false_tool_rate + "- abstain false-memory rate: `{:.4}`\n", + m.abstain_false_memory_rate )); if !m.misclassified.is_empty() { out.push_str("\n## Misclassified Queries\n\n"); - out.push_str("| id | text | gold | pred | T-sim | M-sim | A-sim |\n"); - out.push_str("|---|---|---|---|---|---|---|\n"); + out.push_str("| id | text | gold | pred | M-sim | A-sim |\n"); + out.push_str("|---|---|---|---|---|---|\n"); for trace in &m.misclassified { out.push_str(&format!( - "| `{}` | {} | `{:?}` | `{:?}` | {:.2} | {:.2} | {:.2} |\n", + "| `{}` | {} | `{:?}` | `{:?}` | {:.2} | {:.2} |\n", trace.query_id, trace.text, trace.gold, trace.pred, - trace.sim_tools, trace.sim_memory, trace.sim_abstain )); @@ -3053,10 +1909,10 @@ async fn execute_queries( Some(split) => &query.split == split, None => true, }) - // Tool-lane queries are owned by the router/tooling path, not the memory retrieval path. + // Abstain-lane queries (including legacy tools queries) are owned by the router/tooling path, not the memory retrieval path. // Keeping them in the retrieval benchmark creates misleading "no-hit false answers" // because retrieval has no way to inspect runtime state or read repo files. - .filter(|query| query.route_label.as_ref() != Some(&EvalRouteLabel::Tools)) + .filter(|query| query.route_label.as_ref() != Some(&EvalRouteLabel::Abstain)) .collect::>(); if queries.is_empty() { diff --git a/klbr-core/src/agent.rs b/klbr-core/src/agent.rs index be033ab..7da3228 100644 --- a/klbr-core/src/agent.rs +++ b/klbr-core/src/agent.rs @@ -550,18 +550,11 @@ impl Agent { }; let scores_str = scores .map(|s| { - format!(" (T:{:.2} M:{:.2} A:{:.2})", s.tools, s.memory, s.abstain) + format!(" (M:{:.2} A:{:.2})", s.memory, s.abstain) }) .unwrap_or_default(); match route { - RouteDecision::Tools => { - let _ = self.output.send(AgentEvent::Status(format!( - "routed: tool lane{}", - scores_str - ))); - vec![] - } RouteDecision::Abstain => { let _ = self.output.send(AgentEvent::Status(format!( "routed: abstain lane{}", diff --git a/klbr-core/src/mvp.rs b/klbr-core/src/mvp.rs index eaffd5a..de6433f 100644 --- a/klbr-core/src/mvp.rs +++ b/klbr-core/src/mvp.rs @@ -142,9 +142,8 @@ pub enum ExpectedMemoryAction { pub enum EvalRouteLabel { /// The answer should come from the memory corpus (retrieve + rerank). Memory, - /// The answer depends on runtime state or external inspection and should come from tools. - Tools, /// Neither memory nor available tools can reliably answer; abstain or ask a clarification. + #[serde(alias = "tools")] Abstain, } diff --git a/klbr-core/src/router.rs b/klbr-core/src/router.rs index abb3179..36b5d18 100644 --- a/klbr-core/src/router.rs +++ b/klbr-core/src/router.rs @@ -20,17 +20,12 @@ use serde::{Deserialize, Serialize}; pub enum RouteDecision { /// Answer from memory via the normal recall path. Memory, - /// Runtime state inspection needed — skip memory injection, signal tool lane. - Tools, /// Query is underspecified or out of scope — skip recall, signal abstention. Abstain, } #[derive(Debug, Clone, Copy)] pub struct RouterScores { - /// For centroid routers: max cosine similarity to the class prototypes. - /// For linear routers: softmax probability. - pub tools: f32, pub memory: f32, pub abstain: f32, pub kind: &'static str, @@ -85,22 +80,18 @@ pub struct CentroidRouterModel { pub embed_model: String, pub embed_dim: usize, pub trained_on: String, - pub threshold: f32, // margin for Tools vs others + pub threshold: f32, // margin for Abstain vs Memory // v2+: multiple prototypes per class. #[serde(default)] pub memory_centroids: Vec>, #[serde(default)] - pub tools_centroids: Vec>, - #[serde(default)] pub abstain_centroids: Vec>, // v1 compat (single centroid) #[serde(default)] pub memory_centroid: Vec, #[serde(default)] - pub tools_centroid: Vec, - #[serde(default)] pub abstain_centroid: Vec, } @@ -109,7 +100,6 @@ pub struct CentroidRouter { threshold: f32, embed_dim: usize, memory_centroids_unit: Vec>, - tools_centroids_unit: Vec>, abstain_centroids_unit: Vec>, } @@ -121,10 +111,6 @@ impl CentroidRouter { if mem.is_empty() && !model.memory_centroid.is_empty() { mem.push(model.memory_centroid); } - let mut tools = model.tools_centroids; - if tools.is_empty() && !model.tools_centroid.is_empty() { - tools.push(model.tools_centroid); - } let mut abstain = model.abstain_centroids; if abstain.is_empty() && !model.abstain_centroid.is_empty() { abstain.push(model.abstain_centroid); @@ -134,7 +120,6 @@ impl CentroidRouter { threshold: model.threshold, embed_dim: model.embed_dim, memory_centroids_unit: mem.into_iter().map(|v| normalize(&v)).collect(), - tools_centroids_unit: tools.into_iter().map(|v| normalize(&v)).collect(), abstain_centroids_unit: abstain.into_iter().map(|v| normalize(&v)).collect(), }) } @@ -144,7 +129,6 @@ impl CentroidRouter { return ( RouteDecision::Memory, RouterScores { - tools: 0.0, memory: 0.0, abstain: 0.0, kind: "centroid_v2(mismatch)", @@ -153,24 +137,18 @@ impl CentroidRouter { } let unit = normalize(query_embedding); - let sim_tools = max_sim(&unit, &self.tools_centroids_unit); let sim_memory = max_sim(&unit, &self.memory_centroids_unit); let sim_abstain = max_sim(&unit, &self.abstain_centroids_unit); - let decision = if (sim_tools - sim_memory) >= self.threshold - && (sim_tools - sim_abstain) >= self.threshold - { - RouteDecision::Tools - } else if sim_memory >= sim_abstain { - RouteDecision::Memory - } else { + let decision = if (sim_abstain - sim_memory) >= self.threshold { RouteDecision::Abstain + } else { + RouteDecision::Memory }; ( decision, RouterScores { - tools: sim_tools, memory: sim_memory, abstain: sim_abstain, kind: "centroid_v2(sim)", @@ -192,11 +170,6 @@ fn validate_centroid_model(model: &CentroidRouterModel) -> Result<()> { } else { vec![model.memory_centroid.len()] }; - let tools_lens: Vec = if !model.tools_centroids.is_empty() { - model.tools_centroids.iter().map(|v| v.len()).collect() - } else { - vec![model.tools_centroid.len()] - }; let abstain_lens: Vec = if !model.abstain_centroids.is_empty() { model.abstain_centroids.iter().map(|v| v.len()).collect() } else { @@ -205,7 +178,6 @@ fn validate_centroid_model(model: &CentroidRouterModel) -> Result<()> { for (name, lens) in [ ("memory", mem_lens), - ("tools", tools_lens), ("abstain", abstain_lens), ] { if lens.is_empty() { @@ -230,8 +202,6 @@ fn validate_centroid_model(model: &CentroidRouterModel) -> Result<()> { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LinearDecisionParams { - pub tools_prob_threshold: f32, - pub tools_margin_threshold: f32, pub abstain_prob_threshold: f32, pub abstain_margin_threshold: f32, } @@ -246,7 +216,7 @@ pub struct LinearRouterModel { #[serde(default)] pub dataset_id: Option, #[serde(default)] - pub labels: Vec, // must contain tools/memory/abstain + pub labels: Vec, // must contain memory/abstain pub weights: Vec>, // [C][D] pub bias: Vec, // [C] pub decision: LinearDecisionParams, @@ -255,9 +225,9 @@ pub struct LinearRouterModel { #[derive(Debug, Clone)] pub struct LinearRouter { embed_dim: usize, - // Stored in fixed order: Tools, Memory, Abstain. - w: [Vec; 3], - b: [f32; 3], + // Stored in fixed order: Memory, Abstain. + w: [Vec; 2], + b: [f32; 2], decision: LinearDecisionParams, } @@ -272,9 +242,9 @@ impl LinearRouter { if model.embed_dim == 0 { bail!("embed_dim must be > 0"); } - if model.weights.len() != 3 || model.bias.len() != 3 { + if model.weights.len() != 2 || model.bias.len() != 2 { bail!( - "linear router must have 3 classes; got weights={} bias={}", + "linear router must have 2 classes; got weights={} bias={}", model.weights.len(), model.bias.len() ); @@ -288,44 +258,38 @@ impl LinearRouter { ); } } - if !model.decision.tools_prob_threshold.is_finite() - || !model.decision.tools_margin_threshold.is_finite() - || !model.decision.abstain_prob_threshold.is_finite() + if !model.decision.abstain_prob_threshold.is_finite() || !model.decision.abstain_margin_threshold.is_finite() { bail!("decision thresholds must be finite"); } // Map label order from file into our fixed order. - let mut idx_tools = None; let mut idx_memory = None; let mut idx_abstain = None; if !model.labels.is_empty() { for (i, lab) in model.labels.iter().enumerate() { match lab.as_str() { - "tools" => idx_tools = Some(i), "memory" => idx_memory = Some(i), "abstain" => idx_abstain = Some(i), _ => {} } } } - let (it, im, ia) = match (idx_tools, idx_memory, idx_abstain) { - (Some(it), Some(im), Some(ia)) => (it, im, ia), - _ => (0, 1, 2), // assume canonical order if labels missing + let (im, ia) = match (idx_memory, idx_abstain) { + (Some(im), Some(ia)) => (im, ia), + _ => (0, 1), // assume canonical order if labels missing }; - let w_tools = model.weights[it].clone(); let w_memory = model.weights[im].clone(); let w_abstain = model.weights[ia].clone(); - let b_tools = model.bias[it]; let b_memory = model.bias[im]; let b_abstain = model.bias[ia]; Ok(Self { embed_dim: model.embed_dim, - w: [w_tools, w_memory, w_abstain], - b: [b_tools, b_memory, b_abstain], + w: [w_memory, w_abstain], + b: [b_memory, b_abstain], decision: model.decision, }) } @@ -335,7 +299,6 @@ impl LinearRouter { return ( RouteDecision::Memory, RouterScores { - tools: 0.0, memory: 0.0, abstain: 0.0, kind: "linear_softmax_v1(mismatch)", @@ -344,21 +307,13 @@ impl LinearRouter { } let unit = normalize(query_embedding); let logits = self.logits(&unit); - let probs = softmax3(logits); - let p_tools = probs[0]; - let p_memory = probs[1]; - let p_abstain = probs[2]; - - let best_other_for_tools = p_memory.max(p_abstain); - let tools_margin = p_tools - best_other_for_tools; + let probs = softmax2(logits); + let p_memory = probs[0]; + let p_abstain = probs[1]; let abstain_margin = p_abstain - p_memory; - let decision = if p_tools >= self.decision.tools_prob_threshold - && tools_margin >= self.decision.tools_margin_threshold - { - RouteDecision::Tools - } else if p_abstain >= self.decision.abstain_prob_threshold + let decision = if p_abstain >= self.decision.abstain_prob_threshold && abstain_margin >= self.decision.abstain_margin_threshold { RouteDecision::Abstain @@ -369,7 +324,6 @@ impl LinearRouter { ( decision, RouterScores { - tools: p_tools, memory: p_memory, abstain: p_abstain, kind: "linear_softmax_v1(prob)", @@ -377,22 +331,21 @@ impl LinearRouter { ) } - fn logits(&self, x_unit: &[f32]) -> [f32; 3] { - let mut out = [0.0f32; 3]; - for c in 0..3 { + fn logits(&self, x_unit: &[f32]) -> [f32; 2] { + let mut out = [0.0f32; 2]; + for c in 0..2 { out[c] = dot(x_unit, &self.w[c]) + self.b[c]; } out } } -fn softmax3(logits: [f32; 3]) -> [f32; 3] { - let m = logits[0].max(logits[1]).max(logits[2]); +fn softmax2(logits: [f32; 2]) -> [f32; 2] { + let m = logits[0].max(logits[1]); let e0 = (logits[0] - m).exp(); let e1 = (logits[1] - m).exp(); - let e2 = (logits[2] - m).exp(); - let denom = (e0 + e1 + e2).max(1e-12); - [e0 / denom, e1 / denom, e2 / denom] + let denom = (e0 + e1).max(1e-12); + [e0 / denom, e1 / denom] } // --------------------