use serde_json::Value;
use crate::sampling_loop::PerTokenProbs;
pub(crate) fn render(
per_token: &PerTokenProbs,
top_k: Option<usize>,
decode: &dyn Fn(usize) -> String,
) -> Value {
let Some(top_k) = top_k else {
return Value::Null;
};
if per_token.is_empty() {
return serde_json::json!({
"tokens": [],
"token_logprobs": [],
"top_logprobs": [],
"text_offset": [],
});
}
let mut tokens: Vec<String> = Vec::with_capacity(per_token.len());
let mut token_logprobs: Vec<f64> = Vec::with_capacity(per_token.len());
let mut top_logprobs: Vec<Value> = Vec::with_capacity(per_token.len());
let mut text_offset: Vec<usize> = Vec::with_capacity(per_token.len());
let mut offset = 0usize;
for (id, probs) in per_token {
let piece = decode(*id);
text_offset.push(offset);
offset += piece.len();
tokens.push(piece);
let chosen = probs.get(*id).copied().unwrap_or(0.0);
token_logprobs.push(logprob(chosen));
let mut alternatives: Vec<(usize, f32)> = probs
.iter()
.enumerate()
.filter(|(_, p)| **p > 0.0)
.map(|(i, p)| (i, *p))
.collect();
alternatives.sort_by(|a, b| b.1.total_cmp(&a.1).then(a.0.cmp(&b.0)));
alternatives.truncate(top_k);
let mut map = serde_json::Map::new();
for (alt, p) in alternatives {
map.insert(decode(alt), Value::from(logprob(p)));
}
top_logprobs.push(Value::Object(map));
}
serde_json::json!({
"tokens": tokens,
"token_logprobs": token_logprobs,
"top_logprobs": top_logprobs,
"text_offset": text_offset,
})
}
fn logprob(p: f32) -> f64 {
(p as f64).ln()
}
pub(crate) fn render_chat(
per_token: &PerTokenProbs,
top_k: Option<usize>,
decode: &dyn Fn(usize) -> String,
) -> Value {
let Some(top_k) = top_k else {
return Value::Null;
};
let entry = |id: usize, p: f32, decode: &dyn Fn(usize) -> String| {
let piece = decode(id);
serde_json::json!({
"token": piece,
"logprob": logprob(p),
"bytes": piece.as_bytes(),
})
};
let content: Vec<Value> = per_token
.iter()
.map(|(id, probs)| {
let chosen = probs.get(*id).copied().unwrap_or(0.0);
let mut alternatives: Vec<(usize, f32)> = probs
.iter()
.enumerate()
.filter(|(_, p)| **p > 0.0)
.map(|(i, p)| (i, *p))
.collect();
alternatives.sort_by(|a, b| b.1.total_cmp(&a.1).then(a.0.cmp(&b.0)));
alternatives.truncate(top_k);
let mut e = entry(*id, chosen, decode);
e["top_logprobs"] = Value::Array(
alternatives
.into_iter()
.map(|(alt, p)| entry(alt, p, decode))
.collect(),
);
e
})
.collect();
serde_json::json!({ "content": content })
}
#[cfg(test)]
mod tests {
use super::*;
fn decoder() -> impl Fn(usize) -> String {
|id: usize| format!("t{id}")
}
#[test]
fn the_chat_shape_is_entries_not_parallel_arrays() {
let per_token: PerTokenProbs = vec![(1, vec![0.25, 0.5, 0.25, 0.0])];
let out = render_chat(&per_token, Some(2), &decoder());
assert!(out["tokens"].is_null(), "completions shape leaked: {out}");
assert!(out["text_offset"].is_null(), "{out}");
let content = out["content"].as_array().expect("content");
assert_eq!(content.len(), 1, "one entry per token: {out}");
let e = &content[0];
assert_eq!(e["token"], "t1");
assert_eq!(
e["bytes"],
serde_json::json!([116, 49]),
"the token's UTF-8"
);
let chosen = e["logprob"].as_f64().expect("a real number");
assert!((chosen - 0.5f64.ln()).abs() < 1e-9, "{chosen}");
let top = e["top_logprobs"].as_array().expect("top_logprobs");
assert_eq!(top.len(), 2, "asked for 2: {top:?}");
assert_eq!(top[0]["token"], "t1");
for alt in top {
assert!(alt["bytes"].is_array(), "{alt}");
assert!(
alt["logprob"].as_f64().expect("a real number").is_finite(),
"ln(0) reached the wire: {alt}"
);
}
}
#[test]
fn the_chat_shape_omits_removed_candidates() {
let per_token: PerTokenProbs = vec![(0, vec![0.7, 0.3, 0.0, 0.0])];
let out = render_chat(&per_token, Some(4), &decoder());
let top = out["content"][0]["top_logprobs"].as_array().unwrap();
assert_eq!(top.len(), 2, "a zero was reported: {top:?}");
}
#[test]
fn the_chat_shape_is_absent_when_not_asked_for() {
assert_eq!(render_chat(&Vec::new(), None, &decoder()), Value::Null);
}
#[test]
fn no_request_means_no_object() {
assert_eq!(render(&Vec::new(), None, &decoder()), Value::Null);
}
#[test]
fn the_chosen_token_and_its_alternatives_are_reported() {
let per_token: PerTokenProbs = vec![(1, vec![0.25, 0.5, 0.25, 0.0])];
let out = render(&per_token, Some(3), &decoder());
assert_eq!(out["tokens"], serde_json::json!(["t1"]));
assert_eq!(out["text_offset"], serde_json::json!([0]));
let chosen = out["token_logprobs"][0].as_f64().expect("a real number");
assert!((chosen - 0.5f64.ln()).abs() < 1e-9, "{chosen}");
let top = out["top_logprobs"][0].as_object().expect("an object");
assert_eq!(top.len(), 3, "three survived the filter: {top:?}");
assert!(
!top.contains_key("t3"),
"a zero-probability candidate was reported: {top:?}"
);
assert!((top["t1"].as_f64().unwrap() - 0.5f64.ln()).abs() < 1e-9);
}
#[test]
fn a_removed_candidate_is_omitted_even_with_room_to_spare() {
let per_token: PerTokenProbs = vec![(0, vec![0.7, 0.3, 0.0, 0.0])];
let out = render(&per_token, Some(4), &decoder());
let top = out["top_logprobs"][0].as_object().expect("an object");
assert_eq!(
top.len(),
2,
"a zero-probability candidate was reported: {top:?}"
);
assert!(
!top.contains_key("t2") && !top.contains_key("t3"),
"{top:?}"
);
for v in top.values() {
assert!(
v.as_f64().expect("a real number").is_finite(),
"ln(0) reached the wire: {top:?}"
);
}
}
#[test]
fn top_k_keeps_the_likeliest() {
let per_token: PerTokenProbs = vec![(0, vec![0.6, 0.3, 0.1])];
let out = render(&per_token, Some(2), &decoder());
let top = out["top_logprobs"][0].as_object().unwrap();
assert_eq!(top.len(), 2);
assert!(top.contains_key("t0") && top.contains_key("t1"));
assert!(!top.contains_key("t2"), "kept the least likely: {top:?}");
}
#[test]
fn text_offsets_index_the_text_they_describe() {
let per_token: PerTokenProbs = vec![
(1, vec![0.0, 1.0]),
(0, vec![1.0, 0.0]),
(1, vec![0.0, 1.0]),
];
let out = render(&per_token, Some(1), &decoder());
let tokens: Vec<String> = serde_json::from_value(out["tokens"].clone()).unwrap();
let offsets: Vec<usize> = serde_json::from_value(out["text_offset"].clone()).unwrap();
let text = tokens.concat();
for (i, off) in offsets.iter().enumerate() {
assert!(
text[*off..].starts_with(&tokens[i]),
"offset {off} does not point at {:?} in {text:?}",
tokens[i]
);
}
}
#[test]
fn an_empty_generation_still_reports_the_object() {
let out = render(&Vec::new(), Some(2), &decoder());
assert!(out.is_object(), "{out}");
assert_eq!(out["tokens"], serde_json::json!([]));
}
}