use super::SamplingParams;
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct RecommendedSampling {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<usize>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct RequestedSampling {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<usize>,
}
impl RecommendedSampling {
pub fn is_empty(&self) -> bool {
*self == RecommendedSampling::default()
}
pub fn resolve(
&self,
requested: RequestedSampling,
framework: SamplingParams,
) -> SamplingParams {
SamplingParams {
temperature: requested
.temperature
.or(self.temperature)
.unwrap_or(framework.temperature),
top_p: requested.top_p.or(self.top_p).unwrap_or(framework.top_p),
top_k: requested.top_k.or(self.top_k).unwrap_or(framework.top_k),
..framework
}
}
pub fn from_generation_config(json: &str) -> Self {
let Ok(serde_json::Value::Object(map)) = serde_json::from_str::<serde_json::Value>(json)
else {
return RecommendedSampling::default();
};
if map.get("do_sample").and_then(|v| v.as_bool()) == Some(false) {
return RecommendedSampling {
temperature: Some(0.0),
..RecommendedSampling::default()
};
}
RecommendedSampling {
temperature: map
.get("temperature")
.and_then(|v| v.as_f64())
.map(|v| v as f32),
top_p: map.get("top_p").and_then(|v| v.as_f64()).map(|v| v as f32),
top_k: map
.get("top_k")
.and_then(|v| v.as_u64())
.map(|v| v as usize),
}
}
pub fn from_model_dir(dir: &std::path::Path) -> Self {
match std::fs::read_to_string(dir.join("generation_config.json")) {
Ok(text) => Self::from_generation_config(&text),
Err(_) => RecommendedSampling::default(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_absent_generation_config_key_stays_absent_rather_than_taking_a_default() {
let recommended = RecommendedSampling::from_generation_config(r#"{"temperature": 0.6}"#);
assert_eq!(recommended.temperature, Some(0.6));
assert_eq!(recommended.top_p, None, "top_p was not in the file");
assert_eq!(recommended.top_k, None, "top_k was not in the file");
let nulled = RecommendedSampling::from_generation_config(r#"{"top_p": null}"#);
assert_eq!(nulled, RecommendedSampling::default());
}
#[test]
fn every_generation_config_key_present_is_recommended() {
let recommended = RecommendedSampling::from_generation_config(
r#"{"do_sample": true, "temperature": 1.0, "top_k": 20, "top_p": 0.95}"#,
);
assert_eq!(
recommended,
RecommendedSampling {
temperature: Some(1.0),
top_p: Some(0.95),
top_k: Some(20),
}
);
}
#[test]
fn do_sample_false_recommends_greedy_and_no_other_field() {
let recommended = RecommendedSampling::from_generation_config(
r#"{"do_sample": false, "temperature": 0.7, "top_k": 50, "top_p": 0.9}"#,
);
assert_eq!(recommended.temperature, Some(0.0));
assert_eq!(recommended.top_p, None);
assert_eq!(recommended.top_k, None);
}
#[test]
fn a_malformed_generation_config_recommends_nothing() {
for text in ["", "not json", "[1, 2, 3]", "null"] {
assert!(
RecommendedSampling::from_generation_config(text).is_empty(),
"{text:?} must recommend nothing"
);
}
}
#[test]
fn a_request_outranks_the_recommendation_which_outranks_the_framework_default() {
let recommended = RecommendedSampling {
temperature: Some(1.0),
top_p: Some(0.95),
top_k: Some(20),
};
let resolved = recommended.resolve(
RequestedSampling {
temperature: Some(0.0),
..RequestedSampling::default()
},
SamplingParams::default(),
);
assert_eq!(resolved.temperature, 0.0, "the request asked for greedy");
assert_eq!(resolved.top_p, 0.95, "the request said nothing about top_p");
assert_eq!(resolved.top_k, 20, "the request said nothing about top_k");
assert_eq!(resolved.repetition_penalty, 1.0);
}
#[test]
fn a_checkpoint_that_recommends_nothing_leaves_the_framework_defaults_alone() {
let resolved = RecommendedSampling::default()
.resolve(RequestedSampling::default(), SamplingParams::default());
let default = SamplingParams::default();
assert_eq!(resolved.temperature, default.temperature);
assert_eq!(resolved.top_p, default.top_p);
assert_eq!(resolved.top_k, default.top_k);
}
#[test]
fn a_model_directory_without_a_generation_config_recommends_nothing() {
let dir = std::env::temp_dir().join(format!(
"ferrox_test_no_generation_config_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
assert!(RecommendedSampling::from_model_dir(&dir).is_empty());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn a_model_directory_generation_config_is_read_from_beside_the_weights() {
let dir = std::env::temp_dir().join(format!(
"ferrox_test_generation_config_dir_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(
dir.join("generation_config.json"),
r#"{"temperature": 0.6, "top_p": 0.95}"#,
)
.unwrap();
let recommended = RecommendedSampling::from_model_dir(&dir);
std::fs::remove_dir_all(&dir).ok();
assert_eq!(recommended.temperature, Some(0.6));
assert_eq!(recommended.top_p, Some(0.95));
assert_eq!(recommended.top_k, None);
}
}