{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{- | Evaluation metrics for fitted models. The everyday entry point is
'evaluate', which applies a model's prediction expression and a truth column to
a frame and folds a metric — no manual @interpret@/extract plumbing. Metrics are
plain functions (@type Metric = Vector -> Vector -> Double@), so you pass @mse@
or @accuracy@ directly. Classification metrics handle multiclass via 'Average';
'classificationReport' / 'regressionReport' bundle the common numbers with a
scikit-learn-style 'Show'.
-}
module DataFrame.Metrics (
    -- * Metric type + evaluation
    Metric,
    evaluate,
    predictColumn,
    columnOf,

    -- * Regression metrics
    mse,
    rmse,
    mae,
    r2,

    -- * Classification metrics
    accuracy,
    logLoss,
    Average (..),
    precision,
    recall,
    f1,
    rocAuc,

    -- * Per-class helpers (for reports)
    classCounts,
    precOf,
    recOf,
    f1Of,
) where

import Control.Exception (throw)
import Data.Either (fromRight)
import Data.List (nub, sort, sortBy)
import Data.Ord (comparing)
import qualified Data.Text as T
import qualified Data.Vector.Unboxed as VU

import DataFrame.Internal.Column (TypedColumn (..), toVector)
import DataFrame.Internal.DataFrame (DataFrame)
import DataFrame.Internal.Expression (Expr)
import DataFrame.Internal.Interpreter (interpret)
import DataFrame.Operations.Transformations (derive)

-- | A metric maps predictions and ground truth to a scalar score.
type Metric = VU.Vector Double -> VU.Vector Double -> Double

{- | Evaluate a model's prediction expression against a truth column on a frame.

> evaluate rmse (linearExpr model) (F.col @Double "target") df
> evaluate accuracy (logisticDecisionExpr model) (F.col @Double "label") df
-}
evaluate :: Metric -> Expr Double -> Expr Double -> DataFrame -> Double
evaluate :: Metric -> Expr Double -> Expr Double -> DataFrame -> Double
evaluate Metric
metric Expr Double
predExpr Expr Double
truthExpr DataFrame
df =
    Metric
metric (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
predExpr) (DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
truthExpr)

-- | Add a model's prediction expression to a frame as a named column.
predictColumn :: T.Text -> Expr Double -> DataFrame -> DataFrame
predictColumn :: Text -> Expr Double -> DataFrame -> DataFrame
predictColumn = Text -> Expr Double -> DataFrame -> DataFrame
forall a. Columnable a => Text -> Expr a -> DataFrame -> DataFrame
derive

-- | Interpret an expression to a @Double@ vector over a frame.
columnOf :: DataFrame -> Expr Double -> VU.Vector Double
columnOf :: DataFrame -> Expr Double -> Vector Double
columnOf DataFrame
df Expr Double
e = case forall a.
Columnable a =>
DataFrame -> Expr a -> Either DataFrameException (TypedColumn a)
interpret @Double DataFrame
df Expr Double
e of
    Right (TColumn Column
c) -> Vector Double
-> Either DataFrameException (Vector Double) -> Vector Double
forall b a. b -> Either a b -> b
fromRight Vector Double
forall a. Unbox a => Vector a
VU.empty (forall a (v :: * -> *).
(Vector v a, Columnable a) =>
Column -> Either DataFrameException (v a)
toVector @Double @VU.Vector Column
c)
    Left DataFrameException
err -> DataFrameException -> Vector Double
forall a e. Exception e => e -> a
throw DataFrameException
err

n2 :: VU.Vector Double -> Double
n2 :: Vector Double -> Double
n2 = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Double)
-> (Vector Double -> Int) -> Vector Double -> Double
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length

-- | Mean squared error.
mse :: Metric
mse :: Metric
mse Vector Double
preds Vector Double
truth
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
truth = Double
0
    | Bool
otherwise =
        Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
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
p Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
t) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
preds Vector Double
truth) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
n2 Vector Double
truth

