diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/mod.rs | 32 | ||||
| -rw-r--r-- | src/config/role.rs | 12 | ||||
| -rw-r--r-- | src/config/session.rs | 19 |
3 files changed, 60 insertions, 3 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index e1cec4d..87be519 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -53,6 +53,7 @@ pub struct Config { #[serde(rename(serialize = "model", deserialize = "model"))] pub model_id: Option<String>, pub temperature: Option<f64>, + pub top_p: Option<f64>, pub dry_run: bool, pub save: bool, pub save_session: Option<bool>, @@ -89,6 +90,7 @@ impl Default for Config { Self { model_id: None, temperature: None, + top_p: None, save: true, save_session: None, highlight: true, @@ -297,6 +299,7 @@ impl Config { if let Some(session) = self.session.as_mut() { session.guard_empty()?; session.set_temperature(role.temperature); + session.set_top_p(role.top_p); } self.role = Some(role); Ok(()) @@ -335,6 +338,16 @@ impl Config { } } + pub fn set_top_p(&mut self, value: Option<f64>) { + if let Some(session) = self.session.as_mut() { + session.set_top_p(value); + } else if let Some(role) = self.role.as_mut() { + role.set_top_p(value); + } else { + self.top_p = value; + } + } + pub fn set_save_session(&mut self, value: Option<bool>) { if let Some(session) = self.session.as_mut() { session.set_save_session(value); @@ -411,6 +424,7 @@ impl Config { let items = vec![ ("model", self.model.id()), ("temperature", format_option(&self.temperature)), + ("top_p", format_option(&self.top_p)), ("dry_run", self.dry_run.to_string()), ("save", self.save.to_string()), ("save_session", format_option(&self.save_session)), @@ -478,6 +492,7 @@ impl Config { ".session" => self.list_sessions(), ".set" => vec![ "temperature ", + "top_p ", "compress_threshold", "save ", "save_session ", @@ -529,6 +544,10 @@ impl Config { let value = parse_value(value)?; self.set_temperature(value); } + "top_p" => { + let value = parse_value(value)?; + self.set_top_p(value); + } "compress_threshold" => { let value = parse_value(value)?; self.set_compress_threshold(value); @@ -756,10 +775,18 @@ impl Config { } else { self.temperature }; + let top_p = if let Some(session) = input.session(&self.session) { + session.top_p() + } else if let Some(role) = input.role() { + role.top_p + } else { + self.top_p + }; self.model.max_input_tokens_limit(&messages)?; Ok(SendData { messages, temperature, + top_p, stream, }) } @@ -791,6 +818,11 @@ impl Config { output.insert("temperature", temperature.to_string()); } } + if let Some(top_p) = self.top_p { + if top_p != 0.0 { + output.insert("top_p", top_p.to_string()); + } + } if self.dry_run { output.insert("dry_run", "true".to_string()); } diff --git a/src/config/role.rs b/src/config/role.rs index 50d5b5e..b226622 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -16,12 +16,10 @@ pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Role { - /// Role name pub name: String, - /// Prompt text pub prompt: String, - /// Temperature value pub temperature: Option<f64>, + pub top_p: Option<f64>, } impl Role { @@ -30,6 +28,7 @@ impl Role { name: TEMP_ROLE.into(), prompt: prompt.into(), temperature: None, + top_p: None, } } @@ -67,6 +66,7 @@ If there is a lack of details, provide most logical solution. Output plain text only, without any markdown formatting."# ), temperature: None, + top_p: None, } } @@ -79,6 +79,7 @@ Provide short responses in about 80 words. APPLY MARKDOWN formatting when possible."# .into(), temperature: None, + top_p: None, } } @@ -89,6 +90,7 @@ APPLY MARKDOWN formatting when possible."# If there is a lack of details, provide most logical solution, without requesting further clarification."# .into(), temperature: None, + top_p: None, } } @@ -106,6 +108,10 @@ If there is a lack of details, provide most logical solution, without requesting self.temperature = value; } + pub fn set_top_p(&mut self, value: Option<f64>) { + self.top_p = value; + } + pub fn complete_prompt_args(&mut self, name: &str) { self.name = name.to_string(); self.prompt = complete_prompt_args(&self.prompt, &self.name); diff --git a/src/config/session.rs b/src/config/session.rs index 801615f..4ebf0d3 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -18,6 +18,7 @@ pub struct Session { #[serde(rename(serialize = "model", deserialize = "model"))] model_id: String, temperature: Option<f64>, + top_p: Option<f64>, #[serde(default)] save_session: Option<bool>, messages: Vec<Message>, @@ -43,6 +44,7 @@ impl Session { Self { model_id: config.model.id(), temperature: config.temperature, + top_p: config.top_p, save_session: config.save_session, messages: vec![], compressed_messages: vec![], @@ -80,6 +82,10 @@ impl Session { self.temperature } + pub fn top_p(&self) -> Option<f64> { + self.top_p + } + pub fn save_session(&self) -> Option<bool> { self.save_session } @@ -111,6 +117,9 @@ impl Session { if let Some(temperature) = self.temperature() { data["temperature"] = temperature.into(); } + if let Some(top_p) = self.top_p() { + data["top_p"] = top_p.into(); + } if let Some(save_session) = self.save_session() { data["save_session"] = save_session.into(); } @@ -140,6 +149,9 @@ impl Session { if let Some(temperature) = self.temperature() { items.push(("temperature", temperature.to_string())); } + if let Some(top_p) = self.top_p() { + items.push(("top_p", top_p.to_string())); + } if let Some(save_session) = self.save_session() { items.push(("save_session", save_session.to_string())); @@ -207,6 +219,13 @@ impl Session { } } + pub fn set_top_p(&mut self, value: Option<f64>) { + if self.top_p != value { + self.top_p = value; + self.dirty = true; + } + } + pub fn set_save_session(&mut self, value: Option<bool>) { if self.save_session != value { self.save_session = value; |
