use std::env;
use std::error::Error;
use rml_core::{char_to_index, index_to_char, sample, NGramModel, ALLOWED_CHARS};
const CONTEXT_SIZE: usize = 4;
fn main() -> Result<(), Box<dyn Error>> {
let args: Vec<String> = env::args().collect();
if args.len() < 3 {
eprintln!("Usage: {} <model_path> <seed_text> [length]", args[0]);
return Err("Insufficient arguments".into());
}
let model_path = &args[1];
let seed_text = &args[2];
let generation_length = if args.len() > 3 {
args[3].parse::<usize>().unwrap_or(100)
} else {
100
};
println!("Loading model from {}", model_path);
let mut model = match NGramModel::load(model_path) {
Ok(model) => model,
Err(e) => {
eprintln!("Error loading the model: {}", e);
return Err(e.into());
}
};
if seed_text.len() < CONTEXT_SIZE {
return Err(format!(
"Seed text must be at least {} characters long",
CONTEXT_SIZE
)
.into());
}
println!("Generating text based on: '{}'", seed_text);
let generated_text = generate_text(&mut model, seed_text, generation_length)?;
println!("\nGenerated text:");
println!("{}", generated_text);
Ok(())
}
fn generate_text(
model: &mut NGramModel,
seed: &str,
length: usize,
) -> Result<String, Box<dyn Error>> {
let filtered_seed: String = seed
.chars()
.filter(|&c| ALLOWED_CHARS.contains(c))
.collect();
if filtered_seed.len() < CONTEXT_SIZE {
return Err(format!(
"After filtering, the seed text is too short. Need at least {} characters.",
CONTEXT_SIZE
)
.into());
}
let mut context: Vec<usize> = filtered_seed[filtered_seed.len() - CONTEXT_SIZE..]
.chars()
.map(char_to_index)
.collect();
if context.len() != CONTEXT_SIZE {
return Err(format!(
"Context must be exactly {} characters, but is {}",
CONTEXT_SIZE,
context.len()
)
.into());
}
let mut result = filtered_seed.clone();
for _ in 0..length {
let probs = model.forward(&context);
let next_char_idx = sample(&probs);
let next_char = index_to_char(next_char_idx);
result.push(next_char);
context.remove(0);
context.push(next_char_idx);
}
Ok(result)
}