Skip to main content

autoagents_burn/model/llama/generation/
generate.rs

1use super::super::{tokenizer::Tokenizer, Llama};
2use super::{GenerationContext, Sampler};
3use crate::model::llama::generation::stream_sender::StreamSender;
4use burn::{prelude::*, tensor::activation::softmax};
5use log::debug;
6
7pub(crate) fn temperature_scaled_softmax<B: Backend>(
8    logits: Tensor<B, 2>,
9    temperature: f64,
10) -> Tensor<B, 2> {
11    softmax(logits / temperature, 1)
12}
13
14/// Generated text sample output.
15pub struct GenerationOutput {
16    /// The number of generated tokens.
17    pub tokens: usize,
18    /// The time it took to produce the output tokens (generation + decoding).
19    // pub time: std::time::Duration,
20    pub result: String,
21}
22
23#[derive(Debug)]
24pub enum GenerationError {
25    MaxSequenceLengthExceeded { actual: usize, max: usize },
26}
27
28impl<B: Backend, T: Tokenizer + 'static> Llama<B, T> {
29    /// Generate text sample based on the provided prompt.
30    ///
31    /// # Arguments
32    /// - `prompt`: The prompt string to use for generating the samples.
33    /// - `sample_len`: The number of new tokens to generate (i.e., the number of generation steps to take).
34    /// - `temperature`: Temperature value for controlling randomness in sampling (scales logits by `1 / temperature`).
35    ///   High values result in more random sampling.
36    /// - `sampler`: The sampling strategy to use when selecting the next token based on the predicted probabilities.
37    ///
38    /// # Returns
39    /// The generated text along with some other metadata (see [GenerationOutput]).
40    pub async fn generate(
41        &mut self,
42        prompt: &str,
43        sample_len: usize,
44        temperature: f64,
45        sampler: &mut Sampler,
46        emitter: Option<StreamSender>,
47    ) -> Result<GenerationOutput, GenerationError> {
48        let input_tokens = self.tokenize(prompt);
49        let prompt_len = input_tokens.dims()[0];
50
51        let mut state = GenerationContext::<B, T>::new(
52            prompt_len + sample_len,
53            self.tokenizer.clone(),
54            &self.device,
55            emitter,
56        )
57        .await;
58        state.append(input_tokens);
59
60        let mut input_pos = Tensor::<B, 1, Int>::arange(0..prompt_len as i64, &self.device);
61
62        debug!("Starting Generation Loop");
63        for i in 0..sample_len {
64            debug!("Generation Loop Iter: {i}");
65            if state.should_stop() {
66                break;
67            }
68
69            let x = state
70                .tokens
71                .clone()
72                .select(0, input_pos.clone())
73                .reshape([1, -1]);
74
75            let [_, seq_len] = x.dims();
76
77            // Prepare cache and RoPE for current sequence length and position
78            let mask = self.cache.prepare(seq_len)?;
79            debug!("Prepared Cache");
80            self.pos_encoding.prepare(seq_len);
81            debug!("Prepared Positional Encoding");
82
83            let logits = self
84                .model
85                .forward(x, &mut self.cache, &self.pos_encoding, mask);
86
87            debug!("Model Forwad Pass Completed");
88
89            let [batch_size, seq_len, _vocab_size] = logits.dims();
90            let mut next_token_logits = logits
91                .slice([0..batch_size, seq_len - 1..seq_len])
92                .squeeze_dim(1); // [batch_size=1, vocab_size]
93
94            if temperature > 0.0 {
95                next_token_logits = temperature_scaled_softmax(next_token_logits, temperature);
96            };
97
98            debug!("Sampling Tokens");
99            let next_token = sampler.sample(next_token_logits).await.squeeze_dim(0);
100
101            // Update with the new generated token
102            state.update(next_token.clone()).await;
103            debug!("Update Tokens Complete");
104
105            // Advance
106            let t = input_pos.dims()[0];
107            input_pos = input_pos.slice(t - 1..t) + 1;
108        }
109        debug!("Generation Loop Exited");
110
111        let num_tokens = state.num_tokens_generated();
112
113        // Decode the generated tokens to text
114        let generated_text = state.get_generated_text().await;
115        debug!("Generated Text Extracted");
116
117        Ok(GenerationOutput {
118            tokens: num_tokens,
119            result: generated_text,
120        })
121    }
122}