676 lines
22 KiB
Rust
676 lines
22 KiB
Rust
//! Streaming text generation with token-by-token output and real-time processing
|
|
|
|
use crate::{GenerationConfig, GenerationOutput, ModelInterface, Result, generation::utils};
|
|
use futures::Stream;
|
|
use rtx_tensor::{Device, Tensor};
|
|
use std::collections::VecDeque;
|
|
use std::sync::Arc;
|
|
use std::time::{Duration, Instant};
|
|
use tokio::sync::mpsc;
|
|
|
|
/// Streaming text generator for real-time token-by-token generation
|
|
pub struct StreamingGenerator {
|
|
model: Arc<dyn ModelInterface>,
|
|
config: GenerationConfig,
|
|
streaming_config: StreamingConfig,
|
|
}
|
|
|
|
/// Configuration for streaming generation
|
|
#[derive(Debug, Clone)]
|
|
pub struct StreamingConfig {
|
|
/// Buffer size for streaming tokens
|
|
pub buffer_size: usize,
|
|
/// Maximum latency per token in milliseconds
|
|
pub max_token_latency_ms: u64,
|
|
/// Whether to stream partial tokens (subword pieces)
|
|
pub stream_partial_tokens: bool,
|
|
/// Minimum token confidence to stream
|
|
pub min_confidence: f32,
|
|
/// Whether to include timing information
|
|
pub include_timing: bool,
|
|
/// Whether to include confidence scores
|
|
pub include_confidence: bool,
|
|
/// Batch size for streaming (1 for pure token-by-token)
|
|
pub batch_size: usize,
|
|
/// Whether to use speculative decoding for faster streaming
|
|
pub use_speculative_decoding: bool,
|
|
}
|
|
|
|
impl Default for StreamingConfig {
|
|
fn default() -> Self {
|
|
Self {
|
|
buffer_size: 100,
|
|
max_token_latency_ms: 100,
|
|
stream_partial_tokens: false,
|
|
min_confidence: 0.0,
|
|
include_timing: true,
|
|
include_confidence: false,
|
|
batch_size: 1,
|
|
use_speculative_decoding: false,
|
|
}
|
|
}
|
|
}
|
|
|
|
impl StreamingGenerator {
|
|
/// Create streaming generator with model and generation config
|
|
pub fn with_model(model: Arc<dyn ModelInterface>, config: GenerationConfig) -> Result<Self> {
|
|
Self::with_streaming_config(model, config, StreamingConfig::default())
|
|
}
|
|
|
|
/// Create streaming generator with custom streaming configuration
|
|
pub fn with_streaming_config(
|
|
model: Arc<dyn ModelInterface>,
|
|
config: GenerationConfig,
|
|
streaming_config: StreamingConfig,
|
|
) -> Result<Self> {
|
|
config.validate()?;
|
|
Ok(Self {
|
|
model,
|
|
config,
|
|
streaming_config,
|
|
})
|
|
}
|
|
|
|
/// Generate streaming text with real-time token output
|
|
pub async fn generate_stream(
|
|
&self,
|
|
prompt: &str,
|
|
) -> Result<impl Stream<Item = Result<StreamingToken>>> {
|
|
let (tx, rx) = mpsc::channel(self.streaming_config.buffer_size);
|
|
let model = self.model.clone();
|
|
let config = self.config.clone();
|
|
let streaming_config = self.streaming_config.clone();
|
|
let prompt = prompt.to_string();
|
|
|
|
// Spawn background task for generation
|
|
tokio::spawn(async move {
|
|
if let Err(e) =
|
|
Self::generate_streaming_impl(model, &prompt, config, streaming_config, tx).await
|
|
{
|
|
tracing::error!("Streaming generation failed: {}", e);
|
|
}
|
|
});
|
|
|
|
Ok(tokio_stream::wrappers::ReceiverStream::new(rx))
|
|
}
|
|
|
|
/// Generate multiple streaming completions
|
|
pub async fn generate_multiple_streams(
|
|
&self,
|
|
prompt: &str,
|
|
num_streams: usize,
|
|
) -> Result<Vec<tokio_stream::wrappers::ReceiverStream<Result<StreamingToken>>>> {
|
|
let mut streams = Vec::new();
|
|
let prompt_owned = prompt.to_string();
|
|
|
|
for i in 0..num_streams {
|
|
let seed = self.config.seed.map(|s| s + i as u64);
|
|
let mut modified_config = self.config.clone();
|
|
modified_config.seed = seed;
|
|
|
|
// Clone everything needed for the spawned task
|
|
let model = self.model.clone();
|
|
let streaming_config = self.streaming_config.clone();
|
|
let prompt_clone = prompt_owned.clone();
|
|
|
|
let (tx, rx) = mpsc::channel(self.streaming_config.buffer_size);
|
|
|
|
tokio::spawn(async move {
|
|
if let Err(e) = Self::generate_streaming_impl(
|
|
model,
|
|
&prompt_clone,
|
|
modified_config,
|
|
streaming_config,
|
|
tx,
|
|
)
|
|
.await
|
|
{
|
|
tracing::error!("Streaming generation failed: {}", e);
|
|
}
|
|
});
|
|
|
|
streams.push(tokio_stream::wrappers::ReceiverStream::new(rx));
|
|
}
|
|
|
|
Ok(streams)
|
|
}
|
|
|
|
/// Internal streaming implementation
|
|
async fn generate_streaming_impl(
|
|
model: Arc<dyn ModelInterface>,
|
|
prompt: &str,
|
|
config: GenerationConfig,
|
|
streaming_config: StreamingConfig,
|
|
tx: mpsc::Sender<Result<StreamingToken>>,
|
|
) -> Result<()> {
|
|
let start_time = Instant::now();
|
|
let tokenizer = model.tokenizer();
|
|
let special_tokens = tokenizer.special_tokens();
|
|
|
|
// Tokenize input
|
|
let input_tokens = tokenizer.encode(prompt)?;
|
|
let mut current_sequence = input_tokens.clone();
|
|
|
|
let max_length = config.max_length.unwrap_or(100);
|
|
let max_new_tokens = max_length.saturating_sub(input_tokens.len());
|
|
|
|
// Send initial metadata token
|
|
let metadata_token = StreamingToken {
|
|
text: String::new(),
|
|
token_id: 0,
|
|
position: 0,
|
|
confidence: 1.0,
|
|
timing: TokenTiming {
|
|
generation_time_ms: 0.0,
|
|
cumulative_time_ms: 0.0,
|
|
tokens_per_second: 0.0,
|
|
},
|
|
metadata: Some(StreamingMetadata {
|
|
prompt_tokens: input_tokens.len(),
|
|
max_new_tokens,
|
|
is_start: true,
|
|
is_end: false,
|
|
finish_reason: None,
|
|
}),
|
|
token_type: TokenType::Metadata,
|
|
};
|
|
|
|
if let Err(_) = tx.send(Ok(metadata_token)).await {
|
|
return Ok(()); // Receiver dropped
|
|
}
|
|
|
|
// Main generation loop
|
|
let mut generated_tokens = 0;
|
|
let mut last_token_time = start_time;
|
|
|
|
while generated_tokens < max_new_tokens {
|
|
let token_start = Instant::now();
|
|
|
|
// Create input tensor for current sequence
|
|
let sequence_data: Vec<f32> = current_sequence.iter().map(|&x| x as f32).collect();
|
|
let input_tensor = Tensor::from_data(
|
|
sequence_data,
|
|
[1, current_sequence.len()],
|
|
&Device::default(),
|
|
)?;
|
|
|
|
// Get model predictions (single forward pass)
|
|
let model_output = if streaming_config.use_speculative_decoding {
|
|
// Use speculative decoding for faster generation
|
|
generate_with_speculative_decoding(&*model, &input_tensor, &config).await?
|
|
} else {
|
|
model.generate_tokens(&input_tensor, None, &config)?
|
|
};
|
|
|
|
// Get logits for next token
|
|
let logits = get_next_token_logits(&model_output, &input_tensor)?;
|
|
|
|
// Apply generation strategy (temperature, top-p, top-k, etc.)
|
|
let (next_token, confidence) =
|
|
sample_next_token_with_confidence(&logits, &config, &streaming_config)?;
|
|
|
|
// Check if we should stop generation
|
|
if utils::should_stop_generation(¤t_sequence, &config, special_tokens) {
|
|
break;
|
|
}
|
|
|
|
// Check confidence threshold
|
|
if confidence < streaming_config.min_confidence {
|
|
break;
|
|
}
|
|
|
|
// Add token to sequence
|
|
current_sequence.push(next_token);
|
|
generated_tokens += 1;
|
|
|
|
// Decode token text
|
|
let token_text = if streaming_config.stream_partial_tokens {
|
|
// Stream individual token pieces
|
|
tokenizer.decode(&[next_token])?
|
|
} else {
|
|
// Stream full token only
|
|
let decoded = tokenizer.decode(&[next_token])?;
|
|
if is_complete_token(&decoded) {
|
|
decoded
|
|
} else {
|
|
String::new() // Skip partial tokens
|
|
}
|
|
};
|
|
|
|
let token_end = Instant::now();
|
|
let token_time = token_end.duration_since(token_start);
|
|
let cumulative_time = token_end.duration_since(start_time);
|
|
|
|
// Calculate timing metrics
|
|
let timing = TokenTiming {
|
|
generation_time_ms: token_time.as_millis() as f64,
|
|
cumulative_time_ms: cumulative_time.as_millis() as f64,
|
|
tokens_per_second: generated_tokens as f64 / cumulative_time.as_secs_f64(),
|
|
};
|
|
|
|
// Create streaming token
|
|
let streaming_token = StreamingToken {
|
|
text: token_text,
|
|
token_id: next_token,
|
|
position: input_tokens.len() + generated_tokens - 1,
|
|
confidence: if streaming_config.include_confidence {
|
|
confidence
|
|
} else {
|
|
1.0
|
|
},
|
|
timing: if streaming_config.include_timing {
|
|
timing
|
|
} else {
|
|
TokenTiming::default()
|
|
},
|
|
metadata: None,
|
|
token_type: TokenType::Generated,
|
|
};
|
|
|
|
// Send token if not empty or if it's a special token
|
|
if (!streaming_token.text.is_empty() || is_special_token(next_token, special_tokens))
|
|
&& let Err(_) = tx.send(Ok(streaming_token)).await
|
|
{
|
|
break; // Receiver dropped
|
|
}
|
|
|
|
// Rate limiting - ensure we don't exceed max token latency
|
|
let elapsed_since_last = token_end.duration_since(last_token_time);
|
|
let min_interval = Duration::from_millis(streaming_config.max_token_latency_ms);
|
|
|
|
if elapsed_since_last < min_interval {
|
|
tokio::time::sleep(min_interval.checked_sub(elapsed_since_last).unwrap()).await;
|
|
}
|
|
|
|
last_token_time = Instant::now();
|
|
|
|
// Check for early stopping conditions
|
|
if next_token == special_tokens.eos_token {
|
|
break;
|
|
}
|
|
|
|
// Check timeout
|
|
if let Some(timeout_ms) = config.timeout_ms
|
|
&& cumulative_time.as_millis() > timeout_ms as u128
|
|
{
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Send final metadata token
|
|
let final_time = Instant::now().duration_since(start_time);
|
|
let final_metadata = StreamingToken {
|
|
text: String::new(),
|
|
token_id: 0,
|
|
position: current_sequence.len(),
|
|
confidence: 1.0,
|
|
timing: TokenTiming {
|
|
generation_time_ms: final_time.as_millis() as f64,
|
|
cumulative_time_ms: final_time.as_millis() as f64,
|
|
tokens_per_second: generated_tokens as f64 / final_time.as_secs_f64(),
|
|
},
|
|
metadata: Some(StreamingMetadata {
|
|
prompt_tokens: input_tokens.len(),
|
|
max_new_tokens,
|
|
is_start: false,
|
|
is_end: true,
|
|
finish_reason: Some(determine_finish_reason(
|
|
¤t_sequence,
|
|
&config,
|
|
special_tokens,
|
|
)),
|
|
}),
|
|
token_type: TokenType::Metadata,
|
|
};
|
|
|
|
let _ = tx.send(Ok(final_metadata)).await;
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
/// Streaming token with timing and confidence information
|
|
#[derive(Debug, Clone)]
|
|
pub struct StreamingToken {
|
|
/// Generated text for this token
|
|
pub text: String,
|
|
/// Token ID
|
|
pub token_id: u32,
|
|
/// Position in the sequence
|
|
pub position: usize,
|
|
/// Confidence score for this token
|
|
pub confidence: f32,
|
|
/// Timing information
|
|
pub timing: TokenTiming,
|
|
/// Metadata (for start/end tokens)
|
|
pub metadata: Option<StreamingMetadata>,
|
|
/// Type of token
|
|
pub token_type: TokenType,
|
|
}
|
|
|
|
/// Timing information for each token
|
|
#[derive(Debug, Clone, Default)]
|
|
pub struct TokenTiming {
|
|
/// Time to generate this token in milliseconds
|
|
pub generation_time_ms: f64,
|
|
/// Cumulative generation time in milliseconds
|
|
pub cumulative_time_ms: f64,
|
|
/// Current generation speed in tokens per second
|
|
pub tokens_per_second: f64,
|
|
}
|
|
|
|
/// Metadata for streaming session
|
|
#[derive(Debug, Clone)]
|
|
pub struct StreamingMetadata {
|
|
/// Number of prompt tokens
|
|
pub prompt_tokens: usize,
|
|
/// Maximum new tokens to generate
|
|
pub max_new_tokens: usize,
|
|
/// Whether this is the start of generation
|
|
pub is_start: bool,
|
|
/// Whether this is the end of generation
|
|
pub is_end: bool,
|
|
/// Reason for finishing generation
|
|
pub finish_reason: Option<crate::FinishReason>,
|
|
}
|
|
|
|
/// Type of streaming token
|
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
|
pub enum TokenType {
|
|
/// Regular generated token
|
|
Generated,
|
|
/// Metadata token (start/end)
|
|
Metadata,
|
|
/// Special token (EOS, etc.)
|
|
Special,
|
|
}
|
|
|
|
/// Buffer for accumulating partial tokens
|
|
#[derive(Debug)]
|
|
pub struct TokenBuffer {
|
|
buffer: VecDeque<String>,
|
|
max_size: usize,
|
|
}
|
|
|
|
impl TokenBuffer {
|
|
pub fn new(max_size: usize) -> Self {
|
|
Self {
|
|
buffer: VecDeque::with_capacity(max_size),
|
|
max_size,
|
|
}
|
|
}
|
|
|
|
pub fn add_token(&mut self, token: String) {
|
|
if self.buffer.len() >= self.max_size {
|
|
self.buffer.pop_front();
|
|
}
|
|
self.buffer.push_back(token);
|
|
}
|
|
|
|
pub fn get_text(&self) -> String {
|
|
self.buffer.iter().cloned().collect::<String>()
|
|
}
|
|
|
|
pub fn clear(&mut self) {
|
|
self.buffer.clear();
|
|
}
|
|
}
|
|
|
|
// Helper functions
|
|
|
|
fn get_next_token_logits(_output: &GenerationOutput, _input: &Tensor) -> Result<Tensor> {
|
|
// Mock implementation - in real scenario would extract logits from model output
|
|
let vocab_size = 50000;
|
|
|
|
// Create streaming-optimized logits with quick token selection characteristics
|
|
let mut logits_vec = Vec::with_capacity(vocab_size);
|
|
for i in 0..vocab_size {
|
|
let base_logit = rand::random::<f32>() * 2.5 - 1.25; // Random between -1.25 and 1.25
|
|
// Streaming favors more common tokens for faster generation
|
|
let frequency_boost = if i < 2000 {
|
|
1.0
|
|
} else if i < 10000 {
|
|
0.3
|
|
} else {
|
|
-0.5
|
|
};
|
|
logits_vec.push(base_logit + frequency_boost);
|
|
}
|
|
|
|
let logits = Tensor::from_data(logits_vec, [vocab_size], &Device::default())?;
|
|
Ok(logits)
|
|
}
|
|
|
|
fn sample_next_token_with_confidence(
|
|
logits: &Tensor,
|
|
config: &GenerationConfig,
|
|
_streaming_config: &StreamingConfig,
|
|
) -> Result<(u32, f32)> {
|
|
let mut modified_logits = logits.clone();
|
|
|
|
// Apply temperature
|
|
utils::apply_temperature(&mut modified_logits, config.temperature)?;
|
|
|
|
// Apply sampling based on generation strategy
|
|
match &config.strategy {
|
|
crate::generation::GenerationStrategy::NucleusSampling(nucleus_config) => {
|
|
utils::apply_top_p(&mut modified_logits, nucleus_config.top_p)?;
|
|
}
|
|
crate::generation::GenerationStrategy::TopKSampling(topk_config) => {
|
|
utils::apply_top_k(&mut modified_logits, topk_config.top_k)?;
|
|
}
|
|
_ => {} // Use raw probabilities for other strategies
|
|
}
|
|
|
|
// Sample token and calculate confidence
|
|
let mut rng = rand::thread_rng();
|
|
let next_token = utils::sample_categorical(&modified_logits, &mut rng)?;
|
|
|
|
// Calculate confidence as the probability of the selected token
|
|
let probs = modified_logits.softmax(-1)?;
|
|
let confidence = probs.get(&[next_token as usize])?;
|
|
|
|
Ok((next_token, confidence))
|
|
}
|
|
|
|
async fn generate_with_speculative_decoding(
|
|
model: &dyn ModelInterface,
|
|
input: &Tensor,
|
|
config: &GenerationConfig,
|
|
) -> Result<GenerationOutput> {
|
|
// Simplified speculative decoding - would implement draft model + verification
|
|
model.generate_tokens(input, None, config)
|
|
}
|
|
|
|
fn is_complete_token(text: &str) -> bool {
|
|
// Heuristic to determine if a token is complete (not a subword piece)
|
|
!text.starts_with("##") && !text.is_empty()
|
|
}
|
|
|
|
fn is_special_token(token_id: u32, special_tokens: &crate::SpecialTokens) -> bool {
|
|
token_id == special_tokens.eos_token
|
|
|| token_id == special_tokens.bos_token
|
|
|| token_id == special_tokens.pad_token
|
|
|| token_id == special_tokens.unk_token
|
|
|| (special_tokens.sep_token == Some(token_id))
|
|
|| (special_tokens.cls_token == Some(token_id))
|
|
|| (special_tokens.mask_token == Some(token_id))
|
|
}
|
|
|
|
fn determine_finish_reason(
|
|
sequence: &[u32],
|
|
config: &GenerationConfig,
|
|
special_tokens: &crate::SpecialTokens,
|
|
) -> crate::FinishReason {
|
|
if let Some(&last_token) = sequence.last()
|
|
&& last_token == special_tokens.eos_token
|
|
{
|
|
return crate::FinishReason::EosToken;
|
|
}
|
|
|
|
if let Some(max_length) = config.max_length
|
|
&& sequence.len() >= max_length
|
|
{
|
|
return crate::FinishReason::MaxLength;
|
|
}
|
|
|
|
crate::FinishReason::MaxLength
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use crate::{GenerationConfig, MockModelInterface};
|
|
use futures::StreamExt;
|
|
|
|
#[tokio::test]
|
|
async fn test_streaming_generation() -> Result<()> {
|
|
let model = Arc::new(MockModelInterface::new("/tmp/mock")?);
|
|
let config = GenerationConfig::default().max_length(10);
|
|
let generator = StreamingGenerator::with_model(model, config)?;
|
|
|
|
let mut stream = generator.generate_stream("test prompt").await?;
|
|
let mut tokens = Vec::new();
|
|
|
|
while let Some(result) = stream.next().await {
|
|
tokens.push(result?);
|
|
}
|
|
|
|
assert!(!tokens.is_empty());
|
|
|
|
// Check for metadata tokens
|
|
let start_metadata = tokens.iter().find(|t| {
|
|
matches!(t.token_type, TokenType::Metadata)
|
|
&& t.metadata.as_ref().map_or(false, |m| m.is_start)
|
|
});
|
|
let end_metadata = tokens.iter().find(|t| {
|
|
matches!(t.token_type, TokenType::Metadata)
|
|
&& t.metadata.as_ref().map_or(false, |m| m.is_end)
|
|
});
|
|
|
|
assert!(start_metadata.is_some());
|
|
assert!(end_metadata.is_some());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_streaming_with_custom_config() -> Result<()> {
|
|
let model = Arc::new(MockModelInterface::new("/tmp/mock")?);
|
|
let config = GenerationConfig::default().max_length(5);
|
|
let streaming_config = StreamingConfig {
|
|
buffer_size: 50,
|
|
max_token_latency_ms: 10,
|
|
include_timing: true,
|
|
include_confidence: true,
|
|
..Default::default()
|
|
};
|
|
|
|
let generator = StreamingGenerator::with_streaming_config(model, config, streaming_config)?;
|
|
let mut stream = generator.generate_stream("hello").await?;
|
|
let mut tokens = Vec::new();
|
|
|
|
while let Some(result) = stream.next().await {
|
|
tokens.push(result?);
|
|
}
|
|
|
|
assert!(!tokens.is_empty());
|
|
|
|
// Check that timing and confidence are included
|
|
let generated_tokens: Vec<_> = tokens
|
|
.iter()
|
|
.filter(|t| matches!(t.token_type, TokenType::Generated))
|
|
.collect();
|
|
|
|
for token in generated_tokens {
|
|
assert!(token.timing.generation_time_ms >= 0.0);
|
|
assert!(token.confidence >= 0.0);
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_multiple_streams() -> Result<()> {
|
|
let model = Arc::new(MockModelInterface::new("/tmp/mock")?);
|
|
let config = GenerationConfig::default().max_length(5).seed(123);
|
|
let generator = StreamingGenerator::with_model(model, config)?;
|
|
|
|
let streams = generator.generate_multiple_streams("test", 2).await?;
|
|
assert_eq!(streams.len(), 2);
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_token_buffer() {
|
|
let mut buffer = TokenBuffer::new(3);
|
|
|
|
buffer.add_token("hello".to_string());
|
|
buffer.add_token(" ".to_string());
|
|
buffer.add_token("world".to_string());
|
|
|
|
assert_eq!(buffer.get_text(), "hello world");
|
|
|
|
buffer.add_token("!".to_string());
|
|
assert_eq!(buffer.get_text(), " world!"); // First token evicted
|
|
|
|
buffer.clear();
|
|
assert_eq!(buffer.get_text(), "");
|
|
}
|
|
|
|
#[test]
|
|
fn test_is_complete_token() {
|
|
assert!(is_complete_token("hello"));
|
|
assert!(is_complete_token("world"));
|
|
assert!(!is_complete_token("##ing")); // Subword piece
|
|
assert!(!is_complete_token("")); // Empty
|
|
}
|
|
|
|
#[test]
|
|
fn test_is_special_token() {
|
|
let special_tokens = crate::SpecialTokens {
|
|
pad_token: 0,
|
|
eos_token: 2,
|
|
bos_token: 1,
|
|
unk_token: 3,
|
|
sep_token: Some(4),
|
|
cls_token: Some(5),
|
|
mask_token: Some(6),
|
|
};
|
|
|
|
assert!(is_special_token(0, &special_tokens)); // PAD
|
|
assert!(is_special_token(2, &special_tokens)); // EOS
|
|
assert!(is_special_token(4, &special_tokens)); // SEP
|
|
assert!(!is_special_token(100, &special_tokens)); // Regular token
|
|
}
|
|
|
|
#[test]
|
|
fn test_determine_finish_reason() {
|
|
let special_tokens = crate::SpecialTokens {
|
|
pad_token: 0,
|
|
eos_token: 2,
|
|
bos_token: 1,
|
|
unk_token: 3,
|
|
sep_token: None,
|
|
cls_token: None,
|
|
mask_token: None,
|
|
};
|
|
|
|
let config = GenerationConfig::default().max_length(5);
|
|
|
|
// EOS token finish
|
|
let eos_sequence = vec![1, 10, 20, 2];
|
|
assert!(matches!(
|
|
determine_finish_reason(&eos_sequence, &config, &special_tokens),
|
|
crate::FinishReason::EosToken
|
|
));
|
|
|
|
// Max length finish
|
|
let long_sequence = vec![1, 10, 20, 30, 40];
|
|
assert!(matches!(
|
|
determine_finish_reason(&long_sequence, &config, &special_tokens),
|
|
crate::FinishReason::MaxLength
|
|
));
|
|
}
|
|
}
|