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,
})
}
pub(crate) fn render_echoed(
prompt: &[usize],
prompt_rows: &[Vec<f32>],
prompt_text: &str,
per_token: &PerTokenProbs,
top_k: Option<usize>,
decode: &dyn Fn(usize) -> String,
) -> Value {
let Some(top_k) = top_k else {
return Value::Null;
};
let mut tokens: Vec<Value> = Vec::with_capacity(prompt.len() + per_token.len());
let mut token_logprobs: Vec<Value> = Vec::with_capacity(prompt.len() + per_token.len());
let mut top_logprobs: Vec<Value> = Vec::with_capacity(prompt.len() + per_token.len());
let mut text_offset: Vec<usize> = Vec::with_capacity(prompt.len() + per_token.len());
let mut offset = 0usize;
for (i, id) in prompt.iter().enumerate() {
let piece = decode(*id);
text_offset.push(offset);
offset += piece.len();
tokens.push(Value::from(piece));
match i.checked_sub(1).and_then(|r| prompt_rows.get(r)) {
Some(logits) => {
let probs = softmax(logits);
token_logprobs.push(Value::from(logprob(probs.get(*id).copied().unwrap_or(0.0))));
top_logprobs.push(top_map(&probs, top_k, decode));
}
None => {
token_logprobs.push(Value::Null);
top_logprobs.push(Value::Null);
}
}
}
offset = prompt_text.len();
for (id, probs) in per_token {
let piece = decode(*id);
text_offset.push(offset);
offset += piece.len();
tokens.push(Value::from(piece));
token_logprobs.push(Value::from(logprob(probs.get(*id).copied().unwrap_or(0.0))));
top_logprobs.push(top_map(probs, top_k, decode));
}
serde_json::json!({
"tokens": tokens,
"token_logprobs": token_logprobs,
"top_logprobs": top_logprobs,
"text_offset": text_offset,
})
}
fn top_map(probs: &[f32], top_k: usize, decode: &dyn Fn(usize) -> String) -> Value {
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)));
}
Value::Object(map)
}
pub(crate) fn piece_renderer<'a>(
as_ids: bool,
decode: &'a dyn Fn(usize) -> String,
) -> Box<dyn Fn(usize) -> String + 'a> {
if as_ids {
Box::new(|id| format!("token_id:{id}"))
} else {
Box::new(decode)
}
}
fn logprob(p: f32) -> f64 {
(p as f64).ln()
}
pub(crate) fn render_prompt(
prompt: &[usize],
per_position: &[Vec<f32>],
top_k: usize,
decode: &dyn Fn(usize) -> String,
) -> Value {
let mut out: Vec<Value> = Vec::with_capacity(prompt.len());
out.push(Value::Null);
for (i, id) in prompt.iter().enumerate().skip(1) {
let Some(logits) = per_position.get(i - 1) else {
out.push(Value::Null);
continue;
};
let probs = softmax(logits);
let mut alternatives: Vec<(usize, f32)> = probs
.iter()
.enumerate()
.map(|(t, p)| (t, *p))
.filter(|(_, p)| *p > 0.0)
.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();
let actual = probs.get(*id).copied().unwrap_or(0.0);
map.insert(
decode(*id),
serde_json::json!({
"logprob": logprob(actual),
"rank": rank_of(&probs, *id),
"decoded_token": decode(*id),
}),
);
for (alt, p) in alternatives {
map.entry(decode(alt)).or_insert_with(|| {
serde_json::json!({
"logprob": logprob(p),
"rank": rank_of(&probs, alt),
"decoded_token": decode(alt),
})
});
}
out.push(Value::Object(map));
}
Value::Array(out)
}
fn rank_of(probs: &[f32], id: usize) -> usize {
let p = probs.get(id).copied().unwrap_or(0.0);
1 + probs.iter().filter(|q| **q > p).count()
}
fn softmax(logits: &[f32]) -> Vec<f32> {
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
if !max.is_finite() {
return vec![0.0; logits.len()];
}
let mut out: Vec<f32> = logits.iter().map(|l| (l - max).exp()).collect();
let total: f32 = out.iter().sum();
if total > 0.0 {
for p in &mut out {
*p /= total;
}
}
out
}
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 the_prompt_is_scored_from_the_second_token_on() {
let prompt = vec![0usize, 1, 2];
let rows = vec![
vec![0.0f32, 2.0, 0.0],
vec![3.0f32, 3.0, 0.0],
];
let out = render_prompt(&prompt, &rows, 2, &decoder());
let arr = out.as_array().expect("an array");
assert_eq!(arr.len(), 3, "one entry per prompt token: {out}");
assert!(arr[0].is_null(), "nothing predicted the first token");
let e1 = arr[1].as_object().expect("an object");
assert_eq!(e1["t1"]["rank"], 1, "{e1:?}");
assert!(e1["t1"]["logprob"].as_f64().unwrap() > -0.5, "{e1:?}");
let e2 = arr[2].as_object().expect("an object");
assert!(
e2.contains_key("t2"),
"the token that actually followed was omitted: {e2:?}"
);
assert_eq!(e2["t2"]["rank"], 3, "{e2:?}");
assert!(
e2["t2"]["logprob"].as_f64().unwrap() < -2.0,
"an unlikely token was reported as likely: {e2:?}"
);
}
#[test]
fn a_prompt_token_no_sampler_would_pick_still_gets_a_number() {
let prompt = vec![0usize, 2];
let rows = vec![vec![12.0f32, 0.0, -12.0]];
let out = render_prompt(&prompt, &rows, 1, &decoder());
let e = out[1].as_object().expect("an object");
let v = e["t2"]["logprob"].as_f64().expect("a real number");
assert!(v.is_finite(), "ln(0) reached the wire: {e:?}");
assert!(v < -20.0, "a filtered distribution was used: {e:?}");
assert_eq!(e["t2"]["rank"], 3);
}
#[test]
fn missing_rows_are_null_rather_than_a_shorter_array() {
let out = render_prompt(&[0usize, 1, 2], &[vec![0.0, 1.0, 0.0]], 1, &decoder());
let arr = out.as_array().unwrap();
assert_eq!(arr.len(), 3);
assert!(arr[0].is_null() && arr[2].is_null(), "{out}");
assert!(arr[1].is_object(), "{out}");
}
#[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!([]));
}
}