use burn::tensor::{backend::Backend, Int, Tensor, TensorData};
use super::gemma::{GemmaModel, KvCache, ROPE_CACHE_LEN};
use super::tokenizer::GemmaTokenizer;
const EOS_ID: i64 = 1;
const STOP_MARKER: &str = "\nUser:";
struct StreamEmitter {
sent: String,
stopped: bool,
}
impl StreamEmitter {
fn new() -> Self {
Self {
sent: String::new(),
stopped: false,
}
}
fn step<'a>(&mut self, full: &'a str, final_flush: bool) -> (Option<&'a str>, bool) {
if self.stopped {
return (None, true);
}
if full.get(..self.sent.len()) != Some(self.sent.as_str()) {
return (None, false);
}
let (cut, hit) = match full.find(STOP_MARKER) {
Some(i) => (i, true),
None => (full.len(), false),
};
let end = if hit || final_flush {
cut
} else {
let mut end = full.len().saturating_sub(STOP_MARKER.len() - 1);
while end > 0 && !full.is_char_boundary(end) {
end -= 1;
}
end.min(cut)
};
if hit {
self.stopped = true;
}
let start = self.sent.len();
if end > start {
self.sent.push_str(&full[start..end]);
(Some(&full[start..end]), hit)
} else {
(None, hit)
}
}
}
fn at_context_limit(seq: usize) -> bool {
seq >= ROPE_CACHE_LEN
}
pub async fn generate<B: Backend>(
model: &GemmaModel<B>,
tok: &GemmaTokenizer,
prompt: &str,
max_new: usize,
device: &B::Device,
) -> String {
generate_streamed(model, tok, prompt, max_new, device, |_| true).await
}
pub async fn generate_streamed<B: Backend>(
model: &GemmaModel<B>,
tok: &GemmaTokenizer,
prompt: &str,
max_new: usize,
device: &B::Device,
mut on_delta: impl FnMut(&str) -> bool,
) -> String {
let tokens: Vec<i64> = tok.encode(prompt);
let mut cache: KvCache<B> = model.new_cache();
let mut pending: Vec<i64> = tokens;
let mut total_len = pending.len();
let mut generated: Vec<i64> = Vec::with_capacity(max_new);
let mut emitter = StreamEmitter::new();
#[cfg(target_arch = "wasm32")]
let t_start = js_sys::Date::now();
#[cfg(target_arch = "wasm32")]
web_sys::console::log_1(
&format!("[lh-local] generate: prompt={} tokens, max_new={max_new}", total_len).into(),
);
for _ in 0..max_new {
if at_context_limit(total_len) {
break;
}
let seq = pending.len();
let input = Tensor::<B, 1, Int>::from_data(TensorData::from(pending.as_slice()), device)
.reshape([1, seq]);
let logits = model.forward_cached(input, &mut cache);
let argmax = logits.argmax(2);
let data = match argmax.into_data_async().await {
Ok(d) => d,
Err(_) => break,
};
let ids: Vec<i64> = data.iter::<i64>().collect();
let next = match ids.last().copied() {
Some(id) => id,
None => break, };
if next == EOS_ID {
break;
}
pending = vec![next];
total_len += 1;
generated.push(next);
let full = tok.decode(&generated);
let (delta, hit_marker) = emitter.step(&full, false);
let keep_going = delta.map(&mut on_delta).unwrap_or(true);
#[cfg(target_arch = "wasm32")]
{
let dt = (js_sys::Date::now() - t_start) / 1000.0;
let tps = generated.len() as f64 / dt.max(1e-9);
web_sys::console::log_1(
&format!(
"[lh-local] tok {}/{max_new} ({tps:.2} tok/s) text={:?}",
generated.len(),
full
)
.into(),
);
}
if hit_marker || !keep_going {
break;
}
}
let full = tok.decode(&generated);
if let (Some(delta), _) = emitter.step(&full, true) {
on_delta(delta);
}
full
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn context_limit_guards_the_rope_cache_boundary() {
assert!(!at_context_limit(0));
assert!(!at_context_limit(ROPE_CACHE_LEN - 1)); assert!(at_context_limit(ROPE_CACHE_LEN)); assert!(at_context_limit(ROPE_CACHE_LEN + 1));
}
fn run_emitter(decodes: &[&str], flush: &str) -> (String, bool) {
let mut e = StreamEmitter::new();
let mut out = String::new();
let mut hit = false;
for d in decodes {
let (delta, h) = e.step(d, false);
if let Some(t) = delta {
out.push_str(t);
}
hit |= h;
if hit {
return (out, hit);
}
}
if let (Some(t), h) = e.step(flush, true) {
out.push_str(t);
hit |= h;
}
(out, hit)
}
#[test]
fn emitter_streams_all_text_with_final_flush() {
let (out, hit) = run_emitter(&[" Paris", " Paris is", " Paris is nice."], " Paris is nice.");
assert_eq!(out, " Paris is nice.");
assert!(!hit);
}
#[test]
fn emitter_cuts_at_stop_marker_split_across_tokens() {
let (out, hit) = run_emitter(
&["Paris.", "Paris.\nUser", "Paris.\nUser: and"],
"Paris.\nUser: and",
);
assert_eq!(out, "Paris.");
assert!(hit);
}
#[test]
fn emitter_stops_on_marker_and_goes_quiet() {
let mut e = StreamEmitter::new();
let (d, hit) = e.step("hello there\nUser: hi", false);
assert_eq!(d, Some("hello there"));
assert!(hit);
assert_eq!(e.step("hello there\nUser: hi more", true), (None, true));
}
#[test]
fn emitter_goes_quiet_on_prefix_instability() {
let mut e = StreamEmitter::new();
let (d, _) = e.step("hello world!!", false);
assert_eq!(d, Some("hello wo"));
assert_eq!(e.step("hellO world!! more", false), (None, false));
assert_eq!(e.step("hellO world!! more", true), (None, false));
}
#[tokio::test]
#[ignore]
async fn gemma_native_stream() {
let dir = std::env::var("GEMMA_DIR")
.expect("set GEMMA_DIR to a folder with model.safetensors + tokenizer.json");
let weights = std::fs::read(format!("{dir}/model.safetensors")).expect("read weights");
let tok_bytes =
std::fs::read(format!("{dir}/tokenizer.json")).expect("read tokenizer.json");
let device = burn::backend::wgpu::WgpuDevice::default();
let model = super::super::gemma::GemmaModel::<super::super::LocalBackend>::init(
super::super::gemma::GemmaConfig::gemma_3_270m(),
&device,
);
let model =
super::super::weights::load_gemma(model, &weights, &device).expect("load_gemma");
let tok = super::super::tokenizer::GemmaTokenizer::from_bytes(&tok_bytes)
.expect("load tokenizer");
let prompt = std::env::var("GEMMA_PROMPT")
.unwrap_or_else(|_| "The capital of France is".to_string());
let max_new: usize = std::env::var("GEMMA_MAX_NEW")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(64);
let mut streamed = String::new();
let mut deltas = 0usize;
let t0 = std::time::Instant::now();
let out = generate_streamed(&model, &tok, &prompt, max_new, &device, |d| {
streamed.push_str(d);
deltas += 1;
true
})
.await;
let dt = t0.elapsed().as_secs_f64();
let n_tok = tok.encode(&out).len().saturating_sub(1);
println!(
"\n=== GEMMA NATIVE STREAM ===\nprompt: {prompt:?}\noutput: {out:?}\n\
deltas: {deltas}, tokens: {n_tok}, {dt:.2}s, {:.2} tok/s\n===========================\n",
n_tok as f64 / dt.max(1e-9)
);
let cut = out.find(STOP_MARKER).unwrap_or(out.len());
assert_eq!(
streamed,
&out[..cut],
"streamed deltas must reproduce the returned continuation up to the stop cut"
);
assert!(deltas > 1, "expected incremental deltas, got a single blob");
}
#[tokio::test]
#[ignore]
async fn gemma_kv_parity_and_speed() {
let dir = std::env::var("GEMMA_DIR")
.expect("set GEMMA_DIR to a folder with model.safetensors + tokenizer.json");
let weights = std::fs::read(format!("{dir}/model.safetensors")).expect("read weights");
let tok_bytes =
std::fs::read(format!("{dir}/tokenizer.json")).expect("read tokenizer.json");
let device = burn::backend::wgpu::WgpuDevice::default();
let model = super::super::gemma::GemmaModel::<super::super::LocalBackend>::init(
super::super::gemma::GemmaConfig::gemma_3_270m(),
&device,
);
let model =
super::super::weights::load_gemma(model, &weights, &device).expect("load_gemma");
let tok = super::super::tokenizer::GemmaTokenizer::from_bytes(&tok_bytes)
.expect("load tokenizer");
let prompt = std::env::var("GEMMA_PROMPT")
.unwrap_or_else(|_| "The capital of France is".to_string());
let max_new: usize = std::env::var("GEMMA_MAX_NEW")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(256);
let mut tokens: Vec<i64> = tok.encode(&prompt);
let prompt_len = tokens.len();
let mut uncached: Vec<i64> = Vec::new();
let t0 = std::time::Instant::now();
for _ in 0..max_new {
if at_context_limit(tokens.len()) {
break;
}
let seq = tokens.len();
let input = Tensor::<super::super::LocalBackend, 1, Int>::from_data(
TensorData::from(tokens.as_slice()),
&device,
)
.reshape([1, seq]);
let logits = model.forward(input);
let data = logits.argmax(2).into_data_async().await.expect("read-back");
let next = data.iter::<i64>().last().expect("nonempty");
if next == EOS_ID {
break;
}
tokens.push(next);
uncached.push(next);
}
let dt_uncached = t0.elapsed().as_secs_f64();
let mut cache = model.new_cache();
let mut pending: Vec<i64> = tok.encode(&prompt);
let mut total_len = pending.len();
let mut cached: Vec<i64> = Vec::new();
let t1 = std::time::Instant::now();
for _ in 0..max_new {
if at_context_limit(total_len) {
break;
}
let seq = pending.len();
let input = Tensor::<super::super::LocalBackend, 1, Int>::from_data(
TensorData::from(pending.as_slice()),
&device,
)
.reshape([1, seq]);
let logits = model.forward_cached(input, &mut cache);
let data = logits.argmax(2).into_data_async().await.expect("read-back");
let next = data.iter::<i64>().last().expect("nonempty");
if next == EOS_ID {
break;
}
pending = vec![next];
total_len += 1;
cached.push(next);
}
let dt_cached = t1.elapsed().as_secs_f64();
println!(
"\n=== GEMMA KV PARITY + SPEED ===\nprompt: {prompt:?} ({prompt_len} tokens)\n\
uncached: {} tokens in {dt_uncached:.2}s = {:.2} tok/s\n\
cached: {} tokens in {dt_cached:.2}s = {:.2} tok/s ({:.1}x)\n\
text: {:?}\n===============================\n",
uncached.len(),
uncached.len() as f64 / dt_uncached.max(1e-9),
cached.len(),
cached.len() as f64 / dt_cached.max(1e-9),
dt_uncached / dt_cached.max(1e-9),
tok.decode(&cached),
);
assert_eq!(
uncached, cached,
"greedy token sequences must match between the uncached and KV-cached paths"
);
assert!(!cached.is_empty(), "no tokens generated");
}
fn load_real_model(
device: &burn::backend::wgpu::WgpuDevice,
) -> (
super::super::gemma::GemmaModel<super::super::LocalBackend>,
super::super::tokenizer::GemmaTokenizer,
) {
let dir = std::env::var("GEMMA_DIR")
.expect("set GEMMA_DIR to a folder with model.safetensors + tokenizer.json");
let weights = std::fs::read(format!("{dir}/model.safetensors")).expect("read weights");
let tok_bytes =
std::fs::read(format!("{dir}/tokenizer.json")).expect("read tokenizer.json");
let model = super::super::gemma::GemmaModel::<super::super::LocalBackend>::init(
super::super::gemma::GemmaConfig::gemma_3_270m(),
device,
);
let model =
super::super::weights::load_gemma(model, &weights, device).expect("load_gemma");
let tok = super::super::tokenizer::GemmaTokenizer::from_bytes(&tok_bytes)
.expect("load tokenizer");
(model, tok)
}
async fn greedy_uncached(
model: &super::super::gemma::GemmaModel<super::super::LocalBackend>,
mut tokens: Vec<i64>,
max_new: usize,
device: &burn::backend::wgpu::WgpuDevice,
) -> Vec<i64> {
let mut out = Vec::new();
for _ in 0..max_new {
if at_context_limit(tokens.len()) {
break;
}
let seq = tokens.len();
let input = Tensor::<super::super::LocalBackend, 1, Int>::from_data(
TensorData::from(tokens.as_slice()),
device,
)
.reshape([1, seq]);
let data = model.forward(input).argmax(2).into_data_async().await.expect("read-back");
let next = data.iter::<i64>().last().expect("nonempty");
if next == EOS_ID {
break;
}
tokens.push(next);
out.push(next);
}
out
}
async fn greedy_cached(
model: &super::super::gemma::GemmaModel<super::super::LocalBackend>,
prompt_tokens: Vec<i64>,
max_new: usize,
device: &burn::backend::wgpu::WgpuDevice,
) -> Vec<i64> {
let mut cache = model.new_cache();
let mut pending = prompt_tokens;
let mut total_len = pending.len();
let mut out = Vec::new();
for _ in 0..max_new {
if at_context_limit(total_len) {
break;
}
let seq = pending.len();
let input = Tensor::<super::super::LocalBackend, 1, Int>::from_data(
TensorData::from(pending.as_slice()),
device,
)
.reshape([1, seq]);
let data = model
.forward_cached(input, &mut cache)
.argmax(2)
.into_data_async()
.await
.expect("read-back");
let next = data.iter::<i64>().last().expect("nonempty");
if next == EOS_ID {
break;
}
pending = vec![next];
total_len += 1;
out.push(next);
}
out
}
#[tokio::test]
#[ignore]
async fn gemma_sliding_window_parity() {
let device = burn::backend::wgpu::WgpuDevice::default();
let (model, tok) = load_real_model(&device);
let w = super::super::gemma::GemmaConfig::gemma_3_270m().sliding_window;
let mut prompt = String::from("A guide to the cities of Europe. ");
for i in 0..45 {
prompt.push_str(&format!(
"City number {i} has a river, a market square, an old stone bridge, and {i} towers. "
));
}
let tokens = tok.encode(&prompt);
assert!(
tokens.len() > w + 32,
"prompt must cross the sliding window: {} <= {}",
tokens.len(),
w + 32
);
let max_new: usize = std::env::var("GEMMA_MAX_NEW")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(24);
let t0 = std::time::Instant::now();
let uncached = greedy_uncached(&model, tokens.clone(), max_new, &device).await;
let dt0 = t0.elapsed().as_secs_f64();
let t1 = std::time::Instant::now();
let cached = greedy_cached(&model, tokens.clone(), max_new, &device).await;
let dt1 = t1.elapsed().as_secs_f64();
println!(
"\n=== GEMMA SLIDING-WINDOW PARITY ===\nprompt: {} tokens (window {w})\n\
uncached: {} tokens in {dt0:.2}s · cached: {} tokens in {dt1:.2}s\n\
text: {:?}\n===================================\n",
tokens.len(),
uncached.len(),
cached.len(),
tok.decode(&cached),
);
assert_eq!(
uncached, cached,
"greedy sequences must match across the window boundary"
);
assert!(!cached.is_empty(), "no tokens generated");
}
#[test]
fn emitter_holdback_respects_char_boundaries() {
let mut e = StreamEmitter::new();
let (d, _) = e.step("aééé", false);
assert_eq!(d, Some("a"));
let (d, _) = e.step("aééé", true);
assert_eq!(d, Some("ééé"));
}
}