rml-core 0.1.0

A simple N-gram language model implementation in Rust
Documentation
// generate.rs
use std::env;
use std::error::Error;

// We import the functions and structures from lib.rs
use rml_core::{char_to_index, index_to_char, sample, NGramModel, ALLOWED_CHARS};

// We define the context size here, as it is private in lib.rs
const CONTEXT_SIZE: usize = 4;

fn main() -> Result<(), Box<dyn Error>> {
    // Process arguments
    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];

    // Default length for generation is 100 characters
    let generation_length = if args.len() > 3 {
        args[3].parse::<usize>().unwrap_or(100)
    } else {
        100
    };

    // Load model
    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());
        }
    };

    // Check if the seed text is long enough
    if seed_text.len() < CONTEXT_SIZE {
        return Err(format!(
            "Seed text must be at least {} characters long",
            CONTEXT_SIZE
        )
        .into());
    }

    // Generate the text
    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>> {
    // Only use allowed characters from the seed
    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());
    }

    // Convert the seed text to indices
    let mut context: Vec<usize> = filtered_seed[filtered_seed.len() - CONTEXT_SIZE..]
        .chars()
        .map(char_to_index)
        .collect();

    // Check if the context has the right length
    if context.len() != CONTEXT_SIZE {
        return Err(format!(
            "Context must be exactly {} characters, but is {}",
            CONTEXT_SIZE,
            context.len()
        )
        .into());
    }

    // The generated text starts with the seed
    let mut result = filtered_seed.clone();

    // Generate text
    for _ in 0..length {
        // Perform forward pass
        let probs = model.forward(&context);

        // Sample next character
        let next_char_idx = sample(&probs);
        let next_char = index_to_char(next_char_idx);

        // Add to the result
        result.push(next_char);

        // Update context (remove the oldest character and add the new one)
        context.remove(0);
        context.push(next_char_idx);
    }

    Ok(result)
}