#!/usr/bin/env perl

use strict;
use warnings;
use FindBin qw($Bin);
use lib "$Bin/../lib";
use Encode qw(encode);
use Getopt::Long qw(GetOptions);
use JSON::PP ();
use LaptopNN::Network;
use LaptopNN::TextVectorizer qw(read_labeled_text_csv);
use LaptopNN::Data qw(encode_labels csv_quote);
use LaptopNN::Metrics qw(classification_metrics format_classification_report);

binmode STDOUT, ':encoding(UTF-8)';
binmode STDERR, ':encoding(UTF-8)';
my %opt = (json => 0);
GetOptions(
    'model=s' => \$opt{model}, 'data=s' => \$opt{data}, 'json!' => \$opt{json},
    'output=s' => \$opt{output}, 'predictions=s' => \$opt{predictions},
    'help|h' => \$opt{help},
) or _usage(2);
_usage(0) if $opt{help};
_usage(2, '--model is required') unless defined $opt{model};
_usage(2, '--data is required') unless defined $opt{data};

my ($network, $model) = LaptopNN::Network->load_model($opt{model});
die "model is not a text intent classifier\n"
    unless ($model->{model_type} // '') eq 'text_intent_classifier';
die "model is missing labels or vectorizer\n"
    unless ref($model->{labels}) eq 'ARRAY' && @{$model->{labels}}
        && ref($model->{text_vectorizer}) eq 'HASH';
my $dataset = read_labeled_text_csv(path => $opt{data});
my $vectorizer = LaptopNN::TextVectorizer->from_hash($model->{text_vectorizer});
my $x = $vectorizer->transform_many($dataset->{texts});
my $truth = encode_labels($dataset->{labels}, $model->{labels});
my @probabilities = map { $network->predict_proba($_) } @$x;
my $metrics = classification_metrics($truth, \@probabilities, $model->{labels});
_write_predictions(
    $opt{predictions}, $dataset->{texts}, $dataset->{labels},
    $metrics->{predicted_indices}, \@probabilities, $model->{labels}
) if defined $opt{predictions};
my %public = %$metrics;
delete $public{predicted_indices};
my $report = $opt{json}
    ? JSON::PP->new->canonical(1)->pretty(1)->encode(\%public)
    : format_classification_report($metrics);
if (defined($opt{output}) && $opt{output} ne '-') {
    open my $fh, '>:raw', $opt{output} or die "cannot write '$opt{output}': $!\n";
    print {$fh} encode('UTF-8', $report);
    close $fh or die "cannot close '$opt{output}': $!\n";
} else {
    print $report;
}

sub _write_predictions {
    my ($path, $texts, $actual, $predicted, $probabilities, $labels) = @_;
    open my $fh, '>:raw', $path or die "cannot write predictions '$path': $!\n";
    my @header = ('row', 'utterance', 'actual_intent', 'predicted_intent', 'correct', 'confidence');
    push @header, map { "prob:$_" } @$labels;
    print {$fh} encode('UTF-8', join(',', map { csv_quote($_) } @header) . "\n");
    for my $row (0 .. $#$actual) {
        my $prediction = $predicted->[$row];
        my @fields = (
            $row + 1, $texts->[$row], $actual->[$row], $labels->[$prediction],
            $actual->[$row] eq $labels->[$prediction] ? 1 : 0,
            sprintf('%.10g', $probabilities->[$row][$prediction]),
            map { sprintf('%.10g', $_) } @{$probabilities->[$row]},
        );
        print {$fh} encode('UTF-8', join(',', map { csv_quote($_) } @fields) . "\n");
    }
    close $fh or die "cannot close predictions '$path': $!\n";
}

sub _usage {
    my ($status, $message) = @_;
    print STDERR "$message\n\n" if defined $message;
    print STDERR "Usage: plnn-text-evaluate --model model.json --data benchmark.csv [--json]\n";
    exit $status;
}
