package shader import "core:fmt" // IR optimization passes optimize :: proc(module: ^IR_Module, level: Opt_Level) { if level == .None do return opt_constant_fold(module) opt_copy_propagation(module) opt_cse(module) opt_dead_code_elim(module) if level == .Aggressive { // Second pass catches new opportunities opt_constant_fold(module) opt_copy_propagation(module) opt_cse(module) opt_dead_code_elim(module) } } // -- Constant Folding -- opt_constant_fold :: proc(module: ^IR_Module) { for &fn in module.functions { opt_fold_stmts(fn.body[:]) } } @(private = "file") opt_fold_stmts :: proc(stmts: []IR_Stmt) { for &stmt in stmts { switch s in stmt { case ^IR_Let: s.value = opt_fold_expr(s.value) case ^IR_Assign: s.value = opt_fold_expr(s.value) case ^IR_Return: s.value = opt_fold_expr(s.value) case ^IR_Store_Output: s.value = opt_fold_expr(s.value) case ^IR_If: s.condition = opt_fold_expr(s.condition) opt_fold_stmts(s.then_body[:]) for &ei in s.elseif_clauses { ei.condition = opt_fold_expr(ei.condition) opt_fold_stmts(ei.body[:]) } opt_fold_stmts(s.else_body[:]) case ^IR_For: s.start = opt_fold_expr(s.start) s.stop = opt_fold_expr(s.stop) s.step = opt_fold_expr(s.step) opt_fold_stmts(s.body[:]) case ^IR_While: s.condition = opt_fold_expr(s.condition) opt_fold_stmts(s.body[:]) case ^IR_Expr_Stmt: s.expr = opt_fold_expr(s.expr) case ^IR_Barrier: // no-op case ^IR_Discard: // no-op case ^IR_Break: // no-op case ^IR_Continue: // no-op } } } @(private = "file") opt_fold_expr :: proc(expr: ^IR_Expr) -> ^IR_Expr { if expr == nil do return nil switch d in expr.derived { case ^IR_Binary: d.left = opt_fold_expr(d.left) d.right = opt_fold_expr(d.right) // Try to fold constant binary ops left_lit := get_literal(d.left) right_lit := get_literal(d.right) if left_lit != nil && right_lit != nil { if result, ok := fold_binary(d.op, left_lit, right_lit); ok { return make_literal_expr(result, expr.type) } } // Algebraic simplifications if right_lit != nil { // x * 1.0 -> x if d.op == .Mul { if is_float_one(right_lit) do return d.left if is_float_zero(right_lit) do return make_literal_expr(f64(0.0), expr.type) } // x + 0.0 -> x, x - 0.0 -> x if (d.op == .Add || d.op == .Sub) && is_float_zero(right_lit) { return d.left } } if left_lit != nil { // 1.0 * x -> x if d.op == .Mul && is_float_one(left_lit) { return d.right } // 0.0 * x -> 0 if d.op == .Mul && is_float_zero(left_lit) { return make_literal_expr(f64(0.0), expr.type) } // 0.0 + x -> x if d.op == .Add && is_float_zero(left_lit) { return d.right } } case ^IR_Unary: d.operand = opt_fold_expr(d.operand) lit := get_literal(d.operand) if lit != nil { if result, ok := fold_unary(d.op, lit); ok { return make_literal_expr(result, expr.type) } } case ^IR_Call: for &arg in d.args { arg = opt_fold_expr(arg) } case ^IR_Field_Access: d.object = opt_fold_expr(d.object) case ^IR_Swizzle: d.object = opt_fold_expr(d.object) case ^IR_Composite_Extract: d.object = opt_fold_expr(d.object) case ^IR_Vector_Shuffle: d.object = opt_fold_expr(d.object) case ^IR_Index: d.object = opt_fold_expr(d.object) d.index = opt_fold_expr(d.index) case ^IR_Construct: for &arg in d.args { arg = opt_fold_expr(arg) } case ^IR_Type_Cast: d.value = opt_fold_expr(d.value) case ^IR_Select: d.condition = opt_fold_expr(d.condition) d.true_val = opt_fold_expr(d.true_val) d.false_val = opt_fold_expr(d.false_val) case ^IR_Literal, ^IR_Var_Ref, ^IR_Load_Binding, ^IR_Input_Field, ^IR_Builtin_Var, ^IR_Shared_Ref: // leaf nodes, nothing to fold } return expr } @(private = "file") get_literal :: proc(expr: ^IR_Expr) -> ^IR_Literal { if expr == nil do return nil lit, ok := expr.derived.(^IR_Literal) return ok ? lit : nil } @(private = "file") is_float_zero :: proc(lit: ^IR_Literal) -> bool { switch v in lit.value { case f64: return v == 0.0 case i64: return v == 0 case bool: return false } return false } @(private = "file") is_float_one :: proc(lit: ^IR_Literal) -> bool { switch v in lit.value { case f64: return v == 1.0 case i64: return v == 1 case bool: return false } return false } @(private = "file") fold_binary :: proc(op: IR_Op, left, right: ^IR_Literal) -> (IR_Literal_Value, bool) { // Float-float ops lf, l_is_f := left.value.(f64) rf, r_is_f := right.value.(f64) if l_is_f && r_is_f { #partial switch op { case .Add: return lf + rf, true case .Sub: return lf - rf, true case .Mul: return lf * rf, true case .Div: if rf != 0 do return lf / rf, true case .Mod: if rf != 0 { // SPIR-V FMod semantics: result has same sign as divisor q := lf / rf // Truncate toward zero trunc_q: f64 = q >= 0 ? f64(i64(q)) : -f64(i64(-q)) return lf - trunc_q * rf, true } case .Eq: return lf == rf, true case .Neq: return lf != rf, true case .Lt: return lf < rf, true case .Gt: return lf > rf, true case .Lte: return lf <= rf, true case .Gte: return lf >= rf, true case: // fall through } } // Int-int ops li, l_is_i := left.value.(i64) ri, r_is_i := right.value.(i64) if l_is_i && r_is_i { #partial switch op { case .Add: return li + ri, true case .Sub: return li - ri, true case .Mul: return li * ri, true case .Div: if ri != 0 do return li / ri, true case .Mod: if ri != 0 do return li % ri, true case .Eq: return li == ri, true case .Neq: return li != ri, true case .Lt: return li < ri, true case .Gt: return li > ri, true case .Lte: return li <= ri, true case .Gte: return li >= ri, true case: // fall through } } // Bool-bool ops lb, l_is_b := left.value.(bool) rb, r_is_b := right.value.(bool) if l_is_b && r_is_b { #partial switch op { case .And: return lb && rb, true case .Or: return lb || rb, true case .Eq: return lb == rb, true case .Neq: return lb != rb, true case: // fall through } } return nil, false } @(private = "file") fold_unary :: proc(op: IR_Op, operand: ^IR_Literal) -> (IR_Literal_Value, bool) { #partial switch op { case .Neg: switch v in operand.value { case f64: return -v, true case i64: return -v, true case bool: return nil, false } case .Not: if b, ok := operand.value.(bool); ok { return !b, true } case: // fall through } return nil, false } @(private = "file") make_literal_expr :: proc(value: IR_Literal_Value, type: ^Resolved_Type) -> ^IR_Expr { lit := new(IR_Literal) lit.value = value e := new(IR_Expr) e.kind = .Literal e.type = type e.derived = lit return e } // -- Copy Propagation -- // For immutable `let x = y` where y is a simple Var_Ref, replace uses of x with y. // Conservative: invalidates all mappings at control flow boundaries. opt_copy_propagation :: proc(module: ^IR_Module) { for &fn in module.functions { copy_map := make(map[IR_Var_Id]IR_Var_Id) name_map := make(map[IR_Var_Id]string) // id -> name for updating name field defer delete(copy_map) defer delete(name_map) // Populate name_map from var_decls for decl in fn.var_decls { name_map[decl.id] = decl.name } opt_copy_prop_stmts(fn.body[:], ©_map, &name_map) } } @(private = "file") opt_copy_prop_stmts :: proc(stmts: []IR_Stmt, copy_map: ^map[IR_Var_Id]IR_Var_Id, name_map: ^map[IR_Var_Id]string) { for &stmt in stmts { switch s in stmt { case ^IR_Let: s.value = opt_copy_prop_expr(s.value, copy_map, name_map) // If value is a simple var ref, record the mapping if vr, ok := s.value.derived.(^IR_Var_Ref); ok { copy_map[s.id] = vr.id } case ^IR_Assign: s.value = opt_copy_prop_expr(s.value, copy_map, name_map) // Assignment to a variable invalidates any copies pointing to it if vr, ok := s.target.derived.(^IR_Var_Ref); ok { // Invalidate anything that copies FROM this var keys_to_remove := make([dynamic]IR_Var_Id, context.temp_allocator) for k, v in copy_map { if v == vr.id { append(&keys_to_remove, k) } } for k in keys_to_remove { delete_key(copy_map, k) } // Also invalidate if this var was a copy target delete_key(copy_map, vr.id) } case ^IR_Return: s.value = opt_copy_prop_expr(s.value, copy_map, name_map) case ^IR_Store_Output: s.value = opt_copy_prop_expr(s.value, copy_map, name_map) case ^IR_If: s.condition = opt_copy_prop_expr(s.condition, copy_map, name_map) // Conservative: don't propagate into/across branches then_map := make(map[IR_Var_Id]IR_Var_Id) defer delete(then_map) opt_copy_prop_stmts(s.then_body[:], &then_map, name_map) for &ei in s.elseif_clauses { ei.condition = opt_copy_prop_expr(ei.condition, copy_map, name_map) ei_map := make(map[IR_Var_Id]IR_Var_Id) defer delete(ei_map) opt_copy_prop_stmts(ei.body[:], &ei_map, name_map) } else_map := make(map[IR_Var_Id]IR_Var_Id) defer delete(else_map) opt_copy_prop_stmts(s.else_body[:], &else_map, name_map) case ^IR_For: s.start = opt_copy_prop_expr(s.start, copy_map, name_map) s.stop = opt_copy_prop_expr(s.stop, copy_map, name_map) s.step = opt_copy_prop_expr(s.step, copy_map, name_map) for_map := make(map[IR_Var_Id]IR_Var_Id) defer delete(for_map) opt_copy_prop_stmts(s.body[:], &for_map, name_map) case ^IR_While: s.condition = opt_copy_prop_expr(s.condition, copy_map, name_map) while_map := make(map[IR_Var_Id]IR_Var_Id) defer delete(while_map) opt_copy_prop_stmts(s.body[:], &while_map, name_map) case ^IR_Expr_Stmt: s.expr = opt_copy_prop_expr(s.expr, copy_map, name_map) case ^IR_Barrier: // no-op case ^IR_Discard: // no-op case ^IR_Break: // no-op case ^IR_Continue: // no-op } } } @(private = "file") opt_copy_prop_expr :: proc(expr: ^IR_Expr, copy_map: ^map[IR_Var_Id]IR_Var_Id, name_map: ^map[IR_Var_Id]string) -> ^IR_Expr { if expr == nil do return nil switch d in expr.derived { case ^IR_Var_Ref: // Follow the copy chain to the original resolved := d.id for { next, ok := copy_map[resolved] if !ok do break resolved = next } if resolved != d.id { d.id = resolved if name, ok := name_map[resolved]; ok { d.name = name } } case ^IR_Binary: d.left = opt_copy_prop_expr(d.left, copy_map, name_map) d.right = opt_copy_prop_expr(d.right, copy_map, name_map) case ^IR_Unary: d.operand = opt_copy_prop_expr(d.operand, copy_map, name_map) case ^IR_Call: for &arg in d.args { arg = opt_copy_prop_expr(arg, copy_map, name_map) } case ^IR_Field_Access: d.object = opt_copy_prop_expr(d.object, copy_map, name_map) case ^IR_Swizzle: d.object = opt_copy_prop_expr(d.object, copy_map, name_map) case ^IR_Composite_Extract: d.object = opt_copy_prop_expr(d.object, copy_map, name_map) case ^IR_Vector_Shuffle: d.object = opt_copy_prop_expr(d.object, copy_map, name_map) case ^IR_Index: d.object = opt_copy_prop_expr(d.object, copy_map, name_map) d.index = opt_copy_prop_expr(d.index, copy_map, name_map) case ^IR_Construct: for &arg in d.args { arg = opt_copy_prop_expr(arg, copy_map, name_map) } case ^IR_Type_Cast: d.value = opt_copy_prop_expr(d.value, copy_map, name_map) case ^IR_Select: d.condition = opt_copy_prop_expr(d.condition, copy_map, name_map) d.true_val = opt_copy_prop_expr(d.true_val, copy_map, name_map) d.false_val = opt_copy_prop_expr(d.false_val, copy_map, name_map) case ^IR_Literal, ^IR_Load_Binding, ^IR_Input_Field, ^IR_Builtin_Var, ^IR_Shared_Ref: // leaf nodes } return expr } // -- Common Subexpression Elimination -- // Within flat blocks, hash IR_Let value expressions and replace duplicates // with references to the first computation. DCE cleans up unused originals. opt_cse :: proc(module: ^IR_Module) { for &fn in module.functions { opt_cse_stmts(fn.body[:]) } } @(private = "file") opt_cse_stmts :: proc(stmts: []IR_Stmt) { // Map from expression string → var that holds that value seen_id := make(map[string]IR_Var_Id) seen_name := make(map[string]string) defer delete(seen_id) defer delete(seen_name) for &stmt in stmts { switch s in stmt { case ^IR_Let: if !expr_has_side_effects(s.value) { key := cse_expr_key(s.value) if len(key) > 0 { if existing_id, ok := seen_id[key]; ok { // Replace with reference to existing variable ref := new(IR_Var_Ref) ref.id = existing_id ref.name = seen_name[key] s.value = new(IR_Expr) s.value.kind = .Var_Ref s.value.type = s.type s.value.derived = ref } else { seen_id[key] = s.id seen_name[key] = s.name } } } case ^IR_If: // Recurse into sub-blocks with fresh scope opt_cse_stmts(s.then_body[:]) for &ei in s.elseif_clauses { opt_cse_stmts(ei.body[:]) } opt_cse_stmts(s.else_body[:]) case ^IR_For: opt_cse_stmts(s.body[:]) case ^IR_While: opt_cse_stmts(s.body[:]) case ^IR_Assign, ^IR_Return, ^IR_Store_Output, ^IR_Expr_Stmt, ^IR_Barrier, ^IR_Discard, ^IR_Break, ^IR_Continue: // no CSE opportunities in these } } } // Build a structural key string for an expression for CSE comparison @(private = "file") cse_expr_key :: proc(expr: ^IR_Expr) -> string { if expr == nil do return "" b := make([dynamic]u8, context.temp_allocator) cse_expr_key_build(&b, expr) return string(b[:]) } @(private = "file") cse_expr_key_build :: proc(b: ^[dynamic]u8, expr: ^IR_Expr) { if expr == nil do return switch d in expr.derived { case ^IR_Literal: append(b, 'L') switch v in d.value { case i64: s := fmt.tprintf("%d", v) append(b, ..transmute([]u8)s) case f64: s := fmt.tprintf("%f", v) append(b, ..transmute([]u8)s) case bool: append(b, v ? '1' : '0') } case ^IR_Var_Ref: append(b, 'V') s := fmt.tprintf("%d", int(d.id)) append(b, ..transmute([]u8)s) case ^IR_Binary: append(b, '(') cse_expr_key_build(b, d.left) s := ir_op_to_string(d.op) append(b, ..transmute([]u8)s) cse_expr_key_build(b, d.right) append(b, ')') case ^IR_Unary: append(b, 'U') s := ir_op_to_string(d.op) append(b, ..transmute([]u8)s) cse_expr_key_build(b, d.operand) case ^IR_Call: append(b, 'C') append(b, ..transmute([]u8)d.name) append(b, '(') for arg, i in d.args { if i > 0 do append(b, ',') cse_expr_key_build(b, arg) } append(b, ')') case ^IR_Field_Access: cse_expr_key_build(b, d.object) append(b, '.') append(b, ..transmute([]u8)d.field_name) case ^IR_Swizzle: cse_expr_key_build(b, d.object) append(b, '.') append(b, ..transmute([]u8)d.components) case ^IR_Composite_Extract: cse_expr_key_build(b, d.object) append(b, '@') s := fmt.tprintf("%d", d.index) append(b, ..transmute([]u8)s) case ^IR_Vector_Shuffle: cse_expr_key_build(b, d.object) append(b, 'S') for c, i in d.components { if i > 0 do append(b, ',') s := fmt.tprintf("%d", c) append(b, ..transmute([]u8)s) } case ^IR_Index: cse_expr_key_build(b, d.object) append(b, '[') cse_expr_key_build(b, d.index) append(b, ']') case ^IR_Construct: append(b, 'K') append(b, ..transmute([]u8)d.type_name) append(b, '(') for arg, i in d.args { if i > 0 do append(b, ',') cse_expr_key_build(b, arg) } append(b, ')') case ^IR_Type_Cast: append(b, 'T') s := type_to_string(expr.type) append(b, ..transmute([]u8)s) append(b, '(') cse_expr_key_build(b, d.value) append(b, ')') case ^IR_Load_Binding: append(b, 'B') append(b, ..transmute([]u8)d.name) case ^IR_Input_Field: append(b, 'I') append(b, ..transmute([]u8)d.param_name) append(b, '.') append(b, ..transmute([]u8)d.field_name) case ^IR_Builtin_Var: append(b, 'G') append(b, ..transmute([]u8)d.name) case ^IR_Shared_Ref: append(b, 'H') append(b, ..transmute([]u8)d.name) case ^IR_Select: append(b, 'Q') cse_expr_key_build(b, d.condition) append(b, '?') cse_expr_key_build(b, d.true_val) append(b, ':') cse_expr_key_build(b, d.false_val) } } // -- Dead Code Elimination -- opt_dead_code_elim :: proc(module: ^IR_Module) { for &fn in module.functions { // Count variable uses uses := make(map[IR_Var_Id]int) count_uses_stmts(fn.body[:], &uses) // Remove unused lets (that have no side effects) fn.body = dce_filter_stmts(fn.body, &uses) } } @(private = "file") count_uses_stmts :: proc(stmts: []IR_Stmt, uses: ^map[IR_Var_Id]int) { for stmt in stmts { switch s in stmt { case ^IR_Let: count_uses_expr(s.value, uses) case ^IR_Assign: count_uses_expr(s.target, uses) count_uses_expr(s.value, uses) case ^IR_Return: count_uses_expr(s.value, uses) case ^IR_Store_Output: count_uses_expr(s.value, uses) case ^IR_If: count_uses_expr(s.condition, uses) count_uses_stmts(s.then_body[:], uses) for ei in s.elseif_clauses { count_uses_expr(ei.condition, uses) count_uses_stmts(ei.body[:], uses) } count_uses_stmts(s.else_body[:], uses) case ^IR_For: count_uses_expr(s.start, uses) count_uses_expr(s.stop, uses) count_uses_expr(s.step, uses) count_uses_stmts(s.body[:], uses) case ^IR_While: count_uses_expr(s.condition, uses) count_uses_stmts(s.body[:], uses) case ^IR_Expr_Stmt: count_uses_expr(s.expr, uses) case ^IR_Barrier: // no variable references case ^IR_Discard: // no variable references case ^IR_Break: // no variable references case ^IR_Continue: // no variable references } } } @(private = "file") count_uses_expr :: proc(expr: ^IR_Expr, uses: ^map[IR_Var_Id]int) { if expr == nil do return switch d in expr.derived { case ^IR_Var_Ref: uses[d.id] = (d.id in uses^) ? uses[d.id] + 1 : 1 case ^IR_Binary: count_uses_expr(d.left, uses) count_uses_expr(d.right, uses) case ^IR_Unary: count_uses_expr(d.operand, uses) case ^IR_Call: for arg in d.args do count_uses_expr(arg, uses) case ^IR_Field_Access: count_uses_expr(d.object, uses) case ^IR_Swizzle: count_uses_expr(d.object, uses) case ^IR_Composite_Extract: count_uses_expr(d.object, uses) case ^IR_Vector_Shuffle: count_uses_expr(d.object, uses) case ^IR_Index: count_uses_expr(d.object, uses) count_uses_expr(d.index, uses) case ^IR_Construct: for arg in d.args do count_uses_expr(arg, uses) case ^IR_Type_Cast: count_uses_expr(d.value, uses) case ^IR_Select: count_uses_expr(d.condition, uses) count_uses_expr(d.true_val, uses) count_uses_expr(d.false_val, uses) case ^IR_Literal, ^IR_Load_Binding, ^IR_Input_Field, ^IR_Builtin_Var, ^IR_Shared_Ref: // no variable references } } @(private = "file") dce_filter_stmts :: proc(stmts: [dynamic]IR_Stmt, uses: ^map[IR_Var_Id]int) -> [dynamic]IR_Stmt { result := make([dynamic]IR_Stmt) for stmt in stmts { #partial switch s in stmt { case ^IR_Let: // Remove if unused and no side effects in the value if !(s.id in uses^) && !expr_has_side_effects(s.value) { continue } append(&result, stmt) continue case ^IR_If: s.then_body = dce_filter_stmts(s.then_body, uses) for &ei in s.elseif_clauses { ei.body = dce_filter_stmts(ei.body, uses) } s.else_body = dce_filter_stmts(s.else_body, uses) case ^IR_For: s.body = dce_filter_stmts(s.body, uses) case ^IR_While: s.body = dce_filter_stmts(s.body, uses) } append(&result, stmt) } return result } @(private = "file") expr_has_side_effects :: proc(expr: ^IR_Expr) -> bool { if expr == nil do return false #partial switch d in expr.derived { case ^IR_Call: return true // calls may have side effects case ^IR_Binary: return expr_has_side_effects(d.left) || expr_has_side_effects(d.right) case ^IR_Unary: return expr_has_side_effects(d.operand) case ^IR_Construct: for arg in d.args { if expr_has_side_effects(arg) do return true } return false case: return false } }