0 | module Data.Tensor.Softargmax
2 | import Data.Tensor.Tensor
3 | import Data.Tensor.Utils
13 | logSumExp : {i : Axis} -> Exp a => Ord a => Neg a =>
14 | Foldable (Tensor [i]) =>
15 | (allAlg : AllAlgebra [i] a) =>
16 | Tensor [i] a -> Maybe a
19 | pure $
c + log (reduce (t <&> (\x => exp $
x - c)))
25 | logSoftargmax : {i : Axis} -> Exp a => Ord a => Neg a =>
26 | Foldable (Tensor [i]) =>
27 | (allAlg : AllAlgebra [i] a) =>
28 | Tensor [i] a -> Tensor [i] a
29 | logSoftargmax t = case logSumExp t of
30 | Just lse => t <&> (\x => x - lse)
36 | softargmaxImpl : {i : Axis} -> Fractional a => Exp a => Ord a => Neg a =>
37 | IsFoldable i .cont =>
38 | (allAlg : AllAlgebra [i] a) =>
39 | {default 1 temperature : a} ->
40 | Tensor [i] a -> Tensor [i] a
41 | softargmaxImpl {temperature} t
42 | = exp <$> logSoftargmax (t <&> (/ temperature))