{- | Bundled, pretty-printing evaluation summaries: a labelled confusion matrix
and scikit-learn-style regression / classification reports. The @*Expr@ variants
take a model's prediction expression and a truth column directly, so a full
report is a one-liner after fitting.
-}
module DataFrame.Metrics.Report (
    ConfusionMatrix (..),
    confusionMatrix,
    confusionMatrixExpr,
    RegressionReport (..),
    regressionReport,
    regressionReportExpr,
    ClassStats (..),
    ClassificationReport (..),
    classificationReport,
    classificationReportExpr,
) where

import Data.List (nub, sort, sortBy)
import Data.Ord (Down (..), comparing)
import qualified Data.Vector.Unboxed as VU

import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)
import DataFrame.Metrics (
    Average (..),
    accuracy,
    classCounts,
    columnOf,
    f1,
    f1Of,
    mae,
    mse,
    precOf,
    r2,
    recOf,
    rmse,
 )

-- | A labelled confusion matrix: class order plus row-major @actual×predicted@.
data ConfusionMatrix = ConfusionMatrix
    { ConfusionMatrix -> [Double]
cmClasses :: ![Double]
    , ConfusionMatrix -> [[Int]]
cmCounts :: ![[Int]]
    }
    deriving (ConfusionMatrix -> ConfusionMatrix -> Bool
(ConfusionMatrix -> ConfusionMatrix -> Bool)
-> (ConfusionMatrix -> ConfusionMatrix -> Bool)
-> Eq ConfusionMatrix
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: ConfusionMatrix -> ConfusionMatrix -> Bool
== :: ConfusionMatrix -> ConfusionMatrix -> Bool
$c/= :: ConfusionMatrix -> ConfusionMatrix -> Bool
/= :: ConfusionMatrix -> ConfusionMatrix -> Bool
Eq)

-- | Confusion matrix over the class set of @truth ∪ preds@.
confusionMatrix :: VU.Vector Double -> VU.Vector Double -> ConfusionMatrix
confusionMatrix :: Vector Double -> Vector Double -> ConfusionMatrix
confusionMatrix Vector Double
preds Vector Double
truth = [Double] -> [[Int]] -> ConfusionMatrix
ConfusionMatrix [Double]
classes [[Int]]
counts
  where
    classes :: [Double]
classes = [Double] -> [Double]
forall a. Ord a => [a] -> [a]
sort ([Double] -> [Double]
forall a. Eq a => [a] -> [a]
nub (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
truth [Double] -> [Double] -> [Double]
forall a. [a] -> [a] -> [a]
++ Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
preds))
    counts :: [[Int]]
counts =
        [ [ Vector Bool -> Int
forall a. Unbox a => Vector a -> Int
VU.length ((Bool -> Bool) -> Vector Bool -> Vector Bool
forall a. Unbox a => (a -> Bool) -> Vector a -> Vector a
VU.filter Bool -> Bool
forall a. a -> a
id ((Double -> Double -> Bool)
-> Vector Double -> Vector Double -> Vector Bool
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith (\Double
p Double
t -> Double
t Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
a Bool -> Bool -> Bool
&& Double
p Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c) Vector Double
preds Vector Double
truth))
          | Double
c <- [Double]
classes
          ]
        | Double
a <- [Double]
classes
        ]

-- | Confusion matrix from a prediction expression and a truth column.
confusionMatrixExpr ::
    Expr Double -> Expr Double -> DataFrame -> ConfusionMatrix
confusionMatrixExpr :: Expr Double -> Expr Double -> DataFrame -> ConfusionMatrix
confusionMatrixExpr Expr Double
predExpr Expr Double
truthExpr DataFrame
df =
    Vector Double -> Vector Double -> ConfusionMatrix
confusionMatrix (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
predExpr) (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
truthExpr)

instance Show ConfusionMatrix where
    show :: ConfusionMatrix -> String
show (ConfusionMatrix [Double]
classes [[Int]]
counts) =
        [String] -> String
