@@ -6,7 +6,7 @@ module Vehicle.Backend.Loss.LogicCompilation
66where
77
88import Control.Monad (foldM )
9- import Control.Monad.Except (MonadError (.. ))
9+ import Control.Monad.Except (MonadError (.. ), runExceptT )
1010import Control.Monad.State (MonadState , StateT , execStateT , modify )
1111import Data.Map (Map )
1212import Data.Map qualified as Map
@@ -17,12 +17,11 @@ import Vehicle.Backend.Loss.Core hiding (lookupLogicField)
1717import Vehicle.Backend.Loss.LossCompilation (convertQuantifierlessExprToLoss )
1818import Vehicle.Backend.Prelude (DifferentiableLogicID )
1919import Vehicle.Compile.Error
20- import Vehicle.Compile.Normalise.Builtin
21- import Vehicle.Compile.Normalise.Core
2220import Vehicle.Compile.Normalise.Force
2321import Vehicle.Compile.Normalise.Quote (unnormalise )
2422import Vehicle.Compile.Prelude
2523import Vehicle.Compile.Print
24+ import Vehicle.Compile.Unblock (noUnblocking , unblockBoolExpr )
2625import Vehicle.Data.Builtin.Interface (Accessor (.. ))
2726import Vehicle.Data.Builtin.Loss (ComparisonOp (.. ), LogicDirection , LossBuiltin )
2827import Vehicle.Data.Builtin.Standard (Builtin )
@@ -153,13 +152,17 @@ calculateLogicDirection ::
153152 OMap FieldName (Thunk Builtin ) ->
154153 m LogicDirection
155154calculateLogicDirection 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
164167compileLogicField ::
165168 (MonadLoss m ) =>
0 commit comments