0 | module Data.Num
  1 |
  2 | import Data.Vect
  3 |
  4 | namespace Num
  5 |   %defaulthint
  6 |   public export
  7 |   applicativeNum : Num a => Applicative f => Num (f a)
  8 |   applicativeNum = MkNum
  9 |     (\xs, ys => [| xs + ys |])
 10 |     (\xs, ys => [| xs * ys |])
 11 |     (\n => pure (fromInteger n))
 12 |
 13 |   public export
 14 |   Num Unit where
 15 |     () + () = ()
 16 |     () * () = ()
 17 |     fromInteger x = ()
 18 |   
 19 |   public export
 20 |   Num a => Num b => Num (a, b) where
 21 |     (lFst, lSnd) + (rFst, rSnd) = (lFst + rFst, lSnd + rSnd)
 22 |     (lFst, lSnd) * (rFst, rSnd) = (lFst * rFst, lSnd * rSnd)
 23 |     fromInteger x = (fromInteger x, fromInteger x)
 24 |
 25 |   public export
 26 |   Num a => Num b => Num (DPair a (const b)) where
 27 |     (lFst ** lSnd* (rFst ** rSnd= (lFst * rFst ** lSnd * rSnd)
 28 |     (lFst ** lSnd+ (rFst ** rSnd= (lFst + rFst ** lSnd + rSnd)
 29 |     fromInteger x = (fromInteger x ** fromInteger x)
 30 |
 31 |   %hint
 32 |   public export
 33 |   depFunNum : {k : Fin n -> Type} ->
 34 |     {ss : (i : Fin n) -> Num (k i)} ->
 35 |      Num ((i : Fin n) -> k i)
 36 |   depFunNum = MkNum
 37 |     (\s, t => \i => s i + t i)
 38 |     (\f, g => \i => f i * g i)
 39 |     (\n => \i => fromInteger n)
 40 |
 41 |
 42 | namespace Neg
 43 |   %defaulthint
 44 |   public export
 45 |   applicativeNeg : Neg a => Applicative f => Neg (f a)
 46 |   applicativeNeg = MkNeg @{applicativeNum}
 47 |     (\fa => [| negate fa |])
 48 |     (\fx, fy => [| fx - fy |])
 49 |
 50 |   public export
 51 |   Neg Unit where
 52 |     negate () = ()
 53 |     () - () = ()
 54 |
 55 |   public export
 56 |   Neg a => Neg b => Neg (a, b) where
 57 |     negate (lFst, lSnd) = (negate lFst, negate lSnd)
 58 |     (lFst, lSnd) - (rFst, rSnd) = (lFst - rFst, lSnd - rSnd)
 59 |
 60 |   public export
 61 |   Neg a => Neg b => Neg (DPair a (const b)) where
 62 |     negate (fst ** snd= (negate fst ** negate snd)
 63 |     (fst ** snd- (rFst ** rSnd= (fst - rFst ** snd - rSnd)
 64 |
 65 |   %hint
 66 |   public export
 67 |   depFunNeg : {k : Fin n -> Type} ->
 68 |     {ss : (i : Fin n) -> Neg (k i)} ->
 69 |      Neg ((i : Fin n) -> k i)
 70 |   depFunNeg = MkNeg
 71 |     (\s => \i => negate (s i))
 72 |     (\f, g => \i => f i - g i)
 73 |
 74 | namespace Abs
 75 |   %defaulthint
 76 |   public export
 77 |   applicativeAbs : Abs a => Applicative f => Abs (f a)
 78 |   applicativeAbs = MkAbs @{applicativeNum}
 79 |     (\fa => [| abs fa |])
 80 |
 81 |   public export
 82 |   Abs Unit where
 83 |     abs () = ()
 84 |
 85 |   public export
 86 |   Abs a => Abs b => Abs (a, b) where
 87 |     abs (lFst, lSnd) = (abs lFst, abs lSnd)
 88 |
 89 |   public export
 90 |   Abs a => Abs b => Abs (DPair a (const b)) where
 91 |     abs (fst ** snd= (abs fst ** abs snd)
 92 |
 93 | namespace FromDouble
 94 |   %defaulthint
 95 |   public export
 96 |   applicativeFromDouble : FromDouble a => Applicative f => FromDouble (f a)
 97 |   applicativeFromDouble = MkFromDouble (pure . fromDouble)
 98 |
 99 |   public export
100 |   FromDouble Unit where
101 |     fromDouble x = ()
102 |
103 |   public export
104 |   FromDouble a => FromDouble b => FromDouble (a, b) where
105 |     fromDouble x = (fromDouble x, fromDouble x)
106 |   
107 |   public export
108 |   FromDouble a => FromDouble b => FromDouble (DPair a (const b)) where
109 |     fromDouble x = (fromDouble x ** fromDouble x)
110 |
111 |   %hint
112 |   public export
113 |   depFunFromDouble : {k : Fin n -> Type} ->
114 |     {ss : (i : Fin n) -> FromDouble (k i)} ->
115 |      FromDouble ((i : Fin n) -> k i)
116 |   depFunFromDouble = MkFromDouble
117 |     (\s => \i => fromDouble s)
118 |
119 | namespace Fractional
120 |   %defaulthint
121 |   public export
122 |   applicativeFractional : Fractional a => Applicative f => Fractional (f a)
123 |   applicativeFractional = MkFractional @{applicativeNum}
124 |     (\fx, fy => [| fx / fy |])
125 |     (\x => [| recip x |])
126 |
127 |   public export
128 |   Fractional Unit where
129 |     () / () = ()
130 |     
131 |   public export 
132 |   Fractional a => Fractional b => Fractional (a, b) where
133 |     (lFst, lSnd) / (rFst, rSnd) = (lFst / rFst, lSnd / rSnd)
134 |
135 |   public export
136 |   Fractional a => Fractional b => Fractional (DPair a (const b)) where
137 |     (fst ** snd/ (rFst ** rSnd= (fst / rFst ** snd / rSnd)
138 |
139 |
140 | namespace Exp
141 |   ||| Interface for the Exponential
142 |   ||| We also include minus infinity because of the necessity to compute
143 |   ||| causal masks within the attention mechanism.
144 |   ||| For rules that `exp` should satisfy, see https://arxiv.org/abs/1911.04790
145 |   ||| We also have
146 |   ||| `exp . log = id`, `log . exp = id`, `exp minusInfinity = 0`...
147 |   public export
148 |   interface Num a => Exp a where
149 |     constructor MkExp
150 |     exp : a -> a
151 |     log : a -> a
152 |     minusInfinity : a
153 |   
154 |   public export
155 |   Exp Double where
156 |     exp = Prelude.exp
157 |     log = Prelude.log
158 |     minusInfinity = cast "-inf.0"
159 |
160 |   %defaulthint
161 |   public export
162 |   applicativeExp : Exp a => Applicative f => Exp (f a)
163 |   applicativeExp = MkExp @{applicativeNum}
164 |     (\fa => [| exp fa |])
165 |     (\fa => [| log fa |])
166 |     (pure minusInfinity)
167 |
168 | namespace Sqrt
169 |   public export
170 |   interface Num a => Sqrt a where
171 |     constructor MkSqrt
172 |     sqrt : a -> a
173 |   
174 |   public export
175 |   Sqrt Double where
176 |     sqrt = Prelude.sqrt
177 |
178 |   %defaulthint
179 |   public export
180 |   applicativeSqrt : Sqrt a => Applicative f => Sqrt (f a)
181 |   applicativeSqrt = MkSqrt @{applicativeNum} (\fa => [| sqrt fa |])
182 |
183 |   public export
184 |   Sqrt Unit where
185 |     sqrt () = ()
186 |
187 |   public export
188 |   Sqrt a => Sqrt b => Sqrt (a, b) where
189 |     sqrt (lFst, lSnd) = (sqrt lFst, sqrt lSnd)
190 |
191 |   public export
192 |   Sqrt a => Sqrt b => Sqrt (DPair a (const b)) where
193 |     sqrt (fst ** snd= (sqrt fst ** sqrt snd)