Skip to content

Commit 4df29f6

Browse files
committed
feat: capture RL training logprobs
1 parent 222dab1 commit 4df29f6

20 files changed

Lines changed: 564 additions & 9 deletions

core/src/agent/llm_turn.rs

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -33,17 +33,14 @@ impl AgentLoop {
3333
"LLM completion started"
3434
);
3535

36-
let selected_tool_names =
37-
crate::tools::select_tools_for_messages(&self.config.tools, &state.messages)
38-
.into_iter()
39-
.map(|tool| tool.name)
40-
.collect::<Vec<_>>();
36+
let selected_tools =
37+
crate::tools::select_tools_for_messages(&self.config.tools, &state.messages);
4138
self.config.rl_trajectory_recorder.record_llm_request(
4239
session_id.unwrap_or(""),
4340
turn,
4441
&state.messages,
4542
augmented_system.as_deref(),
46-
&selected_tool_names,
43+
&selected_tools,
4744
estimate_prompt_tokens(&state.messages, augmented_system.as_deref()),
4845
);
4946

core/src/agent/tests.rs

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -263,6 +263,7 @@ impl MockLlmClient {
263263
cache_write_tokens: None,
264264
},
265265
stop_reason: Some("end_turn".to_string()),
266+
token_logprobs: Vec::new(),
266267
meta: None,
267268
}
268269
}
@@ -291,6 +292,7 @@ impl MockLlmClient {
291292
cache_write_tokens: None,
292293
},
293294
stop_reason: Some("tool_use".to_string()),
295+
token_logprobs: Vec::new(),
294296
meta: None,
295297
}
296298
}
@@ -1080,6 +1082,7 @@ async fn test_agent_hitl_multiple_tool_calls() {
10801082
cache_write_tokens: None,
10811083
},
10821084
stop_reason: Some("tool_use".to_string()),
1085+
token_logprobs: Vec::new(),
10831086
meta: None,
10841087
},
10851088
MockLlmClient::text_response("Both executed!"),
@@ -1164,6 +1167,7 @@ async fn test_agent_hitl_partial_approval() {
11641167
cache_write_tokens: None,
11651168
},
11661169
stop_reason: Some("tool_use".to_string()),
1170+
token_logprobs: Vec::new(),
11671171
meta: None,
11681172
},
11691173
MockLlmClient::text_response("First worked, second rejected."),
@@ -1893,6 +1897,7 @@ async fn test_agent_multiple_tools_single_turn() {
18931897
cache_write_tokens: None,
18941898
},
18951899
stop_reason: Some("tool_use".to_string()),
1900+
token_logprobs: Vec::new(),
18961901
meta: None,
18971902
},
18981903
MockLlmClient::text_response("Both commands ran"),

core/src/agent_api.rs

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -215,6 +215,14 @@ pub struct SessionOptions {
215215
/// training or service data collection. If unset, the same config can
216216
/// be enabled by `A3S_CODE_TRAJECTORY_PATH`.
217217
pub rl_trajectory: Option<crate::rl_trajectory::RlTrajectoryConfig>,
218+
/// Request token-level log probabilities from compatible LLM providers.
219+
///
220+
/// This is off by default because many public providers reject logprob
221+
/// requests with tool calls. Training/evaluation harnesses using compatible
222+
/// OpenAI-style backends can enable it explicitly.
223+
pub llm_logprobs: Option<bool>,
224+
/// Number of alternative token logprobs to request per generated token.
225+
pub llm_top_logprobs: Option<usize>,
218226
/// Auto-save after each completed `send()` or default-history `stream()` call.
219227
pub auto_save: bool,
220228
/// Optional artifact retention limits for large tool/program outputs.

core/src/agent_api/session_config.rs

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,13 +74,44 @@ pub(super) fn resolve_session_llm_client(
7474
}
7575
}
7676

