From a17f349daa0609402c51205929b9f2b32f6fb1bb Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 24 Apr 2024 16:12:38 +0800 Subject: feat: support customizing `top_p` parameter (#434) --- src/config/mod.rs | 32 ++++++++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) (limited to 'src/config/mod.rs') 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, pub temperature: Option, + pub top_p: Option, pub dry_run: bool, pub save: bool, pub save_session: Option, @@ -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) { + 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) { 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()); } -- cgit v1.2.3