use std::error::Error;
use std::io::{self, BufRead, Write};
use std::path::PathBuf;
use std::str::FromStr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use clap::{Args, Parser, Subcommand};
use litsea::version;
use litsea::{AdaBoost, AveragedPerceptron, Extractor, Language, PosTrainer, Segmenter, Trainer};
#[derive(Debug, Args)]
#[command(about = "Extract features from a corpus")]
struct ExtractArgs {
#[arg(short, long, default_value = "japanese", value_parser = Language::from_str)]
language: Language,
#[arg(long)]
pos: bool,
corpus_file: PathBuf,
features_file: PathBuf,
}
#[derive(Debug, Args)]
#[command(about = "Train a segmenter")]
struct TrainArgs {
#[arg(short, long, default_value = "0.01")]
threshold: f64,
#[arg(short = 'i', long, default_value = "100")]
num_iterations: usize,
#[arg(short = 'm', long)]
load_model_uri: Option<String>,
#[arg(long)]
pos: bool,
#[arg(long, default_value = "10")]
num_epochs: usize,
features_file: PathBuf,
model_file: PathBuf,
}
#[derive(Debug, Args)]
#[command(about = "Segment a sentence")]
struct SegmentArgs {
#[arg(short, long, default_value = "japanese", value_parser = Language::from_str)]
language: Language,
#[arg(long)]
pos: bool,
model_uri: String,
}
#[derive(Debug, Subcommand)]
enum Commands {
Extract(ExtractArgs),
Train(TrainArgs),
Segment(SegmentArgs),
}
#[derive(Debug, Parser)]
#[command(
name = "litsea",
author,
about = "A morphological analysis command line interface",
version = version(),
propagate_version = true,
)]
struct CommandArgs {
#[command(subcommand)]
command: Commands,
}
fn extract(args: ExtractArgs) -> Result<(), Box<dyn Error>> {
let extractor = Extractor::new(args.language);
if args.pos {
extractor.extract_with_pos(args.corpus_file.as_path(), args.features_file.as_path())?;
} else {
extractor.extract(args.corpus_file.as_path(), args.features_file.as_path())?;
}
eprintln!("Feature extraction completed successfully.");
Ok(())
}
async fn train(args: TrainArgs) -> Result<(), Box<dyn Error>> {
let running = Arc::new(AtomicBool::new(true));
let r = running.clone();
ctrlc::set_handler(move || {
if r.load(Ordering::SeqCst) {
r.store(false, Ordering::SeqCst);
} else {
std::process::exit(0);
}
})?;
if args.pos {
let mut trainer = PosTrainer::new(args.num_epochs, args.features_file.as_path())?;
if let Some(model_uri) = &args.load_model_uri {
trainer.load_model(model_uri).await?;
}
let metrics = trainer.train(&running, args.model_file.as_path())?;
eprintln!("Result Metrics (POS):");
eprintln!(" Accuracy: {:.2}% ( {} )", metrics.accuracy, metrics.num_instances);
eprintln!(" Macro Precision: {:.2}%", metrics.macro_precision);
eprintln!(" Macro Recall: {:.2}%", metrics.macro_recall);
} else {
let mut trainer =
Trainer::new(args.threshold, args.num_iterations, args.features_file.as_path())?;
if let Some(model_uri) = &args.load_model_uri {
trainer.load_model(model_uri).await?;
}
let metrics = trainer.train(&running, args.model_file.as_path())?;
eprintln!("Result Metrics:");
eprintln!(
" Accuracy: {:.2}% ( {} / {} )",
metrics.accuracy,
metrics.true_positives + metrics.true_negatives,
metrics.num_instances
);
eprintln!(
" Precision: {:.2}% ( {} / {} )",
metrics.precision,
metrics.true_positives,
metrics.true_positives + metrics.false_positives
);
eprintln!(
" Recall: {:.2}% ( {} / {} )",
metrics.recall,
metrics.true_positives,
metrics.true_positives + metrics.false_negatives
);
eprintln!(
" Confusion Matrix:\n True Positives: {}\n False Positives: {}\n False Negatives: {}\n True Negatives: {}",
metrics.true_positives,
metrics.false_positives,
metrics.false_negatives,
metrics.true_negatives
);
}
Ok(())
}
fn write_output_line<W: Write>(writer: &mut W, line: &str) -> io::Result<bool> {
match writeln!(writer, "{}", line) {
Ok(()) => Ok(true),
Err(e) if e.kind() == io::ErrorKind::BrokenPipe => Ok(false),
Err(e) => Err(e),
}
}
fn flush_output<W: Write>(writer: &mut W) -> io::Result<()> {
match writer.flush() {
Ok(()) => Ok(()),
Err(e) if e.kind() == io::ErrorKind::BrokenPipe => Ok(()),
Err(e) => Err(e),
}
}
async fn segment(args: SegmentArgs) -> Result<(), Box<dyn Error>> {
let language = args.language;
let stdin = io::stdin();
let stdout = io::stdout();
let mut writer = io::BufWriter::new(stdout.lock());
if args.pos {
let mut pos_learner = AveragedPerceptron::new();
pos_learner.load_model(args.model_uri.as_str()).await?;
let segmenter = Segmenter::with_pos_learner(language, pos_learner);
for line in stdin.lock().lines() {
let line = line?;
let line = line.trim();
if line.is_empty() {
continue;
}
let tokens = segmenter.segment_with_pos(line)?;
let formatted: Vec<String> =
tokens.iter().map(|(word, pos)| format!("{}/{}", word, pos)).collect();
if !write_output_line(&mut writer, &formatted.join(" "))? {
return Ok(());
}
}
} else {
let mut learner = AdaBoost::new(0.01, 100);
learner.load_model(args.model_uri.as_str()).await?;
let segmenter = Segmenter::with_learner(language, learner);
for line in stdin.lock().lines() {
let line = line?;
let line = line.trim();
if line.is_empty() {
continue;
}
let tokens = segmenter.segment(line);
if !write_output_line(&mut writer, &tokens.join(" "))? {
return Ok(());
}
}
}
flush_output(&mut writer)?;
Ok(())
}
async fn run() -> Result<(), Box<dyn std::error::Error>> {
let args = CommandArgs::parse();
match args.command {
Commands::Extract(args) => extract(args),
Commands::Train(args) => train(args).await,
Commands::Segment(args) => segment(args).await,
}
}
#[tokio::main]
async fn main() {
if let Err(e) = run().await {
eprintln!("Error: {}", e);
std::process::exit(1);
}
}