summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-24 16:12:38 +0800
committerGitHub <noreply@github.com>2024-04-24 16:12:38 +0800
commita17f349daa0609402c51205929b9f2b32f6fb1bb (patch)
tree380dd9c158f8d9278355634a8d474df3a90f8d25 /src/config
parent040c48b9b392e3329d3f1de8eeb1b8773129773c (diff)
downloadaichat-a17f349daa0609402c51205929b9f2b32f6fb1bb.tar.gz
feat: support customizing `top_p` parameter (#434)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/mod.rs32
-rw-r--r--src/config/role.rs12
-rw-r--r--src/config/session.rs19
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;