fn dense_stream_tokens(
state: &AppState,
request: &GenerateRequest,
cancel: &CancelToken,
) -> Result<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>), ApiErr> {
let (model, tokenizer) = state
.get_model(request.model_id.as_deref())
.map_err(|e| api_err(super::model_resolution_status(&e), e))?;
let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
let prompt: Vec<usize> = prompt_ids.iter().map(|&id| id as usize).collect();
let prompt_len = prompt.len();
let strategy = match request.strategy.as_str() {
"greedy" => SamplingStrategy::Greedy,
"top_k" => SamplingStrategy::TopK { k: request.top_k },
"top_p" => SamplingStrategy::TopP { p: request.top_p },
other => {
return Err(api_err(
StatusCode::BAD_REQUEST,
format!("Invalid strategy: {other}"),
))
},
};
let mut config = GenerationConfig::default()
.with_max_tokens(request.max_tokens)
.with_temperature(request.temperature)
.with_cancel(cancel.clone());
config.strategy = strategy;
if let Some(seed) = request.seed {
config = config.with_seed(seed);
}
let generated = model
.generate(&prompt, &config)
.map_err(|e| api_err(StatusCode::INTERNAL_SERVER_ERROR, e))?;
let token_ids: Vec<u32> = generated
.iter()
.map(|&id| {
u32::try_from(id).map_err(|_| {
api_err(
StatusCode::BAD_REQUEST,
format!("Token ID {id} exceeds u32 range"),
)
})
})
.collect::<Result<Vec<_>, _>>()?;
Ok((token_ids, prompt_len, tokenizer))
}
fn try_quantized_stream_tokens(
state: &AppState,
request: &GenerateRequest,
cancel: &CancelToken,
) -> Result<Option<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>)>, ApiErr> {
let quantized_model = match state.quantized_model() {
Some(m) => m,
None => return Ok(None),
};
let tokenizer = require_tok(state)?;
let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
let prompt_len = prompt_ids.len();
let sampling = resolve_quantized_sampling(
&request.strategy,
request.top_k,
request.top_p,
request.temperature,
)?;
let q_config = quantized_config(
state,
&tokenizer,
request.max_tokens,
request.temperature,
&sampling,
request.seed,
cancel,
);
let generated = quantized_model
.generate_with_cache(&prompt_ids, &q_config)
.map_err(|e| generation_err(&e))?;
Ok(Some((generated, prompt_len, tokenizer)))
}
fn try_apr_stream_tokens(
state: &AppState,
request: &GenerateRequest,
cancel: &CancelToken,
) -> Result<Option<(Vec<u32>, usize, std::sync::Arc<BPETokenizer>)>, ApiErr> {
use crate::apr_transformer::GenerateConfig;
let apr_transformer = match state.apr_transformer() {
Some(m) => m,
None => return Ok(None),
};
let tokenizer = require_tok(state)?;
let prompt_ids = tokenize_prompt(&tokenizer, &request.prompt)?;
let prompt_len = prompt_ids.len();
let gen_config = GenerateConfig {
max_tokens: request.max_tokens,
temperature: request.temperature,
cancel: cancel.clone(),
..Default::default()
};
let generated = apr_transformer
.generate_with_cache(&prompt_ids, &gen_config)
.map_err(|e| {
api_err(
super::generation_error_status(&e),
format!("APR generation failed: {e}"),
)
})?;
Ok(Some((generated, prompt_len, tokenizer)))
}
pub async fn stream_generate_handler(
State(state): State<AppState>,
Extension(cancel): Extension<CancelToken>,
Json(request): Json<GenerateRequest>,
) -> Result<Sse<impl Stream<Item = Result<Event, Infallible>>>, (StatusCode, Json<ErrorResponse>)> {
let (token_ids, prompt_len, tokenizer_clone) =
if let Some(resolved) = try_quantized_stream_tokens(&state, &request, &cancel)? {
resolved
} else if let Some(resolved) = try_apr_stream_tokens(&state, &request, &cancel)? {
resolved
} else {
dense_stream_tokens(&state, &request, &cancel)?
};
let stream = async_stream::stream! {
let generated_start = prompt_len.min(token_ids.len());
for &token_id in &token_ids[generated_start..] {
let text = match tokenizer_clone.decode(&[token_id]) {
Ok(t) => t,
Err(_) => String::from("<error>"),
};
let event = StreamTokenEvent { token_id, text };
let data = serde_json::to_string(&event)
.unwrap_or_else(|_| r#"{"error":"serialization failed"}"#.to_string());
yield Ok::<_, Infallible>(Event::default().event("token").data(data));
}
let done_event = StreamDoneEvent {
num_generated: token_ids.len().saturating_sub(prompt_len),
};
let data = serde_json::to_string(&done_event)
.unwrap_or_else(|_| r#"{"error":"serialization failed"}"#.to_string());
yield Ok(Event::default().event("done").data(data));
};
Ok(Sse::new(stream))
}
#[cfg(test)]
#[path = "gpu_handlers_tests.rs"]
mod gpu_handlers_tests;