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/session.rs | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) (limited to 'src/config/session.rs') 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, + top_p: Option, #[serde(default)] save_session: Option, messages: Vec, @@ -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 { + self.top_p + } + pub fn save_session(&self) -> Option { 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) { + if self.top_p != value { + self.top_p = value; + self.dirty = true; + } + } + pub fn set_save_session(&mut self, value: Option) { if self.save_session != value { self.save_session = value; -- cgit v1.2.3