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'