unlines (String
header String -> [String] -> [String]
forall a. a -> [a] -> [a]
: [String]
rows)
      where
        lbls :: [String]
lbls = (Double -> String) -> [Double] -> [String]
forall a b. (a -> b) -> [a] -> [b]
map Double -> String
forall a. Show a => a -> String
show [Double]
classes
        w :: Int
w = [Int] -> Int
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum (Int
8 Int -> [Int] -> [Int]
forall a. a -> [a] -> [a]
: (String -> Int) -> [String] -> [Int]
forall a b. (a -> b) -> [a] -> [b]
map String -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [String]
lbls) Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
2
        cell :: ShowS
cell String
s = Int -> Char -> String
forall a. Int -> a -> [a]
replicate (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
- String -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length String
s)) Char
' ' String -> ShowS
forall a. [a] -> [a] -> [a]
++ String
s
        header :: String
header = ShowS
cell String
"a\\p" String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS -> [String] -> String
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap ShowS
cell [String]
lbls
        rows :: [String]
rows = [ShowS
cell String
a String -> ShowS
forall a. [a] -> [a] -> [a]
++ (Int -> String) -> [Int] -> String
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap (ShowS
cell ShowS -> (Int -> String) -> Int -> String
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Int -> String
forall a. Show a => a -> String
show) [Int]
row | (String
a, [Int]
row) <- [String] -> [[Int]] -> [(String, [Int])]
forall a b. [a] -> [b] -> [(a, b)]
zip [String]
lbls [[Int]]
counts]

-- | Regression metrics bundle.
data RegressionReport = RegressionReport
    { RegressionReport -> Double
rrMSE :: !Double
    , RegressionReport -> Double
rrRMSE :: !Double
    , RegressionReport -> Double
rrMAE :: !Double
    , RegressionReport -> Double
rrR2 :: !Double
    }
    deriving (RegressionReport -> RegressionReport -> Bool
(RegressionReport -> RegressionReport -> Bool)
-> (RegressionReport -> RegressionReport -> Bool)
-> Eq RegressionReport
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: RegressionReport -> RegressionReport -> Bool
== :: RegressionReport -> RegressionReport -> Bool
$c/= :: RegressionReport -> RegressionReport -> Bool
/= :: RegressionReport -> RegressionReport -> Bool
Eq)

instance Show RegressionReport where
    show :: RegressionReport -> String
show RegressionReport
r =
        [String] -> String
unlines
            [ String
"Regression report"
            , String
"  mse  = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall a. Show a => a -> String
show (RegressionReport -> Double
rrMSE RegressionReport
r)
            , String
"  rmse = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall a. Show a => a -> String
show (RegressionReport -> Double
rrRMSE RegressionReport
r)
            , String
"  mae  = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall a. Show a => a -> String
show (RegressionReport -> Double
rrMAE RegressionReport
r)
            , String
"  r2   = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall a. Show a => a -> String
show (RegressionReport -> Double
rrR2 RegressionReport
r)
            ]

-- | Regression report from prediction/truth vectors.
regressionReport :: VU.Vector Double -> VU.Vector Double -> RegressionReport
regressionReport :: Vector Double -> Vector Double -> RegressionReport
regressionReport Vector Double
preds Vector Double
truth =
    Double -> Double -> Double -> Double -> RegressionReport
RegressionReport
        (Metric
mse Vector Double
preds Vector Double
truth)
        (Metric
rmse Vector Double
preds Vector Double
truth)
        (Metric
mae Vector Double
preds Vector Double
truth)
        (Metric
r2 Vector Double
preds Vector Double
truth)

-- | Regression report from a prediction expression and a truth column.
regressionReportExpr ::
    Expr Double -> Expr Double -> DataFrame -> RegressionReport
regressionReportExpr :: Expr Double -> Expr Double -> DataFrame -> RegressionReport
regressionReportExpr Expr Double
predExpr Expr Double
truthExpr DataFrame
df =
    Vector Double -> Vector Double -> RegressionReport
regressionReport (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
predExpr) (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
truthExpr)

