rml-core 0.1.0

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

// Import functions from lib.rs
use rml_core::{prepare_training_data, NGramModel};

fn main() -> Result<(), Box<dyn Error>> {
    // Process arguments
    let args: Vec<String> = env::args().collect();
    if args.len() < 3 {
        eprintln!("Usage: {} <text_file> <output_model> [epochs]", args[0]);
        return Err("Insufficient arguments".into());
    }

    let input_file = &args[1];
    let output_model = &args[2];

    // Default number of epochs is 5
    let epochs = if args.len() > 3 {
        args[3].parse::<usize>().unwrap_or(5)
    } else {
        5
    };

    // Load text from file
    println!("Loading text from {}", input_file);
    let text = match fs::read_to_string(input_file) {
        Ok(content) => content,
        Err(e) => {
            eprintln!("Error reading the file: {}", e);
            return Err(e.into());
        }
    };

    // Prepare training data
    println!("Preparing training data...");
    let training_data = prepare_training_data(&text);
    println!("Number of training examples: {}", training_data.len());

    // Create new model
    println!("Initializing new model...");
    let mut model = NGramModel::new();

    // Training
    println!("Starting training for {} epochs...", epochs);
    let start_time = Instant::now();

    for epoch in 1..=epochs {
        let epoch_start = Instant::now();

        // Go through each example
        for (i, (context, target)) in training_data.iter().enumerate() {
            model.train(context, *target);

            // Show progress (every 10000 examples)
            if i % 10000 == 0 {
                print!(
                    "\rEpoch {}/{}: {:.2}% complete",
                    epoch,
                    epochs,
                    (i as f32 / training_data.len() as f32) * 100.0
                );
            }
        }

        let epoch_duration = epoch_start.elapsed();
        println!(
            "\rEpoch {}/{} completed in {:.2?}",
            epoch, epochs, epoch_duration
        );
    }

    let total_duration = start_time.elapsed();
    println!("Training completed in {:.2?}", total_duration);

    // Save model
    println!("Saving model to {}...", output_model);
    match model.save(output_model) {
        Ok(_) => println!("Model successfully saved!"),
        Err(e) => {
            eprintln!("Error saving the model: {}", e);
            return Err(e.into());
        }
    }

    Ok(())
}