Skip to content

Commit 599064b

Browse files
author
kksnowy
committed
Merge dev into counter-example-generation
2 parents abf7490 + bc393f5 commit 599064b

33 files changed

Lines changed: 298 additions & 101 deletions

File tree

ChangeLog.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,8 @@ the ITP backend code can be invoked from any location.
109109

110110
### Loss backend
111111

112+
* `--declaration` now accepts non-property declarations and restricts output to exactly the names listed.
113+
112114
* Added the ability to declare custom Differentiable Logics internally in Vehicle (see documentation for details).
113115

114116
* Fixed a bug where the compiler was erroring on some uses of `forall` for indices.

vehicle/src/Vehicle/Backend/Loss.hs

Lines changed: 19 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@ where
66
import Control.Monad.Reader (ReaderT)
77
import Data.Maybe (maybeToList)
88
import Data.Proxy (Proxy (..))
9+
import Data.Set (Set)
10+
import Data.Set qualified as Set
911
import Vehicle.Backend.Loss.Core
1012
import Vehicle.Backend.Loss.Domain (compileQuantifier)
1113
import Vehicle.Backend.Loss.LogicCompilation (findAndCompileLogic)
@@ -18,6 +20,7 @@ import Vehicle.Compile.Prelude
1820
import Vehicle.Data.Builtin.Loss
1921
import Vehicle.Data.Builtin.Standard
2022
import Vehicle.Data.Builtin.Standard.Normalise ()
23+
import Vehicle.Data.Code.BooleanExpr (unDisjunctAll)
2124
import Vehicle.Data.Code.ForcedValue
2225
import Vehicle.Data.DifferentiableLogic
2326
import Vehicle.Data.Variable.Bound.Context.Tensor (TensorBoundContextT)
@@ -26,49 +29,52 @@ import Vehicle.Data.Variable.Free.Context (MonadFreeContext (..), addDeclEntryTo
2629
convertToLossTensors ::
2730
(MonadCompile m) =>
2831
DifferentiableLogicID ->
32+
Set Name ->
2933
Prog Builtin ->
3034
m (Prog LossBuiltin)
31-
convertToLossTensors logicID prog@(Main ds) = do
35+
convertToLossTensors logicID requestedDecls prog@(Main ds) = do
3236
-- First find and compile the logic
3337
logic <- logCompilerPass LossLogic $ findAndCompileLogic logicID prog
3438

3539
-- Then compile the program using that logic
3640
runFreshFreeContextT (Proxy @Builtin) $ do
3741
runFreshFreeContextT (Proxy @LossBuiltin) $
3842
logCompilerPass Loss $ do
39-
Main <$> convertDecls logicID logic ds
43+
Main <$> convertDecls logicID logic requestedDecls ds
4044

4145
convertDecls ::
4246
(MonadCompile m, MonadFreeContext Builtin m, MonadFreeContext LossBuiltin m) =>
4347
DifferentiableLogicID ->
4448
DifferentiableLogicImplementation ->
49+
Set Name ->
4550
[Decl Builtin] ->
4651
m [Decl LossBuiltin]
47-
convertDecls logicID logic = \case
52+
convertDecls logicID logic requestedDecls = \case
4853
[] -> return []
4954
decl : decls -> do
50-
maybeLossDecl <- convertDecl logicID logic decl
55+
maybeLossDecl <- convertDecl logicID logic requestedDecls decl
5156
decls' <-
5257
maybe id addDeclToContext maybeLossDecl $
5358
addDeclEntryToContext decl $
54-
convertDecls logicID logic decls
59+
convertDecls logicID logic requestedDecls decls
5560
return $ maybeToList maybeLossDecl ++ decls'
5661

5762
convertDecl ::
5863
forall m.
5964
(MonadCompile m, MonadFreeContext Builtin m, MonadFreeContext LossBuiltin m) =>
6065
DifferentiableLogicID ->
6166
DifferentiableLogicImplementation ->
67+
Set Name ->
6268
Decl Builtin ->
6369
m (Maybe (Decl LossBuiltin))
64-
convertDecl logicID logic decl = case decl of
70+
convertDecl logicID logic requestedDecls decl = case decl of
6571
DefAbstract p ident sort typ
6672
| isAnnotatedAsExternalResource sort -> do
6773
let normType = Unforced emptyBoundEnv typ
6874
runConversion $ convertResourceDecl p ident sort normType
6975
| otherwise -> return Nothing
7076
DefFunction p ident ann typ expr
71-
| isAnnotatedAsProperty ann -> do
77+
| isAnnotatedAsProperty ann || nameOf decl `Set.member` requestedDecls -> do
7278
let normType = Unforced emptyBoundEnv typ
7379
let normExpr = Unforced emptyBoundEnv expr
7480
runConversion $ convertPropertyDecl p ident ann normType normExpr
@@ -104,11 +110,14 @@ convertPropertyDecl ::
104110
convertPropertyDecl p ident ann typ body = do
105111
lossType <- convertDeclType typ
106112
lossBody <- convertMultiProperty body
107-
let lossTensorDecl = DefFunction p ident ann lossType lossBody
108-
return lossTensorDecl
113+
return $ DefFunction p ident ann lossType lossBody
109114

110115
convertDeclType :: (MonadLogic m) => UnforcedType Builtin -> m (Type LossBuiltin)
111116
convertDeclType typ = unnormalise 0 <$> convertThunk Nothing typ
112117

113118
convertMultiProperty :: (MonadLogic m) => Thunk Builtin -> m (Expr LossBuiltin)
114-
convertMultiProperty body = unnormalise 0 <$> convertThunk (Just compileQuantifier) body
119+
convertMultiProperty body = do
120+
let compQuantifier args = do
121+
disjuncts <- compileQuantifier args
122+
foldrM1 orLossValue $ unDisjunctAll disjuncts
123+
unnormalise 0 <$> convertThunk (Just compQuantifier) body

vehicle/src/Vehicle/Backend/Loss/Domain.hs

Lines changed: 7 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
module Vehicle.Backend.Loss.Domain
22
( compileQuantifier,
3+
orLossValue,
34
)
45
where
56

@@ -12,7 +13,7 @@ import Data.List.NonEmpty (NonEmpty (..))
1213
import Data.Map (Map)
1314
import Data.Map qualified as Map
1415
import Vehicle.Backend.Loss.Core
15-
import Vehicle.Backend.Loss.Domain.PurifyAssertion (BlockingReason, tryPurifyAssertion, unblockingActions)
16+
import Vehicle.Backend.Loss.Domain.PurifyAssertion (tryPurifyAssertion, unblockingActions)
1617
import Vehicle.Backend.Loss.LossCompilation
1718
import Vehicle.Backend.Solver.UserVariableElimination.ConstraintSearch (findAllBounds)
1819
import Vehicle.Compile.Constants.ForcedValue
@@ -50,17 +51,17 @@ import Vehicle.Prelude.Warning (CompileWarning (..))
5051
compileQuantifier ::
5152
(MonadLogic m) =>
5253
(Quantifier, QuantifyRatTensorArgs (Thunk Builtin) (Closure Builtin)) ->
53-
m (Thunk LossBuiltin)
54+
m (DisjunctAll (Thunk LossBuiltin))
5455
compileQuantifier (q, args) = do
5556
maybePartitions <- compileQuantifierInternal (q, args)
5657
case maybePartitions of
57-
Trivial b ->
58+
Trivial b -> do
5859
-- TODO add a warning
59-
Forced <$> convertBoolTensorLiteral (ZeroDimTensor b)
60+
value <- Forced <$> convertBoolTensorLiteral (ZeroDimTensor b)
61+
return $ DisjunctAll [value]
6062
NonTrivial partitions -> do
6163
let disjunctedPartitions = partitionsToDisjuncts partitions
62-
DisjunctAll (v :| vs) <- traverse checkFinalPartitionUnconstrained disjunctedPartitions
63-
foldrM orLossValue v vs
64+
traverse checkFinalPartitionUnconstrained disjunctedPartitions
6465

6566
checkFinalPartitionUnconstrained ::
6667
(MonadLogic m) =>
@@ -292,21 +293,6 @@ type MonadDomain m =
292293
MonadTensorBoundContext m
293294
)
294295

