use crate::tokenizer::bpe::BpeTokenizer;
use std::collections::HashMap;
pub(crate) fn build_byte_decoder() -> HashMap<char, u8> {
let byte_encoder = bytes_to_unicode();
let mut byte_decoder: HashMap<char, u8> = HashMap::new();
for (byte_val, &ch) in byte_encoder.iter().enumerate() {
byte_decoder.insert(ch, byte_val as u8);
}
byte_decoder
}
fn append_token_bytes(
tokenizer: &BpeTokenizer,
id: u32,
byte_decoder: &HashMap<char, u8>,
out: &mut Vec<u8>,
) {
if let Some(token_str) = tokenizer.token_for_id(id) {
for ch in token_str.chars() {
if let Some(&b) = byte_decoder.get(&ch) {
out.push(b);
}
}
}
}
pub(crate) fn decode_tokens(tokenizer: &BpeTokenizer, ids: &[u32]) -> String {
let byte_decoder = build_byte_decoder();
let mut bytes = Vec::new();
for &id in ids {
append_token_bytes(tokenizer, id, &byte_decoder, &mut bytes);
}
String::from_utf8_lossy(&bytes).to_string()
}
pub(crate) struct IncrementalDetokenizer {
byte_decoder: HashMap<char, u8>,
bytes: Vec<u8>,
flushed: usize,
}
impl IncrementalDetokenizer {
pub(crate) fn new() -> Self {
Self {
byte_decoder: build_byte_decoder(),
bytes: Vec::new(),
flushed: 0,
}
}
pub(crate) fn push(&mut self, tokenizer: &BpeTokenizer, id: u32) -> String {
append_token_bytes(tokenizer, id, &self.byte_decoder, &mut self.bytes);
self.flush_complete()
}
fn flush_complete(&mut self) -> String {
let mut out = String::new();
loop {
match std::str::from_utf8(&self.bytes[self.flushed..]) {
Ok(s) => {
out.push_str(s);
self.flushed = self.bytes.len();
return out;
}
Err(e) => {
let valid = e.valid_up_to();
if valid > 0 {
out.push_str(
String::from_utf8_lossy(
&self.bytes[self.flushed..self.flushed + valid],
)
.as_ref(),
);
self.flushed += valid;
}
match e.error_len() {
None => return out,
Some(len) => {
out.push('\u{FFFD}');
self.flushed += len;
}
}
}
}
}
}
pub(crate) fn finish(&mut self) -> String {
if self.flushed < self.bytes.len() {
let tail = String::from_utf8_lossy(&self.bytes[self.flushed..]).into_owned();
self.flushed = self.bytes.len();
tail
} else {
String::new()
}
}
pub(crate) fn text(&self) -> String {
String::from_utf8_lossy(&self.bytes).to_string()
}
}
pub fn bytes_to_unicode() -> Vec<char> {
let mut bs = Vec::new();
bs.extend(33u16..=126);
bs.extend(161u16..=172);
bs.extend(174u16..=255);
let mut cs = bs.clone();
let mut n = 0u16;
for b in 0u16..=255u16 {
if !bs.contains(&b) {
bs.push(b);
cs.push(256 + n);
n += 1;
}
}
let mut table = vec!['\0'; 256];
for (b, c) in bs.into_iter().zip(cs) {
table[b as usize] =
char::from_u32(c as u32).expect("invariant: byte-to-unicode codepoint is valid");
}
table
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn incremental_flush_reconstructs_split_codepoints() {
let chunks: &[&[u8]] = &[&[0xE5, 0xA5], &[0xBD], &[0xF0, 0x9F], &[0x98, 0x80], b"ok"];
let mut d = IncrementalDetokenizer::new();
let mut out = String::new();
for c in chunks {
d.bytes.extend_from_slice(c);
out.push_str(&d.flush_complete());
}
out.push_str(&d.finish());
assert_eq!(out, "好😀ok");
assert_eq!(out, d.text());
}
#[test]
fn incremental_flush_truncated_mid_codepoint_matches_lossy() {
let mut d = IncrementalDetokenizer::new();
let mut out = String::new();
d.bytes.extend_from_slice(&[b'h', b'i', 0xE5, 0xA5]);
out.push_str(&d.flush_complete());
out.push_str(&d.finish());
assert_eq!(out, String::from_utf8_lossy(&[b'h', b'i', 0xE5, 0xA5]));
assert_eq!(out, d.text());
}
#[test]
fn incremental_flush_invalid_byte_does_not_stall_stream() {
let mut d = IncrementalDetokenizer::new();
d.bytes.extend_from_slice(&[0x80]);
let first = d.flush_complete();
assert_eq!(first, "\u{FFFD}", "invalid byte must flush immediately");
d.bytes.extend_from_slice(b"A");
let second = d.flush_complete();
assert_eq!(second, "A", "valid byte after invalid one must flush");
assert_eq!(
format!("{first}{second}"),
String::from_utf8_lossy(&[0x80, b'A'])
);
}
#[test]
fn incremental_flush_invalid_between_valid_matches_lossy() {
let raw = [b'h', b'i', 0xFF, b'y', b'o'];
let mut d = IncrementalDetokenizer::new();
d.bytes.extend_from_slice(&raw);
let out = d.flush_complete();
assert_eq!(out, String::from_utf8_lossy(&raw));
assert_eq!(out, d.text());
}
#[test]
fn incremental_flush_ascii_is_exact_per_chunk() {
let mut d = IncrementalDetokenizer::new();
let mut out = String::new();
for c in [b"He".as_slice(), b"llo", b"!"] {
d.bytes.extend_from_slice(c);
out.push_str(&d.flush_complete());
}
out.push_str(&d.finish());
assert_eq!(out, "Hello!");
}
}