{-# LANGUAGE UnicodeSyntax #-}
{-# LANGUAGE NoImplicitPrelude #-}
module Net
(
NetParams (..),
netParamsFromArrays,
linear1,
bias1,
relu1,
linear2,
bias2,
model,
forward,
mseLoss,
Boundary,
)
where
import Circuit.Learn.Para (Para (..), runPara)
import Circuit.Mat.Dense (Matrix (..), matVec)
import Control.Category (Category (..))
import Data.These (These (..))
import Data.Vector.Unboxed qualified as VU
import Harpie.Array (Array, arrayAs)
import Harpie.Array qualified as A
import NumHask.Algebra.Additive (Additive, sum, zero)
import NumHask.Algebra.Multiplicative (Multiplicative)
import Prelude hiding (id, sum, (.))
type Boundary a = These a a
data NetParams a = NetParams
{ forall a. NetParams a -> Matrix a
w1 :: Matrix a,
forall a. NetParams a -> Array a
b1 :: Array a,
forall a. NetParams a -> Matrix a
w2 :: Matrix a,
forall a. NetParams a -> Array a
b2 :: Array a
}
deriving (NetParams a -> NetParams a -> Bool
(NetParams a -> NetParams a -> Bool)
-> (NetParams a -> NetParams a -> Bool) -> Eq (NetParams a)
forall a. Eq a => NetParams a -> NetParams a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => NetParams a -> NetParams a -> Bool
== :: NetParams a -> NetParams a -> Bool
$c/= :: forall a. Eq a => NetParams a -> NetParams a -> Bool
/= :: NetParams a -> NetParams a -> Bool
Eq, Int -> NetParams a -> ShowS
[NetParams a] -> ShowS
NetParams a -> String
(Int -> NetParams a -> ShowS)
-> (NetParams a -> String)
-> ([NetParams a] -> ShowS)
-> Show (NetParams a)
forall a. Show a => Int -> NetParams a -> ShowS
forall a. Show a => [NetParams a] -> ShowS
forall a. Show a => NetParams a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> NetParams a -> ShowS
showsPrec :: Int -> NetParams a -> ShowS
$cshow :: forall a. Show a => NetParams a -> String
show :: NetParams a -> String
$cshowList :: forall a. Show a => [NetParams a] -> ShowS
showList :: [NetParams a] -> ShowS
Show)
netParamsFromArrays ::
Array a -> Array a -> Array a -> Array a -> NetParams a
netParamsFromArrays :: forall a. Array a -> Array a -> Array a -> Array a -> NetParams a
netParamsFromArrays Array a
w1a Array a
b1a Array a
w2a Array a
b2a =
Matrix a -> Array a -> Matrix a -> Array a -> NetParams a
forall a. Matrix a -> Array a -> Matrix a -> Array a -> NetParams a
NetParams (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
w1a) Array a
b1a (Array a -> Matrix a
forall a. Array a -> Matrix a
Matrix Array a
w2a) Array a
b2a
rows :: Matrix a -> Int
rows :: forall a. Matrix a -> Int
rows (Matrix Array a
a) = case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a) of [Int
r, Int
_] -> Int
r; [Int]
_ -> Int
0
matVecArr ::
(Additive a, Multiplicative a) => Matrix a -> Array a -> Array a
matVecArr :: forall a.
(Additive a, Multiplicative a) =>
Matrix a -> Array a -> Array a
matVecArr Matrix a
m Array a
v = [Int] -> [a] -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [Matrix a -> Int
forall a. Matrix a -> Int
rows Matrix a
m] (Matrix a -> [a] -> [a]
forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec Matrix a
m (Array a -> [a]
forall t a. FromArray t a => Array a -> t
arrayAs Array a
v))
reluArr :: (Ord a, Additive a) => Array a -> Array a
reluArr :: forall a. (Ord a, Additive a) => Array a -> Array a
reluArr = (a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (\a
x -> a -> a -> a
forall a. Ord a => a -> a -> a
max a
x a
forall a. Additive a => a
zero)
linear1 ::
(Additive a, Multiplicative a) => Para (NetParams a) (Array a) (Array a)
linear1 :: forall a.
(Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
linear1 = ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall p a b. ((p, a) -> b) -> Para p a b
Para (((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a))
-> ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall a b. (a -> b) -> a -> b
$ \(NetParams a
p, Array a
x) -> Matrix a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Matrix a -> Array a -> Array a
matVecArr (NetParams a -> Matrix a
forall a. NetParams a -> Matrix a
w1 NetParams a
p) Array a
x
bias1 :: (Num a) => Para (NetParams a) (Array a) (Array a)
bias1 :: forall a. Num a => Para (NetParams a) (Array a) (Array a)
bias1 = ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall p a b. ((p, a) -> b) -> Para p a b
Para (((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a))
-> ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall a b. (a -> b) -> a -> b
$ \(NetParams a
p, Array a
x) -> (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith a -> a -> a
forall a. Num a => a -> a -> a
(+) Array a
x (NetParams a -> Array a
forall a. NetParams a -> Array a
b1 NetParams a
p)
relu1 :: (Ord a, Additive a) => Para (NetParams a) (Array a) (Array a)
relu1 :: forall a.
(Ord a, Additive a) =>
Para (NetParams a) (Array a) (Array a)
relu1 = ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall p a b. ((p, a) -> b) -> Para p a b
Para (((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a))
-> ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall a b. (a -> b) -> a -> b
$ \(NetParams a
_, Array a
x) -> Array a -> Array a
forall a. (Ord a, Additive a) => Array a -> Array a
reluArr Array a
x
linear2 ::
(Additive a, Multiplicative a) => Para (NetParams a) (Array a) (Array a)
linear2 :: forall a.
(Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
linear2 = ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall p a b. ((p, a) -> b) -> Para p a b
Para (((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a))
-> ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall a b. (a -> b) -> a -> b
$ \(NetParams a
p, Array a
x) -> Matrix a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Matrix a -> Array a -> Array a
matVecArr (NetParams a -> Matrix a
forall a. NetParams a -> Matrix a
w2 NetParams a
p) Array a
x
bias2 :: (Num a) => Para (NetParams a) (Array a) (Array a)
bias2 :: forall a. Num a => Para (NetParams a) (Array a) (Array a)
bias2 = ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall p a b. ((p, a) -> b) -> Para p a b
Para (((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a))
-> ((NetParams a, Array a) -> Array a)
-> Para (NetParams a) (Array a) (Array a)
forall a b. (a -> b) -> a -> b
$ \(NetParams a
p, Array a
x) -> (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith a -> a -> a
forall a. Num a => a -> a -> a
(+) Array a
x (NetParams a -> Array a
forall a. NetParams a -> Array a
b2 NetParams a
p)
model ::
(Ord a, Num a, Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
model :: forall a.
(Ord a, Num a, Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
model = Para (NetParams a) (Array a) (Array a)
forall a. Num a => Para (NetParams a) (Array a) (Array a)
bias2 Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
forall b c a.
Para (NetParams a) b c
-> Para (NetParams a) a b -> Para (NetParams a) a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Para (NetParams a) (Array a) (Array a)
forall a.
(Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
linear2 Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
forall b c a.
Para (NetParams a) b c
-> Para (NetParams a) a b -> Para (NetParams a) a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Para (NetParams a) (Array a) (Array a)
forall a.
(Ord a, Additive a) =>
Para (NetParams a) (Array a) (Array a)
relu1 Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
forall b c a.
Para (NetParams a) b c
-> Para (NetParams a) a b -> Para (NetParams a) a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Para (NetParams a) (Array a) (Array a)
forall a. Num a => Para (NetParams a) (Array a) (Array a)
bias1 Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
-> Para (NetParams a) (Array a) (Array a)
forall b c a.
Para (NetParams a) b c
-> Para (NetParams a) a b -> Para (NetParams a) a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Para (NetParams a) (Array a) (Array a)
forall a.
(Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
linear1
forward ::
(Ord a, Num a, Additive a, Multiplicative a) =>
NetParams a ->
Array a ->
Array a
forward :: forall a.
(Ord a, Num a, Additive a, Multiplicative a) =>
NetParams a -> Array a -> Array a
forward = Para (NetParams a) (Array a) (Array a)
-> NetParams a -> Array a -> Array a
forall p a b. Para p a b -> p -> a -> b
runPara Para (NetParams a) (Array a) (Array a)
forall a.
(Ord a, Num a, Additive a, Multiplicative a) =>
Para (NetParams a) (Array a) (Array a)
model
mseLoss ::
(Fractional a, Additive a) =>
Array a ->
Array a ->
(a, Array a)
mseLoss :: forall a.
(Fractional a, Additive a) =>
Array a -> Array a -> (a, Array a)
mseLoss Array a
y Array a
target =
let diff :: Array a
diff = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith (-) Array a
y Array a
target
n :: a
n = Int -> a
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Array a -> Int
forall a. Array a -> Int
A.size Array a
y)
loss :: a
loss = Array a -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum ((a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (\a
x -> a
x a -> a -> a
forall a. Num a => a -> a -> a
* a
x) Array a
diff) a -> a -> a
forall a. Fractional a => a -> a -> a
/ a
n
grad :: Array a
grad = (a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (a -> a -> a
forall a. Num a => a -> a -> a
* (a
2 a -> a -> a
forall a. Fractional a => a -> a -> a
/ a
n)) Array a
diff
in (a
loss, Array a
grad)