|
1 | | -use std::{ |
2 | | - path::Path, |
3 | | - time::{Instant, SystemTime, UNIX_EPOCH}, |
4 | | -}; |
| 1 | +use std::{path::Path, time::Instant}; |
5 | 2 |
|
6 | 3 | use anyhow::{Context, Result}; |
7 | 4 | use log::info; |
8 | | -use openinfer_engine::engine::{ |
9 | | - EngineHandle, EngineLoadOptions, FinishReason, GenerateRequest, TokenEvent, TokenSink, |
10 | | -}; |
| 5 | +use openinfer_engine::engine::{EngineHandle, EngineLoadOptions}; |
11 | 6 | use tokio::sync::mpsc; |
12 | 7 |
|
13 | | -use crate::runtime::{DeepSeekV2LiteEp2Generator, GenerationResult}; |
| 8 | +use crate::{runtime::DeepSeekV2LiteEp2Generator, scheduler::MixedRequestScheduler}; |
14 | 9 |
|
15 | 10 | pub(crate) fn start_engine(model_path: &Path, options: EngineLoadOptions) -> Result<EngineHandle> { |
16 | 11 | let started = Instant::now(); |
17 | 12 | info!("starting DeepSeek-V2-Lite EP2 engine"); |
18 | | - let mut generator = DeepSeekV2LiteEp2Generator::load(model_path, options)?; |
19 | | - let (submit_tx, mut submit_rx) = mpsc::unbounded_channel(); |
| 13 | + let generator = DeepSeekV2LiteEp2Generator::load(model_path, options)?; |
| 14 | + let servable_len = generator.config().supported_plain_rope_context() as u32; |
| 15 | + let (submit_tx, submit_rx) = mpsc::unbounded_channel(); |
20 | 16 |
|
21 | 17 | let join_handle = std::thread::Builder::new() |
22 | 18 | .name("deepseek-v2-lite-ep2".to_string()) |
23 | | - .spawn(move || { |
24 | | - while let Some(req) = submit_rx.blocking_recv() { |
25 | | - handle_request(&mut generator, &req); |
26 | | - } |
27 | | - }) |
| 19 | + .spawn(move || MixedRequestScheduler::new(generator, submit_rx).run()) |
28 | 20 | .context("spawn DeepSeek-V2-Lite EP=2 engine thread")?; |
29 | 21 |
|
30 | 22 | info!( |
31 | 23 | "DeepSeek-V2-Lite EP2 engine started cost {:.2}s", |
32 | 24 | started.elapsed().as_secs_f64() |
33 | 25 | ); |
34 | | - Ok(EngineHandle::new_with_join_handle(submit_tx, join_handle)) |
35 | | -} |
36 | | - |
37 | | -fn handle_request(generator: &mut DeepSeekV2LiteEp2Generator, req: &GenerateRequest) { |
38 | | - let prompt_tokens = req.prompt_tokens.len(); |
39 | | - let now = unix_time_secs(); |
40 | | - let _ = req.token_tx.send(TokenEvent::Scheduled { |
41 | | - queued_at_unix_s: req.queued_at_unix_s.unwrap_or(now), |
42 | | - scheduled_at_unix_s: now, |
43 | | - prompt_tokens, |
44 | | - cached_tokens: 0, |
45 | | - }); |
46 | | - if req.echo { |
47 | | - let _ = req.token_tx.send(TokenEvent::PromptTokens { |
48 | | - ids: req.prompt_tokens.clone(), |
49 | | - logprobs: vec![None; prompt_tokens], |
50 | | - }); |
51 | | - } |
52 | | - if !req.params.is_greedy() { |
53 | | - reject_request( |
54 | | - req, |
55 | | - prompt_tokens, |
56 | | - format!( |
57 | | - "DeepSeek-V2-Lite EP=2 first gate serves greedy decoding only; requested temperature={}, top_k={}, top_p={}", |
58 | | - req.params.temperature, req.params.top_k, req.params.top_p |
59 | | - ), |
60 | | - ); |
61 | | - return; |
62 | | - } |
63 | | - if req.logprobs > 0 { |
64 | | - reject_request( |
65 | | - req, |
66 | | - prompt_tokens, |
67 | | - "DeepSeek-V2-Lite EP=2 first gate does not return logprobs yet".to_string(), |
68 | | - ); |
69 | | - return; |
70 | | - } |
71 | | - if req.max_tokens == 0 { |
72 | | - let _ = req.token_tx.send(TokenEvent::Finished { |
73 | | - finish_reason: FinishReason::Length, |
74 | | - prompt_tokens, |
75 | | - completion_tokens: 0, |
76 | | - }); |
77 | | - return; |
78 | | - } |
79 | | - |
80 | | - match generator.generate_greedy(&req.prompt_tokens, req.max_tokens, req.params.ignore_eos) { |
81 | | - Ok(result) => { |
82 | | - emit_generation_result(&req.token_tx, prompt_tokens, &result); |
83 | | - } |
84 | | - Err(err) => { |
85 | | - let _ = req.token_tx.send(TokenEvent::Error { |
86 | | - message: err.to_string(), |
87 | | - prompt_tokens, |
88 | | - completion_tokens: 0, |
89 | | - }); |
90 | | - } |
91 | | - } |
92 | | -} |
93 | | - |
94 | | -fn reject_request(req: &GenerateRequest, prompt_tokens: usize, message: String) { |
95 | | - let _ = req.token_tx.send(TokenEvent::Rejected { |
96 | | - message, |
97 | | - prompt_tokens, |
98 | | - completion_tokens: 0, |
99 | | - }); |
100 | | -} |
101 | | - |
102 | | -fn emit_generation_result(token_tx: &TokenSink, prompt_tokens: usize, result: &GenerationResult) { |
103 | | - for token in &result.tokens { |
104 | | - let _ = token_tx.send(TokenEvent::Token { |
105 | | - id: *token, |
106 | | - logprob: None, |
107 | | - }); |
108 | | - } |
109 | | - let _ = token_tx.send(TokenEvent::Finished { |
110 | | - finish_reason: result.finish_reason, |
111 | | - prompt_tokens, |
112 | | - completion_tokens: result.tokens.len(), |
113 | | - }); |
114 | | -} |
115 | | - |
116 | | -fn unix_time_secs() -> f64 { |
117 | | - SystemTime::now() |
118 | | - .duration_since(UNIX_EPOCH) |
119 | | - .map_or(0.0, |duration| duration.as_secs_f64()) |
120 | | -} |
121 | | - |
122 | | -#[cfg(test)] |
123 | | -mod tests { |
124 | | - use super::*; |
125 | | - use crate::runtime::GenerationStats; |
126 | | - |
127 | | - #[test] |
128 | | - fn stop_generation_streams_tokens_and_stop_finish() { |
129 | | - let (tx, mut rx) = TokenSink::standalone(); |
130 | | - |
131 | | - emit_generation_result( |
132 | | - &tx, |
133 | | - 3, |
134 | | - &GenerationResult { |
135 | | - tokens: vec![10, 11], |
136 | | - finish_reason: FinishReason::Stop, |
137 | | - stats: GenerationStats::default(), |
138 | | - }, |
139 | | - ); |
140 | | - drop(tx); |
141 | | - |
142 | | - match rx.try_recv().expect("expected first token").1 { |
143 | | - TokenEvent::Token { id, .. } => assert_eq!(id, 10), |
144 | | - _ => panic!("expected first token event"), |
145 | | - } |
146 | | - match rx.try_recv().expect("expected second token").1 { |
147 | | - TokenEvent::Token { id, .. } => assert_eq!(id, 11), |
148 | | - _ => panic!("expected second token event"), |
149 | | - } |
150 | | - match rx.try_recv().expect("expected finished event").1 { |
151 | | - TokenEvent::Finished { |
152 | | - finish_reason, |
153 | | - prompt_tokens, |
154 | | - completion_tokens, |
155 | | - } => { |
156 | | - assert_eq!(finish_reason, FinishReason::Stop); |
157 | | - assert_eq!(prompt_tokens, 3); |
158 | | - assert_eq!(completion_tokens, 2); |
159 | | - } |
160 | | - _ => panic!("expected finished event"), |
161 | | - } |
162 | | - assert!(rx.try_recv().is_err()); |
163 | | - } |
| 26 | + Ok(EngineHandle::new_with_join_handle(submit_tx, join_handle).with_servable_len(servable_len)) |
164 | 27 | } |
0 commit comments