use eredu_core::{
ModelRuntime, TextContinuationBoundary, TextGenerationBackend, TokenFilterController,
};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, Default, PartialEq, Serialize, Deserialize)]
pub struct SamplingOverride {
pub temperature: Option<f32>,
pub reseed: Option<u64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct SamplingStateFacts {
pub temperature: f32,
pub requires_positive_temperature: bool,
pub has_rng: bool,
}
#[derive(Debug, Clone, Copy)]
pub struct ValidatedSamplingOverride {
temperature: f32,
reseed: Option<u64>,
}
impl ValidatedSamplingOverride {
pub fn temperature(self) -> f32 {
self.temperature
}
pub fn reseed(self) -> Option<u64> {
self.reseed
}
}
pub trait TextSamplingControlBackend: TextGenerationBackend {
fn sampling_control_facts(state: &Self::TextGenerationState) -> SamplingStateFacts;
fn install_sampling_override(
runtime: &mut ModelRuntime<Self>,
state: &mut Self::TextGenerationState,
request: ValidatedSamplingOverride,
) -> Result<(), Self::Error>;
}
#[derive(Debug, thiserror::Error)]
pub enum SamplingOverrideError<E: std::error::Error + 'static> {
#[error("invalid sampling override: {0}")]
Invalid(&'static str),
#[error("native sampling override failed: {0}")]
Backend(#[source] E),
}
fn validate<E: std::error::Error + 'static>(
facts: SamplingStateFacts,
request: SamplingOverride,
) -> Result<ValidatedSamplingOverride, SamplingOverrideError<E>> {
let temperature = request.temperature.unwrap_or(facts.temperature);
if !temperature.is_finite() || temperature < 0.0 {
return Err(SamplingOverrideError::Invalid(
"temperature must be finite and nonnegative",
));
}
if facts.requires_positive_temperature && temperature == 0.0 {
return Err(SamplingOverrideError::Invalid(
"adaptive sampling requires positive temperature",
));
}
if temperature > 0.0 && !facts.has_rng && request.reseed.is_none() {
return Err(SamplingOverrideError::Invalid(
"no inherited RNG exists; an explicit seed is required",
));
}
Ok(ValidatedSamplingOverride {
temperature,
reseed: request.reseed,
})
}
pub fn apply_sampling_override<B: TextSamplingControlBackend, C: TokenFilterController>(
boundary: &mut TextContinuationBoundary<'_, '_, B, C>,
request: SamplingOverride,
) -> Result<SamplingStateFacts, SamplingOverrideError<B::Error>> {
let (runtime, state, _) = boundary.mechanism_parts();
apply_prepared_sampling_override(runtime, state, request)
}
pub fn apply_prepared_sampling_override<B: TextSamplingControlBackend>(
runtime: &mut ModelRuntime<B>,
state: &mut B::TextGenerationState,
request: SamplingOverride,
) -> Result<SamplingStateFacts, SamplingOverrideError<B::Error>> {
if let eredu_core::execution_control::ControlSupport::Unsupported { .. } =
B::text_sampling_control_support(runtime)
{
return Err(SamplingOverrideError::Invalid(
"loaded execution does not support sampling overrides",
));
}
let action = validate(B::sampling_control_facts(state), request)?;
B::install_sampling_override(runtime, state, action).map_err(SamplingOverrideError::Backend)?;
Ok(B::sampling_control_facts(state))
}
#[cfg(test)]
mod tests {
use super::*;
fn check(
facts: SamplingStateFacts,
request: SamplingOverride,
) -> Result<ValidatedSamplingOverride, SamplingOverrideError<std::io::Error>> {
validate(facts, request)
}
#[test]
fn temperature_changes_preserve_rng_and_adaptation_unless_reseeding_is_explicit() {
let initial = SamplingStateFacts {
temperature: 0.0,
requires_positive_temperature: false,
has_rng: false,
};
let stochastic = SamplingOverride {
temperature: Some(0.7),
reseed: None,
};
assert!(check(initial, stochastic).is_err());
let seeded = check(
initial,
SamplingOverride {
reseed: Some(42),
..stochastic
},
)
.unwrap();
assert_eq!(seeded.reseed(), Some(42));
assert_eq!(seeded.temperature(), 0.7);
let retained = SamplingStateFacts {
has_rng: true,
..initial
};
assert_eq!(check(retained, stochastic).unwrap().reseed(), None);
for invalid in [f32::NAN, f32::INFINITY, -1.0] {
assert!(check(
retained,
SamplingOverride {
temperature: Some(invalid),
reseed: None
}
)
.is_err());
}
let adaptive = SamplingStateFacts {
temperature: 0.8,
requires_positive_temperature: true,
has_rng: true,
};
assert!(check(
adaptive,
SamplingOverride {
temperature: Some(0.0),
reseed: None
}
)
.is_err());
assert_eq!(
check(adaptive, SamplingOverride::default())
.unwrap()
.temperature(),
0.8
);
}
}