-- | Root mean squared error.
rmse :: Metric
rmse :: Metric
rmse Vector Double
preds Vector Double
truth = Double -> Double
forall a. Floating a => a -> a
sqrt (Metric
mse Vector Double
preds Vector Double
truth)

-- | Mean absolute error.
mae :: Metric
mae :: Metric
mae Vector Double
preds Vector Double
truth
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
truth = Double
0
    | Bool
otherwise = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
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 -> Double
forall a. Num a => a -> a
abs (Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
t)) Vector Double
preds Vector Double
truth) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
n2 Vector Double
truth

-- | Coefficient of determination @R²@.
r2 :: Metric
r2 :: Metric
r2 Vector Double
preds Vector Double
truth
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
truth Bool -> Bool -> Bool
|| Double
ssTot Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 = Double
0
    | Bool
otherwise = Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
ssRes Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
ssTot
  where
    mean :: Double
mean = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum Vector Double
truth Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
n2 Vector Double
truth
    ssRes :: Double
ssRes = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
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 -> Double
forall a. Num a => a -> a -> a
- Double
p) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
preds Vector Double
truth)
    ssTot :: Double
ssTot = Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum ((Double -> Double) -> Vector Double -> Vector Double
forall a b. (Unbox a, Unbox b) => (a -> b) -> Vector a -> Vector b
VU.map (\Double
t -> (Double
t Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
mean) Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) Vector Double
truth)

-- | Fraction of exact matches.
accuracy :: Metric
accuracy :: Metric
accuracy Vector Double
preds Vector Double
truth
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
truth = Double
0
    | Bool
otherwise =
        Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (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 -> Double -> Bool
forall a. Eq a => a -> a -> Bool
(==) Vector Double
preds Vector Double
truth))) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
n2 Vector Double
truth

-- | Binary log loss; probabilities clamped away from @0@/@1@.
logLoss :: Metric
logLoss :: Metric
logLoss Vector Double
probs Vector Double
truth
    | Vector Double -> Bool
forall a. Unbox a => Vector a -> Bool
VU.null Vector Double
truth = Double
0
    | Bool
otherwise =
        Double -> Double
forall a. Num a => a -> a
negate
            ( Vector Double -> Double
forall a. (Unbox a, Num a) => Vector a -> a
VU.sum
                ( (Double -> Double -> Double)
-> Vector Double -> Vector Double -> Vector Double
forall a b c.
(Unbox a, Unbox b, Unbox c) =>
(a -> b -> c) -> Vector a -> Vector b -> Vector c
VU.zipWith
                    (\Double
p Double
y -> let q :: Double
q = Double -> Double
forall {a}. (Ord a, Fractional a) => a -> a
clampP Double
p in Double
y Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log Double
q Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
y) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double
forall a. Floating a => a -> a
log (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
q))
                    Vector Double
probs
                    Vector Double
truth
                )
            )
            Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Vector Double -> Double
n2 Vector Double
truth
  where
    clampP :: a -> a
clampP a
p = a -> a -> a
forall a. Ord a => a -> a -> a
max a
1e-15 (a -> a -> a
forall a. Ord a => a -> a -> a
min (a
1 a -> a -> a
forall a. Num a => a -> a -> a
- a
1e-15) a
p)

-- | Averaging strategy for multiclass precision/recall/F1.
data Average
    = -- | one class is positive; the rest negative
      Binary Double
    | -- | unweighted mean over classes
      Macro
    | -- | pool per-class counts (equals accuracy for single-label)
      Micro
    | -- | support-weighted mean over classes
      Weighted
    deriving (Average -> Average -> Bool
(Average -> Average -> Bool)
-> (Average -> Average -> Bool) -> Eq Average
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Average -> Average -> Bool
== :: Average -> Average -> Bool
$c/= :: Average -> Average -> Bool
/= :: Average -> Average -> Bool
Eq, Int -> Average -> ShowS
[Average] -> ShowS
Average -> String
(Int -> Average -> ShowS)
-> (Average -> String) -> ([Average] -> ShowS) -> Show Average
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Average -> ShowS
showsPrec :: Int -> Average -> ShowS
$cshow :: Average -> String
show :: Average -> String
$cshowList :: [Average] -> ShowS
showList :: [Average] -> ShowS
Show)

