-- | Ephemeral learning vocabulary adapted to the circuits-learn interface.
--
-- Based on ideas from the @ephemeral@ package: a learning step is a
-- 'Progress' that transduces a parameterised 'Task' given an 'Experience'.
-- Folding progress over an experience set is 'learn'; choosing the better
-- resulting task is 'improve'.
module Circuit.Learn.Ephemeral
  ( -- * Core vocabulary
    Task (..),
    Experience (..),
    Progress (..),
    Learn (..),

    -- * Learning as folding progress
    learn,
    improve,

    -- * Simple gradient descent progress
    sgd,
  )
where

import Data.Foldable (Foldable (..))

-- | A task is a parameterised measurement: given parameters @p@ and an
-- experience @e@, produce a performance value @r@.
--
-- This is the circuits-learn reading of ephemeral's @Task e p@, with the
-- parameter block @p@ made explicit so it can be updated by 'Progress'.
newtype Task p e r = Task
  { forall p e r. Task p e r -> p -> e -> r
measure :: p -> e -> r
  }

-- | An experience is a container of training examples.
newtype Experience f e = Experience
  { forall {k} (f :: k -> *) (e :: k). Experience f e -> f e
set :: f e
  }

-- | A progress step updates parameters from one experience.
--
-- In the original ephemeral vocabulary this transduces the task itself; here
-- the task is fixed and the parameters that define it are changed.
newtype Progress p e = Progress
  { forall p e. Progress p e -> e -> p -> p
step :: e -> p -> p
  }

-- | A learn folds a 'Progress' over an 'Experience' set.
newtype Learn f e p = Learn
  { forall {k} (f :: k -> *) (e :: k) p.
Learn f e p -> Experience f e -> p -> p
change :: Experience f e -> p -> p
  }

-- | Fold a progress step over all experiences in a set.
learn :: (Foldable f) => Progress p e -> Experience f e -> p -> p
learn :: forall (f :: * -> *) p e.
Foldable f =>
Progress p e -> Experience f e -> p -> p
learn Progress p e
p (Experience f e
es) p
params0 = (p -> e -> p) -> p -> f e -> p
forall b a. (b -> a -> b) -> b -> f a -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' (\p
params e
e -> Progress p e -> e -> p -> p
forall p e. Progress p e -> e -> p -> p
step Progress p e
p e
e p
params) p
params0 f e
es

-- | Apply a learn to a task and choose the better parameters.
--
-- Performance is summarised by a user-supplied function @f (r) -> r@ (for
-- example 'sum' or 'mean') so that two parameter blocks can be compared.
improve ::
  (Foldable f, Functor f, Ord r) =>
  (f r -> r) ->
  Progress p e ->
  Experience f e ->
  Task p e r ->
  p ->
  p
improve :: forall (f :: * -> *) r p e.
(Foldable f, Functor f, Ord r) =>
(f r -> r)
-> Progress p e -> Experience f e -> Task p e r -> p -> p
improve f r -> r
summarize Progress p e
prog Experience f e
es (Task p -> e -> r
measure) p
params0 =
  let params1 :: p
params1 = Progress p e -> Experience f e -> p -> p
forall (f :: * -> *) p e.
Foldable f =>
Progress p e -> Experience f e -> p -> p
learn Progress p e
prog Experience f e
es p
params0
      perf0 :: r
perf0 = f r -> r
summarize (p -> e -> r
measure p
params0 (e -> r) -> f e -> f r
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Experience f e -> f e
forall {k} (f :: k -> *) (e :: k). Experience f e -> f e
set Experience f e
es)
      perf1 :: r
perf1 = f r -> r
summarize (p -> e -> r
measure p
params1 (e -> r) -> f e -> f r
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Experience f e -> f e
forall {k} (f :: k -> *) (e :: k). Experience f e -> f e
set Experience f e
es)
   in if r
perf1 r -> r -> Bool
forall a. Ord a => a -> a -> Bool
> r
perf0 then p
params1 else p
params0

-- | Plain stochastic-gradient-descent progress for vector parameters.
--
-- @sgd rate grad@ takes one example, computes the gradient @grad params e@,
-- and steps the parameters in the negative direction.
sgd ::
  -- | learning rate
  Double ->
  -- | gradient of the task at parameters and example
  ([Double] -> e -> [Double]) ->
  Progress [Double] e
sgd :: forall e.
Double -> ([Double] -> e -> [Double]) -> Progress [Double] e
sgd Double
rate [Double] -> e -> [Double]
grad = (e -> [Double] -> [Double]) -> Progress [Double] e
forall p e. (e -> p -> p) -> Progress p e
Progress ((e -> [Double] -> [Double]) -> Progress [Double] e)
-> (e -> [Double] -> [Double]) -> Progress [Double] e
forall a b. (a -> b) -> a -> b
$ \e
e [Double]
params ->
  let g :: [Double]
g = [Double] -> e -> [Double]
grad [Double]
params e
e
   in (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-) [Double]
params ((Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Double
rate Double -> Double -> Double
forall a. Num a => a -> a -> a
*) [Double]
g)