use std::fmt;
use pathwise::Problem;
use pathwise::optimization::{
AsyncOptimizationProblem, BudgetedSearchOptions, budgeted_local_search,
};
use serde::Serialize;
use crate::compiled::CompiledProgram;
use crate::eval::{Comparison, Dataset, EvalOptions, Metric, Report, compare, evaluate_with};
use crate::{Demonstration, Program, Provider, Signature};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct FewShot {
pub k: usize,
pub max_calls: usize,
pub seed: u64,
pub concurrency: usize,
}
impl Default for FewShot {
fn default() -> Self {
Self {
k: 4,
max_calls: 400,
seed: 0,
concurrency: 4,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Trial {
pub demonstrations: Vec<usize>,
pub score: f64,
pub accepted: bool,
}
#[derive(Debug, Clone)]
pub struct Optimized {
pub compiled: CompiledProgram,
pub baseline: Report,
pub best: Report,
pub comparison: Comparison,
pub trials: Vec<Trial>,
pub local_optimum: bool,
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum OptimizeError {
PoolTooSmall {
pool: usize,
k: usize,
},
EmptyValidation,
BudgetTooSmall {
max_calls: usize,
per_evaluation: usize,
},
}
impl fmt::Display for OptimizeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::PoolTooSmall { pool, k } => {
write!(f, "{pool} candidate demonstrations, need at least {k}")
}
Self::EmptyValidation => write!(f, "the validation dataset is empty"),
Self::BudgetTooSmall {
max_calls,
per_evaluation,
} => write!(
f,
"a budget of {max_calls} calls does not cover one evaluation \
({per_evaluation} calls)"
),
}
}
}
impl std::error::Error for OptimizeError {}
pub fn demonstrations_from_labels<S>(dataset: &Dataset<S>) -> Vec<Demonstration>
where
S: Signature,
{
dataset
.examples
.iter()
.filter(|e| e.expected_output().is_some())
.map(|e| Demonstration {
input: serde_json::to_value(&e.input).expect("inputs serialize to JSON"),
output: serde_json::Value::Object(e.expected.clone()),
})
.collect()
}
pub async fn bootstrap_demonstrations<S, P, M>(
teacher: &Program<S, P>,
dataset: &Dataset<S>,
metric: &M,
) -> Vec<Demonstration>
where
S: Signature + Clone,
S::Output: Serialize,
P: Provider,
M: Metric<S>,
{
let mut pool = Vec::new();
for example in &dataset.examples {
let Ok(output) = teacher.run(example.input.clone()).await else {
continue;
};
if metric.score(example, &output) >= 1.0 {
pool.push(Demonstration {
input: serde_json::to_value(&example.input).expect("inputs serialize to JSON"),
output: serde_json::to_value(&output).expect("outputs serialize to JSON"),
});
}
}
pool
}
pub async fn optimize_few_shot<S, P, M>(
program: &Program<S, P>,
pool: &[Demonstration],
validation: &Dataset<S>,
metric: &M,
options: FewShot,
) -> Result<Optimized, OptimizeError>
where
S: Signature + Clone,
S::Output: Serialize,
P: Provider,
M: Metric<S>,
{
if pool.len() < options.k {
return Err(OptimizeError::PoolTooSmall {
pool: pool.len(),
k: options.k,
});
}
if validation.examples.is_empty() {
return Err(OptimizeError::EmptyValidation);
}
let per_evaluation = validation.examples.len();
let max_evaluations = options.max_calls / per_evaluation;
if max_evaluations == 0 {
return Err(OptimizeError::BudgetTooSmall {
max_calls: options.max_calls,
per_evaluation,
});
}
let eval_options = EvalOptions {
concurrency: options.concurrency,
epochs: 1,
};
let baseline = evaluate_with(program, validation, metric, eval_options).await;
let search = DemonstrationSearch {
provider: program.provider(),
base: program.compile(),
pool,
validation,
metric,
k: options.k,
initial: seeded_choice(pool.len(), options.k, options.seed),
eval_options,
};
let solution = budgeted_local_search(
&search,
BudgetedSearchOptions {
max_evaluations,
seed: options.seed,
},
)
.await;
let best = solution.score;
let compiled = search
.with_demonstrations(&solution.state)
.with_provenance(&best);
let comparison = compare(&baseline, &best, 0.0).expect("same program, metric and dataset");
let trials = solution
.evaluations
.into_iter()
.map(|e| Trial {
demonstrations: e.state,
score: e.score.score,
accepted: e.accepted,
})
.collect();
Ok(Optimized {
compiled,
baseline,
best,
comparison,
trials,
local_optimum: solution.local_optimum,
})
}
fn seeded_choice(n: usize, k: usize, seed: u64) -> Vec<usize> {
let mut state = seed;
let mut next = || {
state = state.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
};
let mut indices: Vec<usize> = (0..n).collect();
for i in 0..k {
let j = i + (next() % (n - i) as u64) as usize;
indices.swap(i, j);
}
let mut chosen = indices[..k].to_vec();
chosen.sort_unstable();
chosen
}
struct DemonstrationSearch<'a, S: Signature, P, M> {
provider: &'a P,
base: CompiledProgram,
pool: &'a [Demonstration],
validation: &'a Dataset<S>,
metric: &'a M,
k: usize,
initial: Vec<usize>,
eval_options: EvalOptions,
}
impl<S: Signature, P, M> DemonstrationSearch<'_, S, P, M> {
fn with_demonstrations(&self, chosen: &[usize]) -> CompiledProgram {
let mut compiled = self.base.clone();
compiled.demonstrations = chosen.iter().map(|&i| self.pool[i].clone()).collect();
compiled
}
}
impl<S: Signature, P, M> Problem for DemonstrationSearch<'_, S, P, M> {
type State = Vec<usize>;
type Move = (usize, usize);
fn initial(&self) -> Vec<usize> {
self.initial.clone()
}
fn moves(&self, state: &Vec<usize>) -> impl Iterator<Item = (usize, usize)> {
let unchosen: Vec<usize> = (0..self.pool.len())
.filter(|i| !state.contains(i))
.collect();
(0..self.k).flat_map(move |slot| unchosen.clone().into_iter().map(move |item| (slot, item)))
}
fn apply(&self, state: &Vec<usize>, &(slot, item): &(usize, usize)) -> Vec<usize> {
let mut next = state.clone();
next[slot] = item;
next.sort_unstable();
next
}
fn is_goal(&self, _: &Vec<usize>) -> bool {
false
}
}
impl<S, P, M> AsyncOptimizationProblem for DemonstrationSearch<'_, S, P, M>
where
S: Signature + Clone,
S::Output: Serialize,
P: Provider,
M: Metric<S>,
{
type Score = Report;
async fn evaluate(&self, state: &Vec<usize>) -> Report {
let compiled = self.with_demonstrations(state);
let program = Program::<S, &P>::from_parts(self.provider, &compiled);
evaluate_with(&program, self.validation, self.metric, self.eval_options).await
}
fn is_improvement(&self, candidate: &Report, incumbent: &Report) -> bool {
candidate.score > incumbent.score
}
}
#[cfg(test)]
mod tests {
use super::seeded_choice;
#[test]
fn seeded_choices_are_distinct_sorted_and_reproducible() {
let a = seeded_choice(10, 4, 7);
assert_eq!(a.len(), 4);
assert!(a.windows(2).all(|w| w[0] < w[1]));
assert!(a.iter().all(|&i| i < 10));
assert_eq!(a, seeded_choice(10, 4, 7));
assert_ne!(seeded_choice(10, 4, 7), seeded_choice(10, 4, 8));
assert_eq!(seeded_choice(3, 3, 1), vec![0, 1, 2]);
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Reflective {
pub iterations: usize,
pub max_calls: usize,
pub mistakes_shown: usize,
pub concurrency: usize,
}
impl Default for Reflective {
fn default() -> Self {
Self {
iterations: 6,
max_calls: 400,
mistakes_shown: 5,
concurrency: 4,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct InstructionTrial {
pub instructions: String,
pub score: Option<f64>,
pub accepted: bool,
pub error: Option<String>,
}
#[derive(Debug, Clone)]
pub struct OptimizedInstructions {
pub compiled: CompiledProgram,
pub baseline: Report,
pub best: Report,
pub comparison: Comparison,
pub trials: Vec<InstructionTrial>,
}
#[derive(Debug, Clone, Serialize, serde::Deserialize)]
pub struct Mistake {
pub input: serde_json::Value,
pub expected: serde_json::Value,
pub actual: Option<serde_json::Value>,
pub error: Option<String>,
}
#[derive(Debug, Clone, Serialize, serde::Deserialize)]
pub struct ProposeInstructions {
pub task: String,
pub current_instructions: String,
pub output_schema: serde_json::Value,
pub mistakes: Vec<Mistake>,
}
#[derive(Debug, Clone, Serialize, serde::Deserialize, schemars::JsonSchema)]
pub struct ProposedInstructions {
pub instructions: String,
}
impl Signature for ProposeInstructions {
type Output = ProposedInstructions;
const NAME: &'static str = "ProposeInstructions";
const DESCRIPTION: &'static str = "You improve the instructions of a language-model \
program. You get the program's task, its current instructions, the JSON Schema of \
its output, and examples it got wrong with the expected and the actual output. \
Write new, complete instructions that keep what works and prevent these mistakes. \
State general rules; do not quote or refer to the examples. Do not describe the \
output format beyond what the schema says.";
fn validate(output: &ProposedInstructions) -> Result<(), Vec<String>> {
if output.instructions.trim().is_empty() {
return Err(vec!["instructions must not be empty".into()]);
}
Ok(())
}
}
pub async fn optimize_instructions<S, P, M, T>(
program: &Program<S, P>,
teacher: &T,
validation: &Dataset<S>,
metric: &M,
options: Reflective,
) -> Result<OptimizedInstructions, OptimizeError>
where
S: Signature + Clone,
S::Output: Serialize,
P: Provider,
M: Metric<S>,
T: Provider,
{
if validation.examples.is_empty() {
return Err(OptimizeError::EmptyValidation);
}
let per_evaluation = validation.examples.len();
if options.max_calls < 2 * per_evaluation + 1 {
return Err(OptimizeError::BudgetTooSmall {
max_calls: options.max_calls,
per_evaluation,
});
}
let eval_options = EvalOptions {
concurrency: options.concurrency,
epochs: 1,
};
let base = program.compile();
let schema = crate::output_schema::<S>();
let proposer = Program::<ProposeInstructions, &T>::new(teacher);
let (baseline, answers) =
crate::eval::evaluate_detailed(program, validation, metric, eval_options).await;
let mut calls = per_evaluation;
let mut current = (program.effective_instructions(), baseline.clone(), answers);
let mut trials = Vec::new();
for iteration in 0..options.iterations {
if calls + 1 + per_evaluation > options.max_calls {
break;
}
let mistakes = mistakes(
validation,
¤t.1,
¤t.2,
options.mistakes_shown,
iteration,
);
if mistakes.is_empty() {
break;
}
calls += 1;
let request = ProposeInstructions {
task: S::DESCRIPTION.into(),
current_instructions: current.0.clone(),
output_schema: schema.clone(),
mistakes,
};
let proposal = match proposer.run(request).await {
Ok(proposal) => proposal.instructions,
Err(error) => {
trials.push(InstructionTrial {
instructions: String::new(),
score: None,
accepted: false,
error: Some(error.to_string()),
});
continue;
}
};
let mut candidate = base.clone();
candidate.instructions = Some(proposal.clone());
let candidate_program = Program::<S, &P>::from_parts(program.provider(), &candidate);
let (report, answers) =
crate::eval::evaluate_detailed(&candidate_program, validation, metric, eval_options)
.await;
calls += per_evaluation;
let accepted = report.score > current.1.score;
trials.push(InstructionTrial {
instructions: proposal.clone(),
score: Some(report.score),
accepted,
error: None,
});
if accepted {
current = (proposal, report, answers);
}
}
let (instructions, best, _) = current;
let mut compiled = base;
if best.score > baseline.score {
compiled.instructions = Some(instructions);
}
let compiled = compiled.with_provenance(&best);
let comparison = compare(&baseline, &best, 0.0).expect("same program, metric and dataset");
Ok(OptimizedInstructions {
compiled,
baseline,
best,
comparison,
trials,
})
}
fn mistakes<S: Signature>(
dataset: &Dataset<S>,
report: &Report,
answers: &[(usize, crate::eval::Answer)],
limit: usize,
iteration: usize,
) -> Vec<Mistake> {
let wrong: Vec<&(usize, crate::eval::Answer)> = answers
.iter()
.filter(|(index, _)| report.scores.get(*index).is_some_and(|s| *s < 1.0))
.collect();
if wrong.is_empty() || limit == 0 {
return Vec::new();
}
let start = (iteration * limit) % wrong.len();
wrong
.iter()
.cycle()
.skip(start)
.take(limit.min(wrong.len()))
.map(|(index, answer)| {
let example = &dataset.examples[*index];
Mistake {
input: serde_json::to_value(&example.input).expect("inputs serialize to JSON"),
expected: serde_json::Value::Object(example.expected.clone()),
actual: answer.as_ref().ok().cloned(),
error: answer.as_ref().err().cloned(),
}
})
.collect()
}