如何在Haskell优化编译器中正确实现常量折叠算法?
问题描述
我正在阅读Steven Muchnick所著《Advanced Compiler Design and Implementation》一书,其优化章节开篇即介绍常量折叠。我最初的实现尝试如下,但因需要Integral和Bits类型类而几乎无法完成:
data Value = BoolV Bool | IntegerV Integer | FloatV Double | StringV String deriving (Show, Eq) data Expression = Const Value | Temp Label | List [Expression] | Call Label [Expression] | BinOp Operator Expression Expression deriving (Show, Eq) rerunOverIR :: Expression -> Expression rerunOverIR = \case Const constant -> Const constant Temp temporary -> Temp temporary List list -> List list Call label args -> Call label (map rerunOverIR args) BinOp operator lhs rhs -> case operator of Addition -> folder (+) Subtraction -> folder (-) Multiplication -> folder (*) Modulo -> folder mod Division -> folder div BitwiseAnd -> folder (.&.) BitwiseOr -> folder (.|.) BitwiseXor -> folder xor ShiftRight -> folder shiftR ShiftLeft -> folder shiftL Equal -> folder (==) NotEqual -> folder (/=) Greater -> folder (>) GreaterEqual -> folder (>=) Less -> folder (<) LessEqual -> folder (<=) LogicalAnd -> folder (&&) LogicalOr -> folder (||) _ -> error $ "this operator doesn't exist " ++ show operator where folder op = case lhs of Const c1 -> case rhs of Const c2 -> Const $ op c1 c2 expr -> rerunOverIR $ BinOp operator (Const c1) (rerunOverIR expr) e1 -> case rhs of Const c2 -> rerunOverIR $ BinOp operator (rerunOverIR e1) (Const c2) e2 -> rerunOverIR $ BinOp operator (rerunOverIR e1) (rerunOverIR e2)
随后我尝试修改Expression定义如下,但情况更糟:
data Expression = Bool Bool | Integer Integer | Float Double | String String | Temp Label | List [Expression] | Call Label [Expression] | BinOp Operator Expression Expression deriving (Show, Eq)
我的问题是:用Haskell编写的编译器或解释器在后期阶段如何正确处理常量折叠?我确信自己的思路有误。
解决思路
核心问题在于试图用多态函数直接操作异构的Value类型,Haskell的类型系统不允许这种操作——比如(+)无法直接作用于Value,因为它无法判断你要处理整数还是浮点数。正确的做法是为每种运算符单独处理类型匹配,而非依赖通用的folder函数。
1. 明确运算符的类型约束
首先要清晰划分运算符的适用类型:
- 位运算(
.&.、shiftR等)仅支持整数 - 逻辑运算(
&&、||)仅支持布尔值 - 算术运算(
+、-)支持整数或浮点数 - 比较运算(
==、>)支持同类型的任意值
2. 重构常量折叠逻辑
将递归遍历IR、常量折叠判断、类型安全运算三个环节拆分,代码更清晰且类型安全:
-- 补全Operator定义(根据你的实际场景调整) data Operator = Addition | Subtraction | Multiplication | Modulo | Division | BitwiseAnd | BitwiseOr | BitwiseXor | ShiftRight | ShiftLeft | Equal | NotEqual | Greater | GreaterEqual | Less | LessEqual | LogicalAnd | LogicalOr deriving (Show, Eq) data Value = BoolV Bool | IntegerV Integer | FloatV Double | StringV String deriving (Show, Eq) -- 保留Const Value的封装,区分常量与其他表达式节点 data Expression = Const Value | Temp String -- 假设Label为String,可替换为实际类型 | List [Expression] | Call String [Expression] | BinOp Operator Expression Expression deriving (Show, Eq) -- 递归遍历并处理整个IR rerunOverIR :: Expression -> Expression rerunOverIR = \case Const v -> Const v Temp l -> Temp l List es -> List (map rerunOverIR es) Call l args -> Call l (map rerunOverIR args) BinOp op lhs rhs -> let lhs' = rerunOverIR lhs rhs' = rerunOverIR rhs in foldBinOp op lhs' rhs' -- 专门处理二元运算符的常量折叠 foldBinOp :: Operator -> Expression -> Expression -> Expression foldBinOp op (Const c1) (Const c2) = Const $ applyOp op c1 c2 foldBinOp op lhs rhs = BinOp op lhs rhs -- 非常量组合直接返回原结构 -- 核心:根据运算符与值的类型,执行类型安全的运算 applyOp :: Operator -> Value -> Value -> Value -- 算术运算:整数/浮点数分支 applyOp Addition (IntegerV a) (IntegerV b) = IntegerV (a + b) applyOp Addition (FloatV a) (FloatV b) = FloatV (a + b) applyOp Subtraction (IntegerV a) (IntegerV b) = IntegerV (a - b) applyOp Subtraction (FloatV a) (FloatV b) = FloatV (a - b) applyOp Multiplication (IntegerV a) (IntegerV b) = IntegerV (a * b) applyOp Multiplication (FloatV a) (FloatV b) = FloatV (a * b) applyOp Modulo (IntegerV a) (IntegerV b) = IntegerV (a `mod` b) applyOp Division (IntegerV a) (IntegerV b) = IntegerV (a `div` b) applyOp Division (FloatV a) (FloatV b) = FloatV (a / b) -- 位运算:仅处理整数 applyOp BitwiseAnd (IntegerV a) (IntegerV b) = IntegerV (a .&. b) applyOp BitwiseOr (IntegerV a) (IntegerV b) = IntegerV (a .|. b) applyOp BitwiseXor (IntegerV a) (IntegerV b) = IntegerV (xor a b) applyOp ShiftRight (IntegerV a) (IntegerV b) = IntegerV (a `shiftR` fromIntegral b) applyOp ShiftLeft (IntegerV a) (IntegerV b) = IntegerV (a `shiftL` fromIntegral b) -- 逻辑运算:仅处理布尔值 applyOp LogicalAnd (BoolV a) (BoolV b) = BoolV (a && b) applyOp LogicalOr (BoolV a) (BoolV b) = BoolV (a || b) -- 比较运算:同类型值比较 applyOp Equal (IntegerV a) (IntegerV b) = BoolV (a == b) applyOp Equal (FloatV a) (FloatV b) = BoolV (a == b) applyOp Equal (BoolV a) (BoolV b) = BoolV (a == b) applyOp Equal (StringV a) (StringV b) = BoolV (a == b) applyOp NotEqual c1 c2 = BoolV $ not $ getBool (applyOp Equal c1 c2) where getBool (BoolV b) = b applyOp Greater (IntegerV a) (IntegerV b) = BoolV (a > b) applyOp Greater (FloatV a) (FloatV b) = BoolV (a > b) applyOp GreaterEqual (IntegerV a) (IntegerV b) = BoolV (a >= b) applyOp GreaterEqual (FloatV a) (FloatV b) = BoolV (a >= b) applyOp Less (IntegerV a) (IntegerV b) = BoolV (a < b) applyOp Less (FloatV a) (FloatV b) = BoolV (a < b) applyOp LessEqual (IntegerV a) (IntegerV b) = BoolV (a <= b) applyOp LessEqual (FloatV a) (FloatV b) = BoolV (a <= b) -- 类型不匹配时的错误处理(可根据需求改为保留原表达式) applyOp op c1 c2 = error $ "类型不匹配:运算符" ++ show op ++ "无法作用于" ++ show c1 ++ "和" ++ show c2
3. 方案优势
- 职责拆分清晰:
rerunOverIR负责递归遍历,foldBinOp负责判断是否可折叠,applyOp负责类型安全的运算实现,代码易维护。 - 类型安全:通过模式匹配确保每种运算符只处理合法的
Value类型,编译器会检查遗漏的分支。 - 扩展性强:新增运算符或值类型时,只需在
applyOp中添加对应的模式匹配即可。
4. 关于Expression定义的修改
直接将常量与表达式节点合并的做法(比如用Bool Bool代替Const BoolV Bool)会增加类型判断的复杂度——你需要在大量场景下区分某个Expression是常量还是临时变量/函数调用,反而降低了代码可读性。保持Const Value的封装更清晰,也更符合IR的分层设计。
内容的提问来源于stack exchange,提问作者NotAlfred
相关产品推荐
相关产品推荐