-- | Per-class precision/recall/F1/support.
data ClassStats = ClassStats
    { ClassStats -> Double
csPrecision :: !Double
    , ClassStats -> Double
csRecall :: !Double
    , ClassStats -> Double
csF1 :: !Double
    , ClassStats -> Int
csSupport :: !Int
    }
    deriving (ClassStats -> ClassStats -> Bool
(ClassStats -> ClassStats -> Bool)
-> (ClassStats -> ClassStats -> Bool) -> Eq ClassStats
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: ClassStats -> ClassStats -> Bool
== :: ClassStats -> ClassStats -> Bool
$c/= :: ClassStats -> ClassStats -> Bool
/= :: ClassStats -> ClassStats -> Bool
Eq, Int -> ClassStats -> ShowS
[ClassStats] -> ShowS
ClassStats -> String
(Int -> ClassStats -> ShowS)
-> (ClassStats -> String)
-> ([ClassStats] -> ShowS)
-> Show ClassStats
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> ClassStats -> ShowS
showsPrec :: Int -> ClassStats -> ShowS
$cshow :: ClassStats -> String
show :: ClassStats -> String
$cshowList :: [ClassStats] -> ShowS
showList :: [ClassStats] -> ShowS
Show)

{- | A scikit-learn-style classification report: per-class stats plus accuracy
and macro/weighted F1.
-}
data ClassificationReport = ClassificationReport
    { ClassificationReport -> [(Double, ClassStats)]
crPerClass :: ![(Double, ClassStats)]
    , ClassificationReport -> Double
crAccuracy :: !Double
    , ClassificationReport -> Double
crMacroF1 :: !Double
    , ClassificationReport -> Double
crWeightedF1 :: !Double
    }
    deriving (ClassificationReport -> ClassificationReport -> Bool
(ClassificationReport -> ClassificationReport -> Bool)
-> (ClassificationReport -> ClassificationReport -> Bool)
-> Eq ClassificationReport
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: ClassificationReport -> ClassificationReport -> Bool
== :: ClassificationReport -> ClassificationReport -> Bool
$c/= :: ClassificationReport -> ClassificationReport -> Bool
/= :: ClassificationReport -> ClassificationReport -> Bool
Eq)

instance Show ClassificationReport where
    show :: ClassificationReport -> String
show ClassificationReport
r =
        [String] -> String
unlines ([String] -> String) -> [String] -> String
forall a b. (a -> b) -> a -> b
$
            (ShowS
pad String
"class" String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad String
"precision" String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad String
"recall" String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad String
"f1" String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad String
"support")
                String -> [String] -> [String]
forall a. a -> [a] -> [a]
: [ ShowS
pad (Double -> String
forall a. Show a => a -> String
show Double
c)
                        String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad (Double -> String
forall {a}. RealFrac a => a -> String
num (ClassStats -> Double
csPrecision ClassStats
s))
                        String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad (Double -> String
forall {a}. RealFrac a => a -> String
num (ClassStats -> Double
csRecall ClassStats
s))
                        String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad (Double -> String
forall {a}. RealFrac a => a -> String
num (ClassStats -> Double
csF1 ClassStats
s))
                        String -> ShowS
forall a. [a] -> [a] -> [a]
++ ShowS
pad (Int -> String
forall a. Show a => a -> String
show (ClassStats -> Int
csSupport ClassStats
s))
                  | (Double
c, ClassStats
s) <- ClassificationReport -> [(Double, ClassStats)]
crPerClass ClassificationReport
r
                  ]
                [String] -> [String] -> [String]
forall a. [a] -> [a] -> [a]
++ [ String
""
                   , String
"accuracy    = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall {a}. RealFrac a => a -> String
num (ClassificationReport -> Double
crAccuracy ClassificationReport
r)
                   , String
"macro f1    = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall {a}. RealFrac a => a -> String
num (ClassificationReport -> Double
crMacroF1 ClassificationReport
r)
                   , String
"weighted f1 = " String -> ShowS
forall a. [a] -> [a] -> [a]
++ Double -> String
forall {a}. RealFrac a => a -> String
num (ClassificationReport -> Double
crWeightedF1 ClassificationReport
r)
                   ]
      where
        pad :: ShowS
