diff --git a/ChangeLog.md b/ChangeLog.md index ba9a84e..cbe3bf8 100644 --- a/ChangeLog.md +++ b/ChangeLog.md @@ -2,6 +2,8 @@ ## Unreleased changes +- Share shift and split coefficient-map transformations between objectives and constraints, and simplify access to their fields. ([#23](https://github.com/rasheedja/simplex-method/pull/23)) + - Avoid ambiguous constraint record updates using the existing field lenses, group the substitution helpers together, and test them directly. - Use standard Data.Map selection and deletion operations for solver variable filtering. ([#20](https://github.com/rasheedja/simplex-method/pull/20)) - Remove unused pivot bindings and give the active row-update helper a descriptive name. ([#21](https://github.com/rasheedja/simplex-method/pull/21)) - Extract optimal variable values directly from dictionary constants without an intermediate tableau conversion. ([#22](https://github.com/rasheedja/simplex-method/pull/22)) diff --git a/src/Linear/Simplex/Solver/TwoPhase.hs b/src/Linear/Simplex/Solver/TwoPhase.hs index 47f7c1e..b49763a 100644 --- a/src/Linear/Simplex/Solver/TwoPhase.hs +++ b/src/Linear/Simplex/Solver/TwoPhase.hs @@ -28,6 +28,8 @@ module Linear.Simplex.Solver.TwoPhase , applyShiftToConstraint , applySplitToObjective , applySplitToConstraint + , shiftVarInMap + , splitVarInMap , unapplyTransformsToVarMap , unapplyTransformToVarMap ) where @@ -477,10 +479,7 @@ postprocess originalVars transforms (Optimal varVals) = -- | Compute the value of an objective function given variable values. computeObjective :: ObjectiveFunction -> M.Map Var SimplexNum -> SimplexNum computeObjective objFunction varVals = - let coeffs = case objFunction of - Max m -> m - Min m -> m - in sum $ map (\(var, coeff) -> coeff * M.findWithDefault 0 var varVals) (M.toList coeffs) + sum $ map (\(var, coeff) -> coeff * M.findWithDefault 0 var varVals) (M.toList objFunction.objective) -- | Preprocess the system by applying variable transformations based on domain information. -- Returns the transformed objectives, constraints, and the list of transforms applied. @@ -512,18 +511,9 @@ applyTransformsToConstraints transforms constraints = -- | Collect all variables appearing in the objective functions and constraints collectAllVars :: [ObjectiveFunction] -> [PolyConstraint] -> Set Var collectAllVars objFunctions constraints = - let objVars = Set.unions $ map getObjVars objFunctions - constraintVars = Set.unions $ map getConstraintVars constraints - in Set.union objVars constraintVars - where - getObjVars :: ObjectiveFunction -> Set Var - getObjVars (Max m) = M.keysSet m - getObjVars (Min m) = M.keysSet m - - getConstraintVars :: PolyConstraint -> Set Var - getConstraintVars (LEQ m _) = M.keysSet m - getConstraintVars (GEQ m _) = M.keysSet m - getConstraintVars (EQ m _) = M.keysSet m + Set.unions $ + map (M.keysSet . (.objective)) objFunctions + ++ map (M.keysSet . (.lhs)) constraints -- | Generate a transform for a variable based on its domain. -- Takes the domain map, the variable, and the current (transforms, nextFreshVar). @@ -600,15 +590,7 @@ applyTransform transform (objFunction, constraints) = -- The constant term changes but objectives don't have constants that affect optimization. applyShiftToObjective :: Var -> Var -> SimplexNum -> ObjectiveFunction -> ObjectiveFunction applyShiftToObjective origVar shiftedVar _shiftBy objFunction = - case objFunction of - Max m -> Max (substituteVar origVar shiftedVar m) - Min m -> Min (substituteVar origVar shiftedVar m) - where - substituteVar :: Var -> Var -> VarLitMapSum -> VarLitMapSum - substituteVar oldVar newVar m = - case M.lookup oldVar m of - Nothing -> m - Just coeff -> M.insert newVar coeff (M.delete oldVar m) + objFunction {objective = fst $ shiftVarInMap origVar shiftedVar 0 objFunction.objective} -- | Apply shift transformation to a constraint. -- originalVar = shiftedVar + shiftBy @@ -618,53 +600,38 @@ applyShiftToObjective origVar shiftedVar _shiftBy objFunction = -- So new constraint: (replace originalVar with shiftedVar) REL (rhs - c_j * shiftBy) applyShiftToConstraint :: Var -> Var -> SimplexNum -> PolyConstraint -> PolyConstraint applyShiftToConstraint origVar shiftedVar shiftBy constraint = - case constraint of - LEQ m rhs -> - let (newMap, rhsAdjust) = substituteVarInMap origVar shiftedVar shiftBy m - in LEQ newMap (rhs - rhsAdjust) - GEQ m rhs -> - let (newMap, rhsAdjust) = substituteVarInMap origVar shiftedVar shiftBy m - in GEQ newMap (rhs - rhsAdjust) - EQ m rhs -> - let (newMap, rhsAdjust) = substituteVarInMap origVar shiftedVar shiftBy m - in EQ newMap (rhs - rhsAdjust) - where - substituteVarInMap :: Var -> Var -> SimplexNum -> VarLitMapSum -> (VarLitMapSum, SimplexNum) - substituteVarInMap oldVar newVar shift m = - case M.lookup oldVar m of - Nothing -> (m, 0) - Just coeff -> (M.insert newVar coeff (M.delete oldVar m), coeff * shift) + let (newMap, rhsAdjust) = shiftVarInMap origVar shiftedVar shiftBy constraint.lhs + in constraint & #lhs .~ newMap & #rhs %~ subtract rhsAdjust -- | Apply split transformation to objective function. -- originalVar = posVar - negVar -- coefficient c of originalVar becomes c for posVar and -c for negVar applySplitToObjective :: Var -> Var -> Var -> ObjectiveFunction -> ObjectiveFunction applySplitToObjective origVar posVar negVar objFunction = - case objFunction of - Max m -> Max (splitVar origVar posVar negVar m) - Min m -> Min (splitVar origVar posVar negVar m) - where - splitVar :: Var -> Var -> Var -> VarLitMapSum -> VarLitMapSum - splitVar oldVar pVar nVar m = - case M.lookup oldVar m of - Nothing -> m - Just coeff -> M.insert pVar coeff (M.insert nVar (-coeff) (M.delete oldVar m)) + objFunction {objective = splitVarInMap origVar posVar negVar objFunction.objective} -- | Apply split transformation to a constraint. -- originalVar = posVar - negVar -- coefficient c of originalVar becomes c for posVar and -c for negVar applySplitToConstraint :: Var -> Var -> Var -> PolyConstraint -> PolyConstraint applySplitToConstraint origVar posVar negVar constraint = - case constraint of - LEQ m rhs -> LEQ (splitVarInMap origVar posVar negVar m) rhs - GEQ m rhs -> GEQ (splitVarInMap origVar posVar negVar m) rhs - EQ m rhs -> EQ (splitVarInMap origVar posVar negVar m) rhs - where - splitVarInMap :: Var -> Var -> Var -> VarLitMapSum -> VarLitMapSum - splitVarInMap oldVar pVar nVar m = - case M.lookup oldVar m of - Nothing -> m - Just coeff -> M.insert pVar coeff (M.insert nVar (-coeff) (M.delete oldVar m)) + constraint & #lhs %~ splitVarInMap origVar posVar negVar + +-- | Substitute a shifted variable and return the constant offset introduced. +-- The replacement variable must be fresh. +shiftVarInMap :: Var -> Var -> SimplexNum -> VarLitMapSum -> (VarLitMapSum, SimplexNum) +shiftVarInMap oldVar newVar shift coeffs = + case M.lookup oldVar coeffs of + Nothing -> (coeffs, 0) + Just coeff -> (M.insert newVar coeff (M.delete oldVar coeffs), coeff * shift) + +-- | Substitute oldVar = posVar - negVar in a coefficient map. +-- The replacement variables must be distinct and fresh. +splitVarInMap :: Var -> Var -> Var -> VarLitMapSum -> VarLitMapSum +splitVarInMap oldVar posVar negVar coeffs = + case M.lookup oldVar coeffs of + Nothing -> coeffs + Just coeff -> M.insert posVar coeff (M.insert negVar (-coeff) (M.delete oldVar coeffs)) -- | Unapply transforms to convert a variable value map back to original variables. unapplyTransformsToVarMap :: [VarTransform] -> VarLitMap -> VarLitMap diff --git a/test/Linear/Simplex/Solver/TwoPhaseSpec.hs b/test/Linear/Simplex/Solver/TwoPhaseSpec.hs index 9ca2246..84debe3 100644 --- a/test/Linear/Simplex/Solver/TwoPhaseSpec.hs +++ b/test/Linear/Simplex/Solver/TwoPhaseSpec.hs @@ -28,6 +28,8 @@ import Linear.Simplex.Solver.TwoPhase , getTransform , postprocess , preprocess + , shiftVarInMap + , splitVarInMap , twoPhaseSimplex , unapplyTransformToVarMap , unapplyTransformsToVarMap @@ -2515,6 +2517,66 @@ spec = do let constraint = LEQ (M.fromList [(1, 2), (2, 3)]) 10 applySplitToConstraint 1 10 11 constraint `shouldBe` LEQ (M.fromList [(10, 2), (11, -2), (2, 3)]) 10 + describe "shiftVarInMap" $ do + it "leaves an empty map unchanged with no offset" $ do + shiftVarInMap 1 10 (-5) M.empty `shouldBe` (M.empty, 0) + + it "leaves unrelated variables unchanged with no offset" $ do + let coeffs = M.fromList [(2, 5), (3, -7)] + shiftVarInMap 1 10 (-5) coeffs `shouldBe` (coeffs, 0) + + mapM_ + ( \(label, coeff, shift, offset) -> + it ("substitutes a " ++ label ++ " coefficient and returns its exact offset") $ do + let coeffs = M.fromList [(1, coeff), (2, 5), (3, -7)] + shiftVarInMap 1 10 shift coeffs + `shouldBe` (M.fromList [(10, coeff), (2, 5), (3, -7)], offset) + ) + [ ("positive", 3, -5, -15) + , ("negative", -3, -5, 15) + , ("fractional", 2 % 3, -5 % 7, -10 % 21) + , ("zero", 0, -5, 0) + ] + + it "renames the variable even when the shift is zero" $ do + shiftVarInMap 1 10 0 (M.singleton 1 (2 % 3)) + `shouldBe` (M.singleton 10 (2 % 3), 0) + + it "preserves evaluation under originalVar = shiftedVar + shift" $ + property $ + \(coeff :: Rational) (otherCoeff :: Rational) (shift :: Rational) (shiftedValue :: Rational) (otherValue :: Rational) -> + let coeffs = M.fromList [(1, coeff), (2, otherCoeff)] + (shiftedCoeffs, offset) = shiftVarInMap 1 10 shift coeffs + values = M.fromList [(10, shiftedValue), (2, otherValue)] + in sum (M.intersectionWith (*) shiftedCoeffs values) + offset + == coeff * (shiftedValue + shift) + otherCoeff * otherValue + + describe "splitVarInMap" $ do + it "leaves an empty map unchanged" $ do + splitVarInMap 1 10 11 M.empty `shouldBe` M.empty + + it "leaves unrelated variables unchanged" $ do + let coeffs = M.fromList [(2, 5), (3, -7)] + splitVarInMap 1 10 11 coeffs `shouldBe` coeffs + + mapM_ + ( \(label, coeff) -> + it ("splits a " ++ label ++ " coefficient into opposite signed terms") $ do + let coeffs = M.fromList [(1, coeff), (2, 5), (3, -7)] + splitVarInMap 1 10 11 coeffs + `shouldBe` M.fromList [(10, coeff), (11, -coeff), (2, 5), (3, -7)] + ) + [("positive", 3), ("negative", -3), ("fractional", 2 % 3), ("zero", 0)] + + it "preserves evaluation under originalVar = positiveVar - negativeVar" $ + property $ + \(coeff :: Rational) (otherCoeff :: Rational) (positiveValue :: Rational) (negativeValue :: Rational) (otherValue :: Rational) -> + let coeffs = M.fromList [(1, coeff), (2, otherCoeff)] + splitCoeffs = splitVarInMap 1 10 11 coeffs + values = M.fromList [(10, positiveValue), (11, negativeValue), (2, otherValue)] + in sum (M.intersectionWith (*) splitCoeffs values) + == coeff * (positiveValue - negativeValue) + otherCoeff * otherValue + describe "applyTransform and applyTransforms" $ do describe "Unit tests" $ do it "applyTransform AddLowerBound adds GEQ constraint" $ do