Start the generation loop and call the stream function for every token. `max_tokens` overrides the default sample length if provided.
(
&mut self,
max_tokens: Option<usize>,
mut stream: S,
)
| 107 | /// Start the generation loop and call the stream function for every token. |
| 108 | /// `max_tokens` overrides the default sample length if provided. |
| 109 | pub async fn generate_text<S>( |
| 110 | &mut self, |
| 111 | max_tokens: Option<usize>, |
| 112 | mut stream: S, |
| 113 | ) -> Result<()> |
| 114 | where |
| 115 | S: FnMut(&str), |
| 116 | { |
| 117 | log::info!( |
| 118 | "starting the inference loop (mem={})\n\n", |
| 119 | human_bytes::human_bytes(memory_stats::memory_stats().unwrap().physical_mem as f64) |
| 120 | ); |
| 121 | |
| 122 | let sample_len = max_tokens.unwrap_or(self.ctx.args.sample_len); |
| 123 | log::debug!(" sample_len = {}", sample_len); |
| 124 | |
| 125 | let mut start_gen = std::time::Instant::now(); |
| 126 | let model = self |
| 127 | .model |
| 128 | .as_mut() |
| 129 | .ok_or_else(|| anyhow::anyhow!("No text model loaded"))?; |
| 130 | |
| 131 | for index in 0..sample_len { |
| 132 | if index == 1 { |
| 133 | start_gen = std::time::Instant::now() |
| 134 | } |
| 135 | |
| 136 | let token_start = std::time::Instant::now(); |
| 137 | let token = model.next_token(index).await?; |
| 138 | let token_elapsed = token_start.elapsed(); |
| 139 | |
| 140 | log::debug!( |
| 141 | "token {} generated in {:.1}ms ({:.1} tok/s)", |
| 142 | index, |
| 143 | token_elapsed.as_secs_f64() * 1000.0, |
| 144 | 1.0 / token_elapsed.as_secs_f64(), |
| 145 | ); |
| 146 | |
| 147 | if token.is_end_of_stream { |
| 148 | break; |
| 149 | } else { |
| 150 | stream(&token.to_string()); |
| 151 | // Yield to the runtime so the SSE stream task can flush |
| 152 | // this token to the client before we start the next one. |
| 153 | tokio::task::yield_now().await; |
| 154 | } |
| 155 | } |
| 156 | |
| 157 | // signal end of stream |
| 158 | stream(""); |
| 159 | |
| 160 | let dt = start_gen.elapsed(); |
| 161 | let generated = model.generated_tokens(); |
| 162 | |
| 163 | log::info!( |
| 164 | "{} tokens generated ({:.2} token/s) - mem={}", |
| 165 | generated, |
| 166 | (generated - 1) as f64 / dt.as_secs_f64(), |
no test coverage detected