pad String
s = let w :: Int
w = Int
12 in String
s String -> ShowS
forall a. [a] -> [a] -> [a]
++ Int -> Char -> String
forall a. Int -> a -> [a]
replicate (Int -> Int -> Int
forall a. Ord a => a -> a -> a
max Int
1 (Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
- String -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length String
s)) Char
' '
        num :: a -> String
num a
x = Double -> String
forall a. Show a => a -> String
show (Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (a -> Int
forall b. Integral b => a -> b
forall a b. (RealFrac a, Integral b) => a -> b
round (a
x a -> a -> a
forall a. Num a => a -> a -> a
* a
1000) :: Int) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
1000 :: Double)

-- | Classification report from prediction/truth vectors.
classificationReport ::
    VU.Vector Double -> VU.Vector Double -> ClassificationReport
classificationReport :: Vector Double -> Vector Double -> ClassificationReport
classificationReport Vector Double
preds Vector Double
truth =
    [(Double, ClassStats)]
-> Double -> Double -> Double -> ClassificationReport
ClassificationReport
        [(Double, ClassStats)]
perClass
        (Metric
accuracy Vector Double
preds Vector Double
truth)
        (Average -> Metric
f1 Average
Macro Vector Double
preds Vector Double
truth)
        (Average -> Metric
f1 Average
Weighted Vector Double
preds Vector Double
truth)
  where
    perClass :: [(Double, ClassStats)]
perClass =
        ((Double, ClassStats) -> (Double, ClassStats) -> Ordering)
-> [(Double, ClassStats)] -> [(Double, ClassStats)]
forall a. (a -> a -> Ordering) -> [a] -> [a]
sortBy (((Double, ClassStats) -> Down Int)
-> (Double, ClassStats) -> (Double, ClassStats) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (Int -> Down Int
forall a. a -> Down a
Down (Int -> Down Int)
-> ((Double, ClassStats) -> Int)
-> (Double, ClassStats)
-> Down Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ClassStats -> Int
csSupport (ClassStats -> Int)
-> ((Double, ClassStats) -> ClassStats)
-> (Double, ClassStats)
-> Int
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (Double, ClassStats) -> ClassStats
forall a b. (a, b) -> b
snd)) ([(Double, ClassStats)] -> [(Double, ClassStats)])
-> [(Double, ClassStats)] -> [(Double, ClassStats)]
forall a b. (a -> b) -> a -> b
$
            [ (Double
c, Double -> Double -> Double -> Int -> ClassStats
ClassStats ((Int, Int, Int, Int) -> Double
precOf (Int, Int, Int, Int)
s) ((Int, Int, Int, Int) -> Double
recOf (Int, Int, Int, Int)
s) ((Int, Int, Int, Int) -> Double
f1Of (Int, Int, Int, Int)
s) ((Int, Int, Int, Int) -> Int
forall {a} {b} {c} {d}. (a, b, c, d) -> d
supOf (Int, Int, Int, Int)
s))
            | (Double
c, (Int, Int, Int, Int)
s) <- Vector Double -> Vector Double -> [(Double, (Int, Int, Int, Int))]
classCounts Vector Double
preds Vector Double
truth
            ]
    supOf :: (a, b, c, d) -> d
supOf (a
_, b
_, c
_, d
sup) = d
sup

-- | Classification report from a prediction expression and a truth column.
classificationReportExpr ::
    Expr Double -> Expr Double -> DataFrame -> ClassificationReport
classificationReportExpr :: Expr Double -> Expr Double -> DataFrame -> ClassificationReport
classificationReportExpr Expr Double
predExpr Expr Double
truthExpr DataFrame
df =
    Vector Double -> Vector Double -> ClassificationReport
classificationReport (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
predExpr) (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
truthExpr)