use std::env;
use std::error::Error;
use std::fs;
use std::time::Instant;
use rml_core::{prepare_training_data, NGramModel};
fn main() -> Result<(), Box<dyn Error>> {
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];
let epochs = if args.len() > 3 {
args[3].parse::<usize>().unwrap_or(5)
} else {
5
};
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());
}
};
println!("Preparing training data...");
let training_data = prepare_training_data(&text);
println!("Number of training examples: {}", training_data.len());
println!("Initializing new model...");
let mut model = NGramModel::new();
println!("Starting training for {} epochs...", epochs);
let start_time = Instant::now();
for epoch in 1..=epochs {
let epoch_start = Instant::now();
for (i, (context, target)) in training_data.iter().enumerate() {
model.train(context, *target);
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);
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(())
}