-- | Per-class @(tp, fp, fn, support)@ over the class set of @truth ∪ preds@.
classCounts ::
    VU.Vector Double -> VU.Vector Double -> [(Double, (Int, Int, Int, Int))]
classCounts :: Vector Double -> Vector Double -> [(Double, (Int, Int, Int, Int))]
classCounts Vector Double
preds Vector Double
truth =
    [(Double
c, Double -> (Int, Int, Int, Int)
forall {a} {b} {c} {d}.
(Num a, Num b, Num c, Num d) =>
Double -> (a, b, c, d)
countsFor Double
c) | Double
c <- [Double]
classes]
  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))
    countsFor :: Double -> (a, b, c, d)
countsFor Double
c =
        ((a, b, c, d) -> (Double, Double) -> (a, b, c, d))
-> (a, b, c, d) -> Vector (Double, Double) -> (a, b, c, d)
forall b a. Unbox b => (a -> b -> a) -> a -> Vector b -> a
VU.foldl'
            ( \(a
tp, b
fp, c
fn, d
sup) (Double
p, Double
y) ->
                ( if Double
p Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c Bool -> Bool -> Bool
&& Double
y Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c then a
tp a -> a -> a
forall a. Num a => a -> a -> a
+ a
1 else a
tp
                , if Double
p Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c Bool -> Bool -> Bool
&& Double
y Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
c then b
fp b -> b -> b
forall a. Num a => a -> a -> a
+ b
1 else b
fp
                , if Double
p Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
/= Double
c Bool -> Bool -> Bool
&& Double
y Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c then c
fn c -> c -> c
forall a. Num a => a -> a -> a
+ c
1 else c
fn
                , if Double
y Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
c then d
sup d -> d -> d
forall a. Num a => a -> a -> a
+ d
1 else d
sup
                )
            )
            (a
0, b
0, c
0, d
0)
            (Vector Double -> Vector Double -> Vector (Double, Double)
forall a b.
(Unbox a, Unbox b) =>
Vector a -> Vector b -> Vector (a, b)
VU.zip Vector Double
preds Vector Double
truth)

safeDiv :: Int -> Int -> Double
safeDiv :: Int -> Int -> Double
safeDiv Int
a Int
b = if Int
b Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 then Double
0 else Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
a Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
b

precOf, recOf :: (Int, Int, Int, Int) -> Double
precOf :: (Int, Int, Int, Int) -> Double
precOf (Int
tp, Int
fp, Int
_, Int
_) = Int -> Int -> Double
safeDiv Int
tp (Int
tp Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
fp)
recOf :: (Int, Int, Int, Int) -> Double
recOf (Int
tp, Int
_, Int
fn, Int
_) = Int -> Int -> Double
safeDiv Int
tp (Int
tp Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
fn)

f1Of :: (Int, Int, Int, Int) -> Double
f1Of :: (Int, Int, Int, Int) -> Double
f1Of (Int, Int, Int, Int)
cs =
    let p :: Double
p = (Int, Int, Int, Int) -> Double
precOf (Int, Int, Int, Int)
cs; r :: Double
r = (Int, Int, Int, Int) -> Double
recOf (Int, Int, Int, Int)
cs in if Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
r Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 then Double
0 else Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
r Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
p Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
r)

averaged ::
    ((Int, Int, Int, Int) -> Double) ->
    Average ->
    VU.Vector Double ->
    VU.Vector Double ->
    Double
averaged :: ((Int, Int, Int, Int) -> Double) -> Average -> Metric
averaged (Int, Int, Int, Int) -> Double
stat Average
avg Vector Double
preds Vector Double
truth =
    case Average
avg of
        Binary Double
