Skip to content

Commit 6e47f12

Browse files
Fix #1213 by unblocking rather than evaluating comparison (#1215)
* Fix #1213 by unblocking rather evaluating comparison * Fix test error
1 parent f5be6da commit 6e47f12

14 files changed

Lines changed: 136 additions & 62 deletions

File tree

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@ import Data.List.NonEmpty (NonEmpty (..))
1313
import Data.Map (Map)
1414
import Data.Map qualified as Map
1515
import Vehicle.Backend.Loss.Core
16-
import Vehicle.Backend.Loss.Domain.PurifyAssertion (BlockingReason, tryPurifyAssertion, unblockingActions)
16+
import Vehicle.Backend.Loss.Domain.PurifyAssertion (tryPurifyAssertion, unblockingActions)
1717
import Vehicle.Backend.Loss.LossCompilation
1818
import Vehicle.Backend.Solver.UserVariableElimination.ConstraintSearch (findAllBounds)
1919
import Vehicle.Compile.Constants.ForcedValue

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/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/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 ::

vehicle/src/Vehicle/Compile/Error.hs

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@ module Vehicle.Compile.Error
1818
UnboundedIndices,
1919
ParseLocation,
2020
MonadCompile,
21+
BlockingReason (..),
2122
compilerDeveloperError,
2223
unsupportedTensorLikeQuantifier,
2324
)
@@ -230,7 +231,7 @@ data CompileError
230231
| UnsupportedLossOperation DeclProvenance (Doc Void)
231232
| UnableToLiftLogicFieldToTensors DifferentiableLogicID TensorDifferentiableLogicField (BooleanDifferentiableLogicField, Thunk Builtin) NamedBoundCtx (Thunk Builtin)
232233
| NoQuantifierDomainFound DeclProvenance (UnforcedBinder Builtin) (These (NonEmpty TensorIndices) (NonEmpty TensorIndices))
233-
| UnorderableDifferentiableLogic DeclProvenance (ForcedValue Builtin)
234+
| UnorderableDifferentiableLogic DeclProvenance (Thunk Builtin) (Either BlockingReason (ForcedValue Builtin))
234235
| -- ITP backend errors
235236
UnimplementedFeature Provenance (Doc Void)
236237
| UnusedMonomorphisableDeclaration Provenance Identifier
@@ -240,6 +241,12 @@ data CompileError
240241

241242
deriving instance Show CompileError
242243

244+
data BlockingReason
245+
= BlockingNetwork Identifier
246+
| BlockingDatasetOrParameter Identifier
247+
248+
deriving instance Show BlockingReason
249+
243250
--------------------------------------------------------------------------------
244251
-- Some useful developer errors
245252

vehicle/src/Vehicle/Compile/Normalise/Core.hs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -48,11 +48,11 @@ data EvalScheme meta builtin m
4848
| TypeClassOp
4949
| None
5050

51-
class (Monad m, HasBuiltinConstructor expr thunk) => NormalisableExpr expr thunk builtin m | thunk -> expr where
51+
class (Monad m, HasBuiltinConstructor expr thunk, Show (expr builtin)) => NormalisableExpr expr thunk builtin m | thunk -> expr where
5252
force :: thunk builtin -> m (expr builtin)
5353
forceApp :: thunk builtin -> [GenericArg (thunk builtin)] -> m (expr builtin)
5454

55-
instance (Monad m) => NormalisableExpr Expr Expr builtin m where
55+
instance (Monad m, Show builtin) => NormalisableExpr Expr Expr builtin m where
5656
force = return
5757
forceApp fun args = return $ normAppList fun args
5858

vehicle/src/Vehicle/Compile/Normalise/Force.hs

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -112,7 +112,8 @@ instance
112112

113113
-- Merge into `TypedEvalScheme`?
114114
instance
115-
( MonadNorm builtin m,
115+
( Show meta,
116+
MonadNorm builtin m,
116117
TypedEvalScheme meta builtin m
117118
) =>
118119
NormalisableExpr (GenericForcedValue meta) (GenericThunk meta) builtin m

vehicle/src/Vehicle/Compile/Print/Error.hs

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -983,7 +983,7 @@ formatCompileError = \case
983983
Just
984984
"declare the logic directly as a record literal."
985985
}
986-
UnorderableDifferentiableLogic (ident, p) value ->
986+
UnorderableDifferentiableLogic (ident, p) expr reason ->
987987
VehicleUserError
988988
{ provenance = Just p,
989989
problem =
@@ -995,14 +995,21 @@ formatCompileError = \case
995995
<> line
996996
<> "in order to work out whether the loss should be maximised or minimised."
997997
<> line
998-
<> "However, Vehicle was unable to establish the truth value of the result:"
999-
<+> lineIndent (prettyFriendlyEmptyCtx value),
998+
<> "However, Vehicle was unable to establish the truth value of"
999+
<> lineIndent (prettyFriendlyEmptyCtx expr)
1000+
<> line
1001+
<> "because it could not evaluate" <+> case reason of
1002+
Right value ->
1003+
":"
1004+
<+> lineIndent (prettyFriendlyEmptyCtx value)
1005+
Left (BlockingDatasetOrParameter blockingIdent) -> quotePretty (nameOf blockingIdent)
1006+
Left (BlockingNetwork blockingIdent) -> quotePretty (nameOf blockingIdent),
10001007
fix =
10011008
Just $
10021009
"ensure that the expression" <+> squotes comp <+> "evaluates to either `true` or `false`."
10031010
}
10041011
where
1005-
comp = pretty TruthityElement <+> "<=" <+> pretty FalsityElement
1012+
comp = pretty TruthityElement <+> "<" <+> pretty FalsityElement
10061013

10071014
datasetDimensionsFix :: Doc a -> Identifier -> FilePath -> Doc a
10081015
datasetDimensionsFix feature ident file =

0 commit comments

Comments
 (0)