packages feed

futhark-0.26.4: src/Futhark/AD/Rev.hs

{-# LANGUAGE TypeFamilies #-}

-- Naming scheme:
--
-- An adjoint-related object for "x" is named "x_adj".  This means
-- both actual adjoints and statements.
--
-- Do not assume "x'" means anything related to derivatives.
module Futhark.AD.Rev (revVJP) where

import Control.Monad
import Data.List.NonEmpty (NonEmpty (..))
import Data.Map qualified as M
import Data.Tuple
import Futhark.AD.Derivatives
import Futhark.AD.Rev.Acc
import Futhark.AD.Rev.Loop
import Futhark.AD.Rev.Monad
import Futhark.AD.Rev.SOAC
import Futhark.AD.Shared
import Futhark.Analysis.PrimExp.Convert
import Futhark.Builder
import Futhark.IR.SOACS
import Futhark.Tools
import Futhark.Transform.Rename
import Futhark.Transform.Substitute
import Futhark.Util (takeLast)

patName :: Pat Type -> ADM VName
patName (Pat [pe]) = pure $ patElemName pe
patName pat = error $ "Expected single-element pattern: " ++ prettyString pat

copyIfArray :: VName -> ADM VName
copyIfArray v = do
  v_t <- lookupType v
  case v_t of
    Array {} ->
      letExp (baseName v <> "_copy") . BasicOp $ Replicate mempty (Var v)
    _ -> pure v

-- The vast majority of BasicOps require no special treatment in the
-- forward pass and produce one value (and hence one adjoint).  We
-- deal with that case here.
commonBasicOp :: Pat Type -> StmAux () -> BasicOp -> ADM () -> ADM (VName, VName)
commonBasicOp pat aux op m = do
  addStm $ Let pat aux $ BasicOp op
  m
  pat_v <- patName pat
  pat_adj <- lookupAdjVal pat_v
  pure (pat_v, pat_adj)

diffBasicOp :: Pat Type -> StmAux () -> BasicOp -> ADM () -> ADM ()
diffBasicOp pat aux e m =
  case e of
    CmpOp {} ->
      void $ commonBasicOp pat aux e m
    --
    ConvOp op x -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        adj_shape <- askShape

        contrib <- letExp "convop_contrib" <=< mapNest adj_shape (MkSolo (Var pat_adj)) $
          \(MkSolo pat_adj') ->
            pure $ BasicOp $ ConvOp (flipConvOp op) pat_adj'

        updateSubExpAdj x contrib
    --
    UnOp op x -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m

      returnSweepCode $ do
        let t = unOpType op

        adj_shape <- askShape

        contrib <- letExp "unop_contrib" <=< mapNest adj_shape (MkSolo (Var pat_adj)) $
          \(MkSolo pat_adj') ->
            toExp $ primExpFromSubExp t pat_adj' ~*~ pdUnOp op (primExpFromSubExp t x)

        updateSubExpAdj x contrib
    --
    BinOp op x y -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m

      returnSweepCode $ do
        let t = binOpType op
            (wrt_x, wrt_y) =
              pdBinOp op (primExpFromSubExp t x) (primExpFromSubExp t y)

        adj_shape <- askShape

        adj_x <- letExp "binop_x_adj"
          <=< mapNest adj_shape (MkSolo (Var pat_adj))
          $ \(MkSolo pat_adj') ->
            let pat_adj'' = primExpFromSubExp t pat_adj'
             in toExp $ pat_adj'' ~*~ wrt_x

        adj_y <- letExp "binop_y_adj"
          <=< mapNest adj_shape (MkSolo (Var pat_adj))
          $ \(MkSolo pat_adj') ->
            let pat_adj'' = primExpFromSubExp t pat_adj'
             in toExp $ pat_adj'' ~*~ wrt_y

        updateSubExpAdj x adj_x
        updateSubExpAdj y adj_y
    --
    SubExp se -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ updateSubExpAdj se pat_adj
    --
    Assert {} ->
      void $ commonBasicOp pat aux e m
    --
    ArrayVal {} ->
      void $ commonBasicOp pat aux e m
    --
    ArrayLit elems _ -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      t <- lookupType pat_adj
      returnSweepCode $ do
        forM_ (zip [(0 :: Int64) ..] elems) $ \(i, se) -> do
          let slice = fullSlice t [DimFix (constant i)]
          updateSubExpAdj se <=< letExp "elem_adj" $ BasicOp $ Index pat_adj slice
    --
    Index arr slice -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ void $ updateAdjSlice slice arr pat_adj
    FlatIndex {} -> error "FlatIndex not handled by AD yet."
    FlatUpdate {} -> error "FlatUpdate not handled by AD yet."
    --
    Opaque _ se -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ updateSubExpAdj se pat_adj
    --
    Reshape arr newshape -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        arr_shape <- arrayShape <$> lookupType arr
        void $
          updateAdj arr <=< letExp "adj_reshape" . BasicOp $
            Reshape pat_adj (reshapeAll (newShape newshape) arr_shape)
    --
    Rearrange arr perm -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      r <- shapeRank <$> askShape
      returnSweepCode $
        void . updateAdj arr <=< letExp "adj_rearrange" . BasicOp $
          Rearrange pat_adj ([0 .. r - 1] <> map (+ r) (rearrangeInverse perm))
    --
    Replicate (Shape []) (Var se) -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ void $ updateAdj se pat_adj
    --
    Replicate (Shape ns) x -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        x_t <- subExpType x
        lam <- addLambda x_t
        ne <- letSubExp "zero" $ zeroExp x_t
        n <- letSubExp "rep_size" =<< foldBinOp (Mul Int64 OverflowUndef) (intConst Int64 1) ns
        pat_adj_flat <-
          letExp (baseName pat_adj <> "_flat") . BasicOp $
            Reshape pat_adj (reshapeAll (Shape ns) (Shape $ n : arrayDims x_t))
        reduce <- reduceSOAC [Reduce Commutative lam [ne]]
        updateSubExpAdj x
          =<< letExp "rep_contrib" (Op $ Screma n [pat_adj_flat] reduce)
    --
    Concat d (arr :| arrs) _ -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        let sliceAdj _ [] = pure []
            sliceAdj start (v : vs) = do
              v_t <- lookupType v
              let w = arraySize 0 v_t
                  slice = DimSlice start w (intConst Int64 1)
              pat_adj_slice <-
                letExp (baseName pat_adj <> "_slice") $
                  BasicOp $
                    Index pat_adj (sliceAt v_t d [slice])
              start' <- letSubExp "start" $ BasicOp $ BinOp (Add Int64 OverflowUndef) start w
              slices <- sliceAdj start' vs
              pure $ pat_adj_slice : slices

        slices <- sliceAdj (intConst Int64 0) $ arr : arrs

        zipWithM_ updateAdj (arr : arrs) slices
    --
    Manifest se _ -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ void $ updateAdj se pat_adj
    --
    Scratch {} ->
      void $ commonBasicOp pat aux e m
    --
    Iota n _ _ t -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        ne <- letSubExp "zero" $ zeroExp $ Prim $ IntType t
        lam <- addLambda $ Prim $ IntType t
        reduce <- reduceSOAC [Reduce Commutative lam [ne]]
        updateSubExpAdj n
          =<< letExp "iota_contrib" (Op $ Screma n [pat_adj] reduce)
    --
    Update safety arr slice v -> do
      (_pat_v, pat_adj) <- commonBasicOp pat aux e m
      returnSweepCode $ do
        adj_shape <- askShape
        let adj_slice = Slice $ map sliceDim (shapeDims adj_shape) ++ unSlice slice
        v_adj <- letExp "update_val_adj" $ BasicOp $ Index pat_adj adj_slice
        v_adj_copy <- copyIfArray v_adj
        updateSubExpAdj v v_adj_copy
        v_adj_t <- lookupType v_adj
        zeroes <- letSubExp "update_zero" $ zeroExp v_adj_t
        void $
          updateAdj arr
            =<< letExp "update_src_adj" (BasicOp $ Update safety pat_adj adj_slice zeroes)
    UpdateAcc safety acc is vs ->
      diffUpdateAcc pat aux safety acc is vs m
    --
    UserParam {} ->
      void $ commonBasicOp pat aux e m

vjpOps :: VjpOps
vjpOps =
  VjpOps
    { vjpLambda = diffLambda,
      vjpStm = diffStm,
      vjpBody = diffBody
    }

diffStm :: Stm SOACS -> ADM () -> ADM ()
diffStm (Let pat aux (BasicOp e)) m =
  diffBasicOp pat aux e m
diffStm stm@(Let pat _ (Apply f args _ _)) m
  | Just (ret, argts) <- M.lookup f builtInFunctions = do
      addStm stm
      m

      pat_adj <- lookupAdjVal =<< patName pat
      let arg_pes = zipWith primExpFromSubExp argts (map fst args)
          convert ft tt
            | ft == tt = id
          convert (IntType ft) (IntType tt) = ConvOpExp (SExt ft tt)
          convert (FloatType ft) (FloatType tt) = ConvOpExp (FPConv ft tt)
          convert Bool (FloatType tt) = ConvOpExp (BToF tt)
          convert (FloatType ft) Bool = ConvOpExp (FToB ft)
          convert ft tt = error $ "diffStm.convert: " ++ prettyString (f, ft, tt)

      adj_shape <- askShape

      contribs <-
        case pdBuiltin f arg_pes of
          Nothing ->
            error $ "No partial derivative defined for builtin function: " ++ prettyString f
          Just derivs ->
            forM (zip derivs argts) $ \(deriv, argt) ->
              letExp "apply_contrib" <=< mapNest adj_shape (MkSolo (Var pat_adj)) $
                \(MkSolo pat_adj') ->
                  toExp $ convert ret argt $ primExpFromSubExp ret pat_adj' ~*~ deriv

      zipWithM_ updateSubExpAdj (map fst args) contribs
diffStm stm@(Let pat _ (Match ses cases defbody _)) m = do
  addStm stm
  m
  returnSweepCode $ do
    let cases_free = map freeIn cases
        defbody_free = freeIn defbody
        branches_free = namesToList $ mconcat $ defbody_free : cases_free

    adjs <- mapM lookupAdj $ patNames pat

    branches_free_adj <-
      ( pure . takeLast (length branches_free)
          <=< letTupExp "branch_adj"
          <=< renameExp
      )
        =<< eMatch
          ses
          (map (fmap $ diffBody adjs branches_free) cases)
          (diffBody adjs branches_free defbody)
    -- See Note [Array Adjoints of Match]
    forM_ (zip branches_free branches_free_adj) $ \(v, v_adj) ->
      insAdj v =<< copyIfArray v_adj
diffStm (Let pat aux (Op soac)) m =
  -- We add the attributes from 'aux' to every SOAC (but only SOAC) produced. We
  -- could do this on *every* stm, but it would be very verbose.
  censorStms (fmap addAttrs) $ vjpSOAC vjpOps pat aux soac m
  where
    addAttrs stm
      | Op _ <- stmExp stm =
          attribute (stmAuxAttrs aux) stm
      | otherwise = stm
diffStm (Let pat aux loop@Loop {}) m =
  diffLoop diffStms pat aux loop m
-- See Note [Adjoints of accumulators]
diffStm (Let pat aux (WithAcc inputs lam)) m =
  diffWithAcc vjpOps pat aux inputs lam m
diffStm stm _ = error $ "diffStm unhandled:\n" ++ prettyString stm

diffStms :: Stms SOACS -> ADM ()
diffStms all_stms
  | Just (stm, stms) <- stmsHead all_stms = do
      (subst, copy_stms) <- copyConsumedArrsInStm stm
      let (stm', stms') = substituteNames subst (stm, stms)
      diffStms copy_stms >> diffStm stm' (diffStms stms')
      forM_ (M.toList subst) $ \(from, to) ->
        setAdj from =<< lookupAdj to
  | otherwise =
      pure ()

-- | Preprocess statements before differentiating.
-- For now, it's just stripmining.
preprocess :: Stms SOACS -> ADM (Stms SOACS)
preprocess = stripmineStms

diffBody :: [Adj] -> [VName] -> Body SOACS -> ADM (Body SOACS)
diffBody res_adjs get_adjs_for (Body () stms res) = subAD $
  subSubsts $ do
    let onResult (SubExpRes _ (Constant _)) _ = pure ()
        onResult (SubExpRes _ (Var v)) v_adj = void $ updateAdj v =<< adjVal v_adj
    (adjs, stms') <- collectStms $ do
      zipWithM_ onResult (takeLast (length res_adjs) res) res_adjs
      diffStms =<< preprocess stms
      mapM lookupAdjVal get_adjs_for
    pure $ Body () stms' $ res <> varsRes adjs

diffLambda :: [Adj] -> [VName] -> Lambda SOACS -> ADM (Lambda SOACS)
diffLambda res_adjs get_adjs_for (Lambda params _ body) =
  mkLambda params $ do
    res <- bodyBind =<< diffBody res_adjs get_adjs_for body
    pure $ takeLast (length get_adjs_for) res

revVJP ::
  (MonadFreshNames m) =>
  Scope SOACS ->
  Shape ->
  Attrs ->
  Lambda SOACS ->
  m (Lambda SOACS)
revVJP scope shape attrs (Lambda params ts body) = do
  runADM shape attrs . localScope (scope <> scopeOfLParams params) $ do
    adj_shape <- askShape
    params_adj <- forM (zip (map resSubExp (bodyResult body)) ts) $ \(se, t) ->
      Param mempty
        <$> maybe (newVName "const_res_adj") adjVName (subExpVar se)
        <*> pure (t `arrayOfShape` adj_shape)

    body' <-
      localScope (scopeOfLParams params_adj) $
        diffBody
          (map adjFromParam params_adj)
          (map paramName params)
          body

    pure $ Lambda (params ++ params_adj) (ts <> map paramType params) body'