pos -> Double
-> ((Int, Int, Int, Int) -> Double)
-> Maybe (Int, Int, Int, Int)
-> Double
forall b a. b -> (a -> b) -> Maybe a -> b
maybe Double
0 (Int, Int, Int, Int) -> Double
stat (Double
-> [(Double, (Int, Int, Int, Int))] -> Maybe (Int, Int, Int, Int)
forall a b. Eq a => a -> [(a, b)] -> Maybe b
lookup Double
pos [(Double, (Int, Int, Int, Int))]
cc)
        Average
Macro -> [Double] -> Double
forall {t :: * -> *} {a}. (Foldable t, Fractional a) => t a -> a
meanOf [(Int, Int, Int, Int) -> Double
stat (Int, Int, Int, Int)
c | (Double
_, (Int, Int, Int, Int)
c) <- [(Double, (Int, Int, Int, Int))]
cc]
        Average
Weighted ->
            let total :: Int
total = [Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int
sup | (Double
_, (Int
_, Int
_, Int
_, Int
sup)) <- [(Double, (Int, Int, Int, Int))]
cc]
             in if Int
total Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0
                    then Double
0
                    else
                        [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
sup Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Int, Int, Int, Int) -> Double
stat (Int, Int, Int, Int)
c | (Double
_, c :: (Int, Int, Int, Int)
c@(Int
_, Int
_, Int
_, Int
sup)) <- [(Double, (Int, Int, Int, Int))]
cc]
                            Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
total
        Average
Micro ->
            let tp :: Int
tp = [Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int
t | (Double
_, (Int
t, Int
_, Int
_, Int
_)) <- [(Double, (Int, Int, Int, Int))]
cc]
                fp :: Int
fp = [Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int
x | (Double
_, (Int
_, Int
x, Int
_, Int
_)) <- [(Double, (Int, Int, Int, Int))]
cc]
                fn :: Int
fn = [Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int
x | (Double
_, (Int
_, Int
_, Int
x, Int
_)) <- [(Double, (Int, Int, Int, Int))]
cc]
             in (Int, Int, Int, Int) -> Double
stat (Int
tp, Int
fp, Int
fn, Int
0)
  where
    cc :: [(Double, (Int, Int, Int, Int))]
cc = Vector Double -> Vector Double -> [(Double, (Int, Int, Int, Int))]
classCounts Vector Double
preds Vector Double
truth
    meanOf :: t a -> a
meanOf t a
xs = if t a -> Bool
forall a. t a -> Bool
forall (t :: * -> *) a. Foldable t => t a -> Bool
null t a
xs then a
0 else t a -> a
forall a. Num a => t a -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum t a
xs a -> a -> a
forall a. Fractional a => a -> a -> a
/ Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral (t a -> Int
forall a. t a -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length t a
xs)

-- | Precision with the given averaging.
precision :: Average -> VU.Vector Double -> VU.Vector Double -> Double
precision :: Average -> Metric
precision = ((Int, Int, Int, Int) -> Double) -> Average -> Metric
averaged (Int, Int, Int, Int) -> Double
precOf

-- | Recall with the given averaging.
recall :: Average -> VU.Vector Double -> VU.Vector Double -> Double
recall :: Average -> Metric
recall = ((Int, Int, Int, Int) -> Double) -> Average -> Metric
averaged (Int, Int, Int, Int) -> Double
recOf

-- | F1 with the given averaging.
f1 :: Average -> VU.Vector Double -> VU.Vector Double -> Double
f1 :: Average -> Metric
f1 = ((Int, Int, Int, Int) -> Double) -> Average -> Metric
averaged (Int, Int, Int, Int) -> Double
f1Of

{- | Binary ROC-AUC (Mann–Whitney). @scores@ are predicted probabilities, @truth@
is @0@/@1@.
-}
rocAuc :: VU.Vector Double -> VU.Vector Double -> Double
rocAuc :: Metric
rocAuc Vector Double
scores Vector Double
truth
    | Double
nPos Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 Bool -> Bool -> Bool
|| Double
nNeg Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
0 = Double
0.5
    | Bool
otherwise = (Double
rankSum Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
nPos Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
nPos Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
1) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
2) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
nPos Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
nNeg)
  where
    ranked :: [Double]
