MLToolBench: Learning Tool-Augmented Agents for Machine Learning Development
Abstract
Machine learning engineering (MLE) agents have made substantial progress, but learning through ML experimentation remains costly in time and computation. Synthetic environments reduce these costs while introducing variations in data and experimental settings that require task-specific diagnosis. Access to diagnostic tools alone does not ensure that agents learn when to use them or how to act on their findings. We introduce ToolMLBench, a suite of executable tools for data inspection, code verification, and experiment diagnosis, together with an SFT and RL pipeline for learning their use. Diagnostic calls acquire evidence whose value depends on subsequent decisions, so final outcomes provide limited guidance on which calls to reinforce. We address this challenge with SPICE, which measures how privileged context changes the likelihood of a sampled tool action and uses this difference as a turn-level reward alongside the final outcome. We train on 80 synthetic tasks and evaluate on 25 in-domain and 10 out-of-domain tasks. Providing tool interfaces and descriptions alone yields inconsistent gains across unadapted models. With the same diagnostic interface, our training pipeline raises in-domain success from 24.8% to 52.4% for Qwen3-8B and from 35.6% to 69.2% for Qwen3.5-35B-A3B. The latter also improves from 31% to 48% out-of-domain, supporting learned diagnostic tool use on held-out sources and targets.
Figures & tables
| In-domain | Out-of-domain | ||||
| Backbone | Model / training | Original | + MLToolBench | Original | + MLToolBench |
| Closed-Source Models | |||||
| Gemini 3.6 Flash | 65.20% | 61.20% | 52.00% | 46.00% | |
| Claude Sonnet 5 | 64.80% | 62.40% | 45.00% | 43.00% | |
| Open-Source Models | |||||
| Qwen2.5-Coder-14B | 45.20% | 40.40% | 34.00% | 32.00% | |
Appendix figures & tables15 assets
Supplementary material from the paper’s appendix.
Appendix
| Partition | Instances | Runs per task | Total / purpose |
|---|---|---|---|
| Training | 80 | 100 | 8,000 raw demonstrations per SFT condition |
| ID evaluation | 25 | 10 | 250 evaluation runs per condition |
| OOD evaluation | 10 | 10 | 100 evaluation runs per condition |
| RL configuration | Value |
|---|---|
| Optimizer updates / reported checkpoint | 100 / final step 100 |
| Rollout batch | 8 tasks 8 rollouts = 64 trajectories |
| Optimizer / learning rate / scheduler | AdamW / / cosine |
| Minibatch size / update epochs / accumulation | 128 / 1 / 1 |
| Advantage estimation | Task-group mean centering; no standard-deviation scaling |
| Clip parameter / KL coefficient | 0.2 / 0 (KL disabled) |
| 35B training condition | ID success (%) | OOD success (%) |
|---|---|---|
| Outcome-only RL | 48.00 (120/250) | 44.00 (44/100) |
| SPICE, full solution summaries | 69.20 (173/250) | 48.00 (48/100) |
| SPICE, reduced solution summaries | 67.20 (168/250) | 41.00 (41/100) |
| SPICE, shuffled scores | 59.20 (148/250) | 45.00 (45/100) |
| Coverage (out of 5) | Mean per rollout | |||
|---|---|---|---|---|
| Model / training | Data management | Verification | Calls | Distinct tools |
| 8B Full SFT | 2.5 | 3.4 | 5 | 2 |
| 8B RL + Spice | 2.5 | 3.3 | 11 | 6 |
| 35B Full SFT | 2.6 | 3.9 | 7 | 4 |
| 35B Outcome-only RL | 0.8 | 1.3 | 9 | 2 |
| 35B RL + TIPS | 1.4 | 1.6 | 10 | 3 |
| 35B training objective | Observation | ID success (%) |
|---|---|---|
| Outcome-only RL | True | 48.00 |
| Masked | 46.40 | |
| Paraphrased | 47.60 | |
| RL + Spice | True | 69.20 |
| Masked | 54.40 | |
| Paraphrased | 68.40 |
| Tool | Returned evidence / behavior |
|---|---|
| Discover Log Metrics | Metric names, sample values, and occurrence counts in logs or captured stdout. |
| Extract Training Metric Series | Value–step series for requested logged metrics. |
| Verify Hyperparam Effective | Static parameter-definition sites, conflicting values, and environment overrides. |
| Estimate Peak Memory | Observed peak GPU memory during a short training-script probe; OOM status. |
| Trace Tensor Shapes | Per-layer tensor shapes or failure traces from a CPU dummy forward pass. |
| Check Config Consistency | Rule-based checks of configuration relations and incompatible settings. |
| Tool | Returned evidence / behavior |
|---|---|
| Tabular (12 tools) | |
| Detect Missing Values | Per-column missingness and suggested imputation strategies. |
| Identify Feature Types | Inferred column types, unique counts, and ID/high-cardinality flags. |
| Detect Label Imbalance | Class counts, imbalance ratios, and resampling suggestions. |
| Check Feature Normalization | Numeric ranges, means, standard deviations, and scaling suggestions. |
| Detect Outliers | Per-column outlier counts and IQR bounds. |
| Task | Task-specific diagnostic guidance |
|---|---|
| ogbn-arxiv | Context. Improve paper-topic node classification from an MLP, following the diagnostic example in Figure 2 . The presence of graph edges alone does not establish that replacing the MLP with GraphSAGE will improve performance. Inspect. Check node-feature ranges, per-column scales, and the active preprocessing pipeline. Compare training and validation curves, then assess whether graph connectivity and label homophily on permitted training labels motivate neighborhood aggregation. Act on evidence. If features have disproportionate scales or missing normalization, fit the transformation on permitted training data and re-run the MLP first. Compare the corrected MLP with the original baseline under the same split, evaluator, and budget. Only then test GraphSAGE if graph diagnostics motivate it; retain the same feature preprocessing and assess the additional gain from neighborhood information. Do not use hidden labels for diagnosis or treat a more complex architecture as an automatic improvement. |
| cifar10 | Context. Improve ten-class image classification from a small CNN trained for five epochs; configured instances vary input scaling, training class balance, or cropping strength. Inspect. Check pixel ranges and channel statistics after preprocessing, inspect transformed images for retained object content, and compare training class counts with per-class validation performance. Read training and validation curves to distinguish optimization problems from overfitting. Act on evidence. If inputs retain an unintended integer scale, correct scaling and use training-derived normalization. If crops remove most object content, reduce cropping strength while preserving the expected input shape. If minority classes are underrepresented and perform poorly, test balanced sampling or class-weighted loss. If these checks are healthy, retain preprocessing and investigate optimization. Re-run under the same split and budget, comparing validation accuracy with the starter before accepting an edit. |
| SFT configuration | Value |
|---|---|
| Adaptation | Full parameters |
| Epochs / optimizer | 3 / AdamW |
| Learning rate / scheduler / warmup | / cosine / 0 |
| Effective batch size | 64, using gradient accumulation |
| Maximum sequence length | 10,240 tokens |
| Hardware / precision | 8 H200 / BF16 |
| Condition | Agent-visible prompt | Training signal |
|---|---|---|
| Outcome-only RL | Shared system + task + registry + rollout history | Outcome reward only |
| RL + TIPS | Same agent prompt as outcome-only RL | Outcome + TIPS shaping |
| RL + SPICE | Same agent prompt as outcome-only RL | Outcome + SPICE shaping |
| SFT demonstration collection | Shared system + task + registry + history; guidance with MLToolBench | Full SFT on filtered trajectories |
| Unguided evaluation, both interfaces | Shared system + task + condition-specific registry + history | None |
| Guided Gemini 3.6 Flash evaluation | Shared system + task + registry + teacher guidance | None |
| Instance suffix | Source | Held-out source / target |
|---|---|---|
| cv-ood-01 | APTOS 2019 | Unseen medical-image source and ordinal target |
| cv-ood-02 | RSNA Breast Cancer | Unseen mammography source and probabilistic-F1 metric |
| nlp-ood-01 | AI4Code | Unseen code-notebook source and ranking target |
| nlp-ood-02 | BabyLM | Unseen language-modeling source and next-token target |
| tabular-ood-01 | California Housing | Unseen dataset and reconstruction target; downstream R² is secondary |
| tabular-ood-02 | Breast Cancer Wisconsin | Unseen clinical tabular source and calibration target |
| Instance suffix | Template | Configuration | Seed |
|---|---|---|---|
| cv-train-01 | cv-cifar10 | healthy | 10000 |
| cv-train-02 | cv-cifar10 | input scale | 10001 |
| cv-train-03 | cv-cifar10 | rare class sampling | 10002 |
| cv-train-04 | cv-cifar10 | overaggressive crop | 10003 |
| cv-train-05 | cv-cifar100 | healthy | 10010 |
| cv-train-06 | cv-cifar100 | normalization scale | 10011 |