77+
let logprobs = opts
78+
.llm_logprobs
79+
.or_else(|| env_bool("A3S_CODE_LLM_LOGPROBS"))
80+
.or_else(|| env_bool("A3S_CODE_OPENAI_LOGPROBS"));
81+
if let Some(enabled) = logprobs {
82+
llm_config = llm_config.with_logprobs(enabled);
83+
}
84+
85+
let top_logprobs = opts
86+
.llm_top_logprobs
87+
.or_else(|| env_usize("A3S_CODE_LLM_TOP_LOGPROBS"))
88+
.or_else(|| env_usize("A3S_CODE_OPENAI_TOP_LOGPROBS"));
89+
if let Some(top_logprobs) = top_logprobs {
90+
llm_config = llm_config.with_top_logprobs(top_logprobs);
91+
}
92+
7793
if let Some(session_id) = session_id {
7894
llm_config = llm_config.with_session_id(session_id);
7995
}
8096

8197
Ok(crate::llm::create_client_with_config(llm_config))
8298
}
8399

100+
fn env_bool(name: &str) -> Option<bool> {
101+
let value = std::env::var(name).ok()?;
102+
match value.trim().to_ascii_lowercase().as_str() {
103+
"1" | "true" | "yes" | "on" => Some(true),
104+
"0" | "false" | "no" | "off" => Some(false),
105+
_ => None,
106+
}
107+
}
108+
109+
fn env_usize(name: &str) -> Option<usize> {
110+
std::env::var(name)
111+
.ok()
112+
.and_then(|value| value.trim().parse::<usize>().ok())
113+
}
114+
84115
pub(super) struct ResolvedSessionMemory {
85116
pub(super) memory: Option<Arc<crate::memory::AgentMemory>>,
86117
pub(super) init_warning: Option<String>,

core/src/agent_api/session_options.rs

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,8 @@ impl std::fmt::Debug for SessionOptions {
4343
.field("session_store", &self.session_store.is_some())
4444
.field("session_id", &self.session_id)
4545
.field("rl_trajectory", &self.rl_trajectory)
46+
.field("llm_logprobs", &self.llm_logprobs)
47+
.field("llm_top_logprobs", &self.llm_top_logprobs)
4648
.field("auto_save", &self.auto_save)
4749
.field("artifact_store_limits", &self.artifact_store_limits)
4850
.field("max_parse_retries", &self.max_parse_retries)
@@ -376,6 +378,19 @@ impl SessionOptions {
376378
self
377379
}
378380

381+
/// Request token-level log probabilities from compatible LLM providers.
382+
pub fn with_llm_logprobs(mut self, enabled: bool) -> Self {
383+
self.llm_logprobs = Some(enabled);
384+
self
385+
}
386+
387+
/// Request up to `top_logprobs` alternative logprobs per generated token.
388+
pub fn with_llm_top_logprobs(mut self, top_logprobs: usize) -> Self {
389+
self.llm_logprobs = Some(true);
390+
self.llm_top_logprobs = Some(top_logprobs);
391+
self
392+
}
393+
379394
/// Enable auto-save after each `send()` call
380395
pub fn with_auto_save(mut self, enabled: bool) -> Self {
381396
self.auto_save = enabled;

core/src/agent_api/tests.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ impl StaticStreamingClient {
3030
cache_write_tokens: None,
3131
},
3232
stop_reason: Some("end_turn".to_string()),
33+
token_logprobs: Vec::new(),
3334
meta: None,
3435
}
3536
}

core/src/llm/anthropic.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -245,6 +245,7 @@ impl AnthropicClient {
245245
cache_write_tokens: parsed.usage.cache_creation_input_tokens,
246246
},
247247
stop_reason: Some(parsed.stop_reason),
248+
token_logprobs: Vec::new(),
248249
meta: Some(LlmResponseMeta {
249250
provider: Some(self.provider_name.clone()),
250251
request_model: Some(self.model.clone()),
@@ -581,6 +582,7 @@ impl AnthropicClient {
581582
},
582583
usage: usage.clone(),
583584
stop_reason: stop_reason.clone(),
585+
token_logprobs: Vec::new(),
584586
meta: Some(LlmResponseMeta {
585587
provider: Some(provider_name.clone()),
586588
request_model: Some(request_model.clone()),

core/src/llm/factory.rs

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,10 @@ pub struct LlmConfig {
2626
pub max_tokens: Option<usize>,
2727
/// Extended thinking budget in tokens (Anthropic only).
2828
pub thinking_budget: Option<usize>,
29+
/// Request token-level log probabilities from OpenAI-compatible providers.
30+
pub logprobs: Option<bool>,
31+
/// Number of alternative logprobs per token when logprobs are requested.
32+
pub top_logprobs: Option<usize>,
2933
/// When true, temperature is never sent to the API (e.g., o1 models).
3034
pub disable_temperature: bool,
3135
}
@@ -47,6 +51,8 @@ impl std::fmt::Debug for LlmConfig {
4751
.field("temperature", &self.temperature)
4852
.field("max_tokens", &self.max_tokens)
4953
.field("thinking_budget", &self.thinking_budget)
54+
.field("logprobs", &self.logprobs)
55+
.field("top_logprobs", &self.top_logprobs)
5056
.field("disable_temperature", &self.disable_temperature)
5157
.finish()
5258
}
@@ -70,6 +76,8 @@ impl LlmConfig {
7076
temperature: None,
7177
max_tokens: None,
7278
thinking_budget: None,
79+
logprobs: None,
80+
top_logprobs: None,
7381
disable_temperature: false,
7482
}
7583
}
@@ -114,6 +122,17 @@ impl LlmConfig {
114122
self
115123
}
116124

125+
pub fn with_logprobs(mut self, enabled: bool) -> Self {
126+
self.logprobs = Some(enabled);
127+
self
128+
}
129+
130+
pub fn with_top_logprobs(mut self, top_logprobs: usize) -> Self {
131+
self.logprobs = Some(true);
132+
self.top_logprobs = Some(top_logprobs);
133+
self
134+
}
135+
117136
pub(crate) fn resolved_headers(&self) -> HashMap<String, String> {
118137
let mut headers = self.headers.clone();
119138
if let (Some(header_name), Some(session_id)) = (&self.session_id_header, &self.session_id) {
@@ -168,6 +187,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc<dyn LlmClient> {
168187
if let Some(max) = config.max_tokens {
169188
client = client.with_max_tokens(max);
170189
}
190+
if let Some(enabled) = config.logprobs {
191+
client = client.with_logprobs(enabled);
192+
}
193+
if let Some(top_logprobs) = config.top_logprobs {
194+
client = client.with_top_logprobs(top_logprobs);
195+
}
171196
Arc::new(client)
172197
}
173198
"glm" | "zhipu" | "bigmodel" => {
@@ -183,6 +208,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc<dyn LlmClient> {
183208
if let Some(max) = config.max_tokens {
184209
client = client.with_max_tokens(max);
185210
}
211+
if let Some(enabled) = config.logprobs {
212+
client = client.with_logprobs(enabled);
213+
}
214+
if let Some(top_logprobs) = config.top_logprobs {
215+
client = client.with_top_logprobs(top_logprobs);
216+
}
186217
Arc::new(client)
187218
}
188219
// OpenAI-compatible providers (deepseek, groq, together, ollama, etc.)
@@ -208,6 +239,12 @@ pub fn create_client_with_config(config: LlmConfig) -> Arc<dyn LlmClient> {
208239
if let Some(max) = config.max_tokens {
209240
client = client.with_max_tokens(max);
210241
}
242+
if let Some(enabled) = config.logprobs {
243+
client = client.with_logprobs(enabled);
244+
}
245+
if let Some(top_logprobs) = config.top_logprobs {
246+
client = client.with_top_logprobs(top_logprobs);
247+
}
211248
Arc::new(client)
212249
}
213250
}

0 commit comments

Comments
 (0)