ranked = [Double] -> [Double]
rankAverages (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
scores)
    pairs :: [(Double, Double)]
pairs = [Double] -> [Double] -> [(Double, Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
truth) [Double]
ranked
    rankSum :: Double
rankSum = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Double
r | (Double
y, Double
r) <- [(Double, Double)]
pairs, Double
y Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
1]
    nPos :: Double
nPos = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length ((Double -> Bool) -> [Double] -> [Double]
forall a. (a -> Bool) -> [a] -> [a]
filter (Double -> Double -> Bool
forall a. Eq a => a -> a -> Bool
== Double
1) (Vector Double -> [Double]
forall a. Unbox a => Vector a -> [a]
VU.toList Vector Double
truth)))
    nNeg :: Double
nNeg = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Vector Double -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Double
truth) Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
nPos

-- | Average ranks (ties share the mean rank), returned in input order.
rankAverages :: [Double] -> [Double]
rankAverages :: [Double] -> [Double]
rankAverages [Double]
xs =
    let indexed :: [(Int, Double)]
indexed = [Int] -> [Double] -> [(Int, Double)]
forall a b. [a] -> [b] -> [(a, b)]
zip [Int
0 :: Int ..] [Double]
xs
        sorted :: [(Int, Double)]
sorted = ((Int, Double) -> (Int, Double) -> Ordering)
-> [(Int, Double)] -> [(Int, Double)]
forall a. (a -> a -> Ordering) -> [a] -> [a]
sortBy (((Int, Double) -> Double)
-> (Int, Double) -> (Int, Double) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (Int, Double) -> Double
forall a b. (a, b) -> b
snd) [(Int, Double)]
indexed
        ranked :: [(Int, Double)]
ranked = Int -> [(Int, Double)] -> [(Int, Double)]
forall {b} {b} {a}.
(Eq b, Fractional b) =>
Int -> [(a, b)] -> [(a, b)]
assignRanks Int
1 [(Int, Double)]
sorted
     in ((Int, Double) -> Double) -> [(Int, Double)] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Int, Double) -> Double
forall a b. (a, b) -> b
snd (((Int, Double) -> (Int, Double) -> Ordering)
-> [(Int, Double)] -> [(Int, Double)]
forall a. (a -> a -> Ordering) -> [a] -> [a]
sortBy (((Int, Double) -> Int)
-> (Int, Double) -> (Int, Double) -> Ordering
forall a b. Ord a => (b -> a) -> b -> b -> Ordering
comparing (Int, Double) -> Int
forall a b. (a, b) -> a
fst) [(Int, Double)]
ranked)
  where
    assignRanks :: Int -> [(a, b)] -> [(a, b)]
assignRanks Int
_ [] = []
    assignRanks Int
start [(a, b)]
grp =
        let v :: b
v = (a, b) -> b
forall a b. (a, b) -> b
snd ([(a, b)] -> (a, b)
forall a. HasCallStack => [a] -> a
head [(a, b)]
grp)
            ([(a, b)]
tied, [(a, b)]
rest) = ((a, b) -> Bool) -> [(a, b)] -> ([(a, b)], [(a, b)])
forall a. (a -> Bool) -> [a] -> ([a], [a])
span ((b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== b
v) (b -> Bool) -> ((a, b) -> b) -> (a, b) -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (a, b) -> b
forall a b. (a, b) -> b
snd) [(a, b)]
grp
            k :: Int
k = [(a, b)] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [(a, b)]
tied
            avgRank :: b
avgRank = Int -> b
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([Int] -> Int
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [Int
start .. Int
start Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]) b -> b -> b
forall a. Fractional a => a -> a -> a
/ Int -> b
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
k
         in [(a
i, b
avgRank) | (a
i, b
_) <- [(a, b)]
tied] [(a, b)] -> [(a, b)] -> [(a, b)]
forall a. [a] -> [a] -> [a]
++ Int -> [(a, b)] -> [(a, b)]
assignRanks (Int
start Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
k) [(a, b)]
rest