295-
orLossValue :: (MonadDomain m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
296-
orLossValue e1 e2 =
297-
convertBooleanOp PointwiseDisjunction $
298-
mkExpr accessSpine (TensorOp2Args (Forced IDimNil) e1 e2)
299-
300-
andLossValue :: (MonadDomain m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
301-
andLossValue e1 e2 =
302-
convertBooleanOp PointwiseConjunction $
303-
mkExpr accessSpine (TensorOp2Args (Forced IDimNil) e1 e2)
304-
305-
notLossValue :: (MonadDomain m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
306-
notLossValue dims e =
307-
convertBooleanOp PointwiseNegation $
308-
mkExpr accessSpine (TensorOp1Args dims e)
309-
310296
notConstraint :: (MonadDomain m) => UserVariableConstraint LossBuiltin -> m (BooleanExpr (UserVariableConstraint LossBuiltin))
311297
notConstraint (NormalisedRelation rel expr) = do
312298
negExpr <- scaleExpr (-1) expr

vehicle/src/Vehicle/Backend/Loss/Domain/PurifyAssertion.hs

Lines changed: 8 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ import Control.Applicative (liftA2)
1414
import Control.Monad (liftM2)
1515
import Control.Monad.Except (MonadError (..), runExceptT)
1616
import Vehicle.Compile.Constants.ForcedValue (TensorValueLinearExpr)
17+
import Vehicle.Compile.Error
1718
import Vehicle.Compile.Normalise.Force
1819
import Vehicle.Compile.Normalise.RewriteRules (forceAndRewriteTensor)
1920
import Vehicle.Compile.Normalise.TypedValue
@@ -281,10 +282,6 @@ tryAndUnblock dims expr = do
281282
Right unblocked -> forIfTreeM unblocked $ \unblockedExpr ->
282283
compileLinearExpr dims unblockedExpr
283284

284-
data BlockingReason
285-
= BlockingNetwork Identifier
286-
| BlockingDatasetOrParameter Identifier
287-
288285
unblockingActions ::
289286
(MonadPurifyAssertion m, MonadError BlockingReason m) =>
290287
UnblockingActions m
@@ -293,18 +290,19 @@ unblockingActions =
293290
{ unblockRatTensorBoundVar = purifyBoundVar,
294291
unblockRecordBoundVar = purifyBoundVar,
295292
unblockNetworkApp = \_ _ ident _ -> throwError $ BlockingNetwork ident,
296-
unblockDatasetOrParameter = \ident -> throwError $ BlockingDatasetOrParameter ident
293+
unblockDatasetOrParameter = \_ ident -> throwError $ BlockingDatasetOrParameter ident
297294
}
298295

299296
purifyBoundVar ::
300-
(MonadLogger m, MonadReadableTensorBoundContext m) =>
297+
(MonadPurifyAssertion m) =>
298+
TypeUnblockingFunction (Thunk Builtin) m ->
301299
Lv ->
302-
m (Thunk Builtin)
303-
purifyBoundVar lv = do
300+
m (IfTree (Thunk Builtin) (Thunk Builtin))
301+
purifyBoundVar unblock lv = do
304302
(_, maybeChildVars) <- lookupVariableInNestedCtx lv
305303
case maybeChildVars of
306-
Nothing -> return $ Forced $ VBoundVar lv []
307-
Just (_tensorVar, sliceVar) -> replaceTensorVariableWithStackedChildren sliceVar
304+
Nothing -> return $ IfLeaf $ Forced $ VBoundVar lv []
305+
Just (_tensorVar, sliceVar) -> unblock =<< replaceTensorVariableWithStackedChildren sliceVar
308306

309307
--------------------------------------------------------------------------------
310308
-- Utility functions

vehicle/src/Vehicle/Backend/Loss/JSON.hs

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ module Vehicle.Backend.Loss.JSON
77
)
88
where
99

10+
import Control.Monad.Except (MonadError (..))
1011
import Data.Aeson (ToJSON (..), genericToJSON)
1112
import Data.List (elemIndex)
1213
import Data.Map (Map)
@@ -35,7 +36,7 @@ import Vehicle.Data.Code.Interface.Args
3536
import Vehicle.Data.Tensor (ExtendedRatTensor)
3637
import Vehicle.Data.Variable.Bound.Context.Name
3738
import Vehicle.Data.Variable.Free.Context (MonadFreeContext, addDeclToContext, runFreshFreeContextT)
38-
import Vehicle.Prelude (Doc, GenericArg (..), HasName (..), HasType (..), Identifier (..), Name, Provenance, explicit, indent, jsonOptions, line, mkExplicitBinder, resolutionError, squotes, userModulePath)
39+
import Vehicle.Prelude (Doc, GenericArg (..), HasName (..), HasType (..), Identifier (..), Name, Provenance, explicit, indent, jsonOptions, line, mkExplicitBinder, resolutionError, squotes, stdlibIdentifier, userModulePath)
3940
import Vehicle.Prelude.Error (developerError)
4041
import Vehicle.Prelude.Logging.Class
4142

@@ -166,8 +167,8 @@ type MonadJSON m =
166167
MonadFreeContext LossBuiltin m
167168
)
168169

169-
unsupportedError :: (Pretty a) => a -> b
170-
unsupportedError b = developerError $ "Conversion of" <+> pretty b <+> "is not yet implemented"
170+
unsupportedError :: (MonadJSON m, Pretty a) => a -> m b
171+
unsupportedError a = throwError $ UnsupportedLossOperation (stdlibIdentifier "unknown", mempty) (pretty a)
171172

172173
dependentTypesError :: (Pretty a) => a -> b
173174
dependentTypesError b = developerError $ "Conversion of" <+> pretty b <+> "is not yet implemented"

vehicle/src/Vehicle/Backend/Loss/LogicCompilation.hs

Lines changed: 13 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ module Vehicle.Backend.Loss.LogicCompilation
66
where
77

88
import Control.Monad (foldM)
9-
import Control.Monad.Except (MonadError (..))
9+
import Control.Monad.Except (MonadError (..), runExceptT)
1010
import Control.Monad.State (MonadState, StateT, execStateT, modify)
1111
import Data.Map (Map)
1212
import Data.Map qualified as Map
@@ -17,12 +17,11 @@ import Vehicle.Backend.Loss.Core hiding (lookupLogicField)
1717
import Vehicle.Backend.Loss.LossCompilation (convertQuantifierlessExprToLoss)
1818
import Vehicle.Backend.Prelude (DifferentiableLogicID)
1919
import Vehicle.Compile.Error
20-
import Vehicle.Compile.Normalise.Builtin
21-
import Vehicle.Compile.Normalise.Core
2220
import Vehicle.Compile.Normalise.Force
2321
import Vehicle.Compile.Normalise.Quote (unnormalise)
2422
import Vehicle.Compile.Prelude
2523
import Vehicle.Compile.Print
24+
import Vehicle.Compile.Unblock (noUnblocking, unblockBoolExpr)
2625
import Vehicle.Data.Builtin.Interface (Accessor (..))
2726
import Vehicle.Data.Builtin.Loss (ComparisonOp (..), LogicDirection, LossBuiltin)
2827
import Vehicle.Data.Builtin.Standard (Builtin)
@@ -153,13 +152,17 @@ calculateLogicDirection ::
153152
OMap FieldName (Thunk Builtin) ->
154153
m LogicDirection
155154
calculateLogicDirection declProv fields = do
156-
let trueValue = lookupLogicField TruthityElement fields
157-
let falseValue = lookupLogicField FalsityElement fields
158-
let args = TensorComparisonArgs (Forced IDimNil) (Forced IDimNil) trueValue falseValue
159-
result <- runFreshNameBoundContextT $ evalCompareRatTensor @ForcedValue Le args
160-
case result of
161-
Evaluated (Forced (IBoolLiteral b)) -> return b
162-
_ -> throwError $ UnorderableDifferentiableLogic declProv (mkExpr accessCompareRatTensor (Le, args))
155+
let expr = do
156+
let trueValue = lookupLogicField TruthityElement fields
157+
let falseValue = lookupLogicField FalsityElement fields
158+
let args = TensorComparisonArgs (Forced IDimNil) (Forced IDimNil) trueValue falseValue
159+
Forced $ mkExpr accessCompareRatTensor (Le, args)
160+
161+
errorOrResult <- runExceptT $ runFreshNameBoundContextT $ forceThunk =<< unblockBoolExpr noUnblocking expr
162+
case errorOrResult of
163+
Left blockingErr -> throwError $ UnorderableDifferentiableLogic declProv expr (Left blockingErr)
164+
Right (IBoolLiteral b) -> return b
165+
Right result -> throwError $ UnorderableDifferentiableLogic declProv expr (Right result)
163166

164167
compileLogicField ::
165168
(MonadLoss m) =>

vehicle/src/Vehicle/Backend/Loss/LossCompilation.hs

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,9 @@ module Vehicle.Backend.Loss.LossCompilation
33
convertBoolTensorLiteral,
44
convertQuantifierlessExprToLoss,
55
convertBooleanOp,
6+
orLossValue,
7+
andLossValue,
8+
notLossValue,
69
)
710
where
811

