use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use crate::{FerrumError, Result, TokenId};
pub const DEFAULT_CHAT_REPETITION_PENALTY: f32 = 1.1;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SamplingParams {
pub max_tokens: usize,
pub temperature: f32,
pub top_p: f32,
pub top_k: Option<usize>,
pub repetition_penalty: f32,
pub presence_penalty: f32,
pub frequency_penalty: f32,
pub stop_sequences: Vec<String>,
pub seed: Option<u64>,
pub min_p: Option<f32>,
pub tfs: Option<f32>,
pub typical_p: Option<f32>,
pub mirostat: Option<MirostatParams>,
#[serde(default)]
pub response_format: ResponseFormat,
#[serde(default)]
pub structured_output_start: StructuredOutputStart,
#[serde(default)]
pub response_completion_boundary: ResponseCompletionBoundary,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(tag = "mode", content = "delimiter", rename_all = "snake_case")]
pub enum StructuredOutputStart {
#[default]
Immediate,
AfterDelimiter(String),
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ResponseCompletionEnvelope {
pub open_token_text: String,
pub close_token_text: String,
pub max_envelopes: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
#[serde(tag = "mode", rename_all = "snake_case")]
pub enum ResponseCompletionBoundary {
#[default]
Immediate,
AfterDelimiterAndPayload {
delimiter: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
alternate_envelope: Option<ResponseCompletionEnvelope>,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(tag = "type", content = "schema")]
#[derive(Default)]
pub enum ResponseFormat {
#[default]
Text,
JsonObject,
JsonSchema(String),
}
impl Default for SamplingParams {
fn default() -> Self {
Self {
max_tokens: 512,
temperature: 1.0,
top_p: 1.0,
top_k: None,
repetition_penalty: 1.0,
presence_penalty: 0.0,
frequency_penalty: 0.0,
stop_sequences: vec![],
seed: None,
min_p: None,
tfs: None,
typical_p: None,
mirostat: None,
response_format: ResponseFormat::default(),
structured_output_start: StructuredOutputStart::default(),
response_completion_boundary: ResponseCompletionBoundary::default(),
}
}
}
impl SamplingParams {
pub fn greedy() -> Self {
Self {
temperature: 0.0,
top_p: 1.0,
top_k: None,
..Default::default()
}
}
pub fn with_temperature(temperature: f32) -> Self {
Self {
temperature,
..Default::default()
}
}
pub fn validate(&self) -> Result<()> {
if !self.temperature.is_finite() || self.temperature < 0.0 {
return Err(FerrumError::invalid_request(
"temperature must be finite and non-negative".to_string(),
));
}
if !self.top_p.is_finite() || self.top_p <= 0.0 || self.top_p > 1.0 {
return Err(FerrumError::invalid_request(
"top_p must be in range (0, 1]".to_string(),
));
}
if let Some(top_k) = self.top_k {
if top_k == 0 {
return Err(FerrumError::invalid_request(
"top_k must be positive".to_string(),
));
}
}
if !self.repetition_penalty.is_finite() || self.repetition_penalty <= 0.0 {
return Err(FerrumError::invalid_request(
"repetition_penalty must be finite and positive".to_string(),
));
}
if !self.presence_penalty.is_finite() || !(-2.0..=2.0).contains(&self.presence_penalty) {
return Err(FerrumError::invalid_request(
"presence_penalty must be in range [-2, 2]".to_string(),
));
}
if !self.frequency_penalty.is_finite() || !(-2.0..=2.0).contains(&self.frequency_penalty) {
return Err(FerrumError::invalid_request(
"frequency_penalty must be in range [-2, 2]".to_string(),
));
}
if let Some(min_p) = self.min_p {
if !min_p.is_finite() || min_p <= 0.0 || min_p > 1.0 {
return Err(FerrumError::invalid_request(
"min_p must be in range (0, 1]".to_string(),
));
}
}
if let Some(tfs) = self.tfs {
if !tfs.is_finite() || tfs <= 0.0 || tfs > 1.0 {
return Err(FerrumError::invalid_request(
"tfs must be in range (0, 1]".to_string(),
));
}
}
if let Some(typical_p) = self.typical_p {
if !typical_p.is_finite() || typical_p <= 0.0 || typical_p > 1.0 {
return Err(FerrumError::invalid_request(
"typical_p must be in range (0, 1]".to_string(),
));
}
}
if let ResponseCompletionBoundary::AfterDelimiterAndPayload {
delimiter,
alternate_envelope,
} = &self.response_completion_boundary
{
if delimiter.is_empty() {
return Err(FerrumError::invalid_request(
"response completion delimiter must not be empty".to_string(),
));
}
if let Some(envelope) = alternate_envelope {
if envelope.open_token_text.is_empty() || envelope.close_token_text.is_empty() {
return Err(FerrumError::invalid_request(
"response completion envelope tokens must not be empty".to_string(),
));
}
if envelope.max_envelopes == 0 {
return Err(FerrumError::invalid_request(
"response completion envelope limit must be greater than zero".to_string(),
));
}
}
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MirostatParams {
pub mode: u8,
pub tau: f32,
pub eta: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SamplingPresets {
pub presets: HashMap<String, SamplingParams>,
}
impl Default for SamplingPresets {
fn default() -> Self {
let mut presets = HashMap::new();
presets.insert("greedy".to_string(), SamplingParams::greedy());
presets.insert(
"creative".to_string(),
SamplingParams {
temperature: 1.2,
top_p: 0.9,
top_k: Some(50),
repetition_penalty: 1.1,
..Default::default()
},
);
presets.insert(
"precise".to_string(),
SamplingParams {
temperature: 0.3,
top_p: 0.95,
top_k: Some(20),
repetition_penalty: 1.05,
..Default::default()
},
);
Self { presets }
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize, Default,
)]
pub enum Priority {
Low = 0,
#[default]
Normal = 1,
High = 2,
Critical = 3,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum FinishReason {
Length,
Stop,
EOS,
Cancelled,
Error,
ContentFilter,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SpecialTokens {
pub bos_token: Option<TokenId>,
pub eos_token: Option<TokenId>,
pub unk_token: Option<TokenId>,
pub pad_token: Option<TokenId>,
pub sep_token: Option<TokenId>,
pub cls_token: Option<TokenId>,
pub mask_token: Option<TokenId>,
#[serde(default)]
pub extra_eos_tokens: Vec<TokenId>,
}