{-# LANGUAGE UnicodeSyntax #-}
{-# LANGUAGE NoImplicitPrelude #-}

-- | Forward-only neural-network layer definitions built on the circuits
-- ecosystem.
--
-- * 'NetParams' stores weights as 'Circuit.Mat.Dense.Matrix' and biases as
--   'Harpie.Array' vectors.
-- * Linear layers use 'Circuit.Mat.Dense.matVec'.
-- * The full model is a 'Circuit.Learn.Para' composition, so inference
--   reuses the same parameter threading as the rest of circuits-learn.
module Net
  ( -- * Parameter bundle
    NetParams (..),
    netParamsFromArrays,

    -- * Layer primitives
    linear1,
    bias1,
    relu1,
    linear2,
    bias2,

    -- * Model
    model,
    forward,

    -- * Loss
    mseLoss,

    -- * Boundary type
    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, (.))

-- | A 'These'-based boundary distinguishes inference-only, gradient-only,
-- and combined inference-with-gradient traffic at a layer boundary.
--
-- * 'This' prediction — inference only, no gradient.
-- * 'That' gradient — training signal only, no prediction.
-- * 'These' prediction gradient — training with prediction (common
--   supervised case).
type Boundary a = These a a

-- | Network parameters.  Weights are stored as dense matrices so that
-- 'Circuit.Mat.Dense' can be used for the linear maps; biases and
-- activations remain as 'Array' vectors.
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)

-- | Build 'NetParams' from plain 'Array' weights and biases.
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

-- | Row count of a dense matrix.
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

-- | Matrix–vector product, returning an 'Array'.
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))

-- | Rectified linear unit.
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)

-- | First linear layer.
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

-- | First bias layer.
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)

-- | ReLU activation.
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

-- | Second linear layer.
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

-- | Second bias layer.
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)

-- | Full 2-layer MLP: linear1 → bias1 → relu → linear2 → bias2.
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

-- | Run the model with explicit parameters.
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

-- | Mean-squared-error loss and its gradient w.r.t. the prediction.
--
-- @L = (1/n) Σ (y - target)²@, @dL/dy = (2/n)(y - target)@.
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)