@@ -269,6 +272,24 @@ convertBoolTensorLiteral tensor = do
269272
--------------------------------------------------------------------------------
270273
-- Utils
271274

275+
orLossValue :: (MonadLogic m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
276+
orLossValue e1 e2 =
277+
convertBooleanOp PointwiseDisjunction $
278+
mkExpr accessSpine (TensorOp2Args (Forced IDimNil) e1 e2)
279+
280+
andLossValue :: (MonadLogic m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
281+
andLossValue e1 e2 =
282+
convertBooleanOp PointwiseConjunction $
283+
mkExpr accessSpine (TensorOp2Args (Forced IDimNil) e1 e2)
284+
285+
notLossValue :: (MonadLogic m) => Thunk LossBuiltin -> Thunk LossBuiltin -> m (Thunk LossBuiltin)
286+
notLossValue dims e =
287+
convertBooleanOp PointwiseNegation $
288+
mkExpr accessSpine (TensorOp1Args dims e)
289+
290+
--------------------------------------------------------------------------------
291+
-- Utils
292+
272293
currentPass :: Doc a
273294
currentPass = "logic translation"
274295

vehicle/src/Vehicle/Backend/Solver.hs

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ import Vehicle.Compile.Print (prettyFriendly, prettyFriendlyEmptyCtx)
3030
import Vehicle.Compile.Print.Warning ()
3131
import Vehicle.Compile.Property (traverseMultiProperty)
3232
import Vehicle.Compile.Unblock (UnblockingActions (..), unblockBoolExpr)
33+
import Vehicle.Data.Builtin.Interface (Accessor (..))
3334
import Vehicle.Data.Builtin.Standard
3435
import Vehicle.Data.Code.BooleanExpr
3536
import Vehicle.Data.Code.ForcedValue
@@ -317,10 +318,10 @@ compileQuerySetPartitions globalCtx isPropertyNegated maybePartitions = case may
317318
topLevelUnblockingActions :: (Monad m) => UnblockingActions m
318319
topLevelUnblockingActions =
319320
UnblockingActions
320-
{ unblockRatTensorBoundVar = developerError "No bound variables should exist at top-level",
321-
unblockRecordBoundVar = developerError "No bound variables should exist at top-level",
322-
unblockNetworkApp = \_ _ _ -> developerError "Unblocking of constant network functions at top-level not yet supported",
323-
unblockDatasetOrParameter = developerError "Should not be unblocking datasets or parameters"
321+
{ unblockRatTensorBoundVar = \_ _ -> developerError "No bound variables should exist at top-level",
322+
unblockRecordBoundVar = \_ _ -> developerError "No bound variables should exist at top-level",
323+
unblockNetworkApp = \_ _ ident args -> return $ IfLeaf $ Forced $ VFreeVar ident (mkExpr accessSpine args),
324+
unblockDatasetOrParameter = \_ _ -> developerError "Should not be unblocking datasets or parameters"
324325
}
325326

326327
handlePropertyCompileError ::

vehicle/src/Vehicle/Backend/Solver/UserVariableElimination.hs

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,7 @@ import Vehicle.Compile.Unblock qualified as Unblocking
3838
import Vehicle.Compile.Variable (createUserVar)
3939
import Vehicle.Data.Builtin.Interface
4040
import Vehicle.Data.Builtin.Standard
41-
import Vehicle.Data.Code.BooleanExpr (elimIfTree)
41+
import Vehicle.Data.Code.BooleanExpr (IfTree, elimIfTree)
4242
import Vehicle.Data.Code.ForcedValue
4343
import Vehicle.Data.Code.Interface
4444
import Vehicle.Data.MaybeTrivial
@@ -271,15 +271,16 @@ unblockingActions =
271271
{ unblockRatTensorBoundVar = unblockQuantifiedBoundVar,
272272
unblockRecordBoundVar = unblockQuantifiedBoundVar,
273273
unblockNetworkApp = unblockNetworkApplication,
274-
unblockDatasetOrParameter = unexpectedExprError "solver compilation" "dataset or parameter"
274+
unblockDatasetOrParameter = \_ _ -> unexpectedExprError "solver compilation" "dataset or parameter"
275275
}
276276

277277
unblockQuantifiedBoundVar ::
278278
(MonadQuantifierBody m) =>
279+
TypeUnblockingFunction (Thunk Builtin) m ->
279280
Lv ->
280-
m (Thunk Builtin)
281-
unblockQuantifiedBoundVar lv =
282-
replaceTensorVariableWithStackedChildren (SliceVariable lv)
281+
m (IfTree (Thunk Builtin) (Thunk Builtin))
282+
unblockQuantifiedBoundVar unblock lv =
283+
unblock =<< replaceTensorVariableWithStackedChildren (SliceVariable lv)
283284

284285
unblockNetworkApplication ::
285286
(MonadQuantifierBody m) =>

vehicle/src/Vehicle/Backend/Solver/UserVariableElimination/PurifyAssertion.hs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -104,7 +104,7 @@ purifyBoundVar ::
104104
Lv ->
105105
PurifyFn m
106106
purifyBoundVar v actions@UnblockingActions {..} incrDims
107-
| incrDims > 0 = purifyExpr actions incrDims =<< unblockRatTensorBoundVar v
107+
| incrDims > 0 = unblockRatTensorBoundVar (purifyExpr actions incrDims) v
108108
| otherwise = return $ IfLeaf $ Forced $ VBoundVar v []
109109

110110
purifyNetworkVar ::

0 commit comments

Comments
 (0)