已开启
[Feature]: 重构 CJBarrierOpt 的 BarrierNeed 算法以提升编译性能 #185
睡觉对我很重要创建于  27 天前
睡觉对我很重要
睡觉对我很重要仓颉Developer
27 天前 创建

新需求提供了什么功能? | What functionality does the new feature provide?

重构 CJBarrierOptBarrierNeed 的 safepoint 路径分析算法,并移除该 pass 对 LoopInfo 的依赖。

  • 旧实现:依赖 LoopInfoWrapperPass;在 scanSuccPath / scanPreForSafepoint 中反复调用 isPotentiallyReachable(每次 O(V+E) DFS),整体近似 O(E·(V+E))。
  • 新实现:先一次反向 BFS 求解 CanReachEnd(从 End 往回标记能到达的 BB),再一次前向 worklist 单调 OR-join 不动点(SafepointFlow),总共 O(V+E)。
  • 移除 LoopInfo 依赖:legacy pass 的 getAnalysisUsage 不再 addRequired<LoopInfoWrapperPass>;new-PM run 不再消费 ModuleAnalysisManager

涉及文件:llvm/lib/Transforms/Scalar/CJBarrierOpt.cpp(+72 / -215,净减 143 行)。

该需求带来的价值、应用场景? | Value and application scenarios?

CJ 编译器后端 cj-barrier-opt pass 对每个 gc "cangjie" 函数运行;重构后编译性能显著提升。

实测(合成压测模块:100 个 gc "cangjie" 函数 × 3 层嵌套循环 × 每层 4 菱形,共 7700 基本块;触发 barriersCheckFail → BarrierNeed 的真实路径,--safepoint-mode none 使其全遍历不 abort)。-time-passesCJBarrierOpt pass user time,各 5 轮:

版本 5 轮耗时(秒,升序)
重构前 0.0014 / 0.0052 / 0.0053 / 0.0056 / 0.0057
重构后 0.0009 / 0.0009 / 0.0010 / 0.0011 / 0.0011

等价性已验证:CJBarrierOpt 测试套件(llvm/test/Transforms/CJBarrierOpt/)重构前后均 22/23 通过,输出 byte-identical(同一用例 barrierarray5.ll 失败与本次重构无关,源于更早的 intrinsic 重载名变更)。


import argparse
import subprocess
import sys
import tempfile
from pathlib import Path


class FunctionBuilder:
    def __init__(self, function_id, args):
        self.function_id = function_id
        self.args = args
        self.blocks = []
        self.labels = set()
        self.edges = []
        self.phis = []
        self.safepoints = 0

    def add_block(self, label, instructions, targets=()):
        if label in self.labels:
            raise ValueError("duplicate basic block: {0}".format(label))
        self.labels.add(label)
        self.blocks.append((label, instructions))
        for target in targets:
            self.edges.append((label, target))

    def maybe_safepoint(self, arm):
        mode = self.args.safepoint_mode
        if mode == "both" or mode == arm:
            self.safepoints += 1
            return ["  call cangjiegccc void @MCC_SafepointStub()"]
        return []

    def emit_diamonds(self, level, predecessor):
        current = predecessor
        for diamond in range(self.args.diamonds):
            prefix = "l{0}.d{1}".format(level, diamond)
            branch = prefix + ".branch"
            then = prefix + ".then"
            otherwise = prefix + ".else"
            merge = prefix + ".merge"
            self.add_block(current, ["  br label %{0}".format(branch)], [branch])
            self.add_block(
                branch,
                ["  br i1 %diamond.cond, label %{0}, label %{1}".format(
                    then, otherwise
                )],
                [then, otherwise],
            )
            then_insts = self.maybe_safepoint("then")
            then_insts.append("  br label %{0}".format(merge))
            self.add_block(then, then_insts, [merge])
            else_insts = self.maybe_safepoint("else")
            else_insts.append("  br label %{0}".format(merge))
            self.add_block(otherwise, else_insts, [merge])
            current = merge
        return current

    def emit_loop(self, level, entry, normal_exit):
        header = "l{0}.header".format(level)
        dispatch = "l{0}.dispatch".format(level)
        body = "l{0}.body".format(level)
        latch = "l{0}.latch".format(level)
        after = "l{0}.after".format(level)

        self.add_block(entry, ["  br label %{0}".format(header)], [header])
        self.add_block(
            header,
            [
                "  %iv{0} = phi i32 [ 0, %{1} ], [ %iv{0}.next, %{2} ]".format(
                    level, entry, latch
                ),
                "  %keep{0} = icmp ult i32 %iv{0}, %trip.count".format(level),
                "  br i1 %keep{0}, label %{1}, label %{2}".format(
                    level, dispatch, after
                ),
            ],
            [dispatch, after],
        )
        self.phis.append((header, [entry, latch]))

        exceptional_exits = []
        for exit_index in range(1, self.args.exits):
            exceptional_exits.append("l{0}.exit{1}".format(level, exit_index))

        if exceptional_exits:
            cases = []
            for exit_index, target in enumerate(exceptional_exits, 1):
                case_value = level * 100000 + exit_index
                cases.append("    i32 {0}, label %{1}".format(case_value, target))
            switch = [
                "  switch i32 %exit.selector, label %{0} [".format(body)
            ] + cases + ["  ]"]
            self.add_block(dispatch, switch, [body] + exceptional_exits)
        else:
            self.add_block(dispatch, ["  br label %{0}".format(body)], [body])

        for target in exceptional_exits:
            self.add_block(
                target,
                ["  br label %final.merge"],
                ["final.merge"],
            )

        continuation = self.emit_diamonds(level, body)
        if level + 1 < self.args.nested_loops:
            self.emit_loop(level + 1, continuation, latch)
        else:
            self.add_block(
                continuation, ["  br label %{0}".format(latch)], [latch]
            )

        latch_insts = []
        if self.args.safepoint_mode == "loop-latch":
            self.safepoints += 1
            latch_insts.append("  call cangjiegccc void @MCC_SafepointStub()")
        latch_insts.extend(
            [
                "  %iv{0}.next = add nuw i32 %iv{0}, 1".format(level),
                "  br label %{0}".format(header),
            ]
        )
        self.add_block(latch, latch_insts, [header])
        self.add_block(after, ["  br label %{0}".format(normal_exit)], [normal_exit])

    def validate(self):
        missing = sorted(
            {target for _, target in self.edges if target not in self.labels}
        )
        if missing:
            raise ValueError("branches target missing blocks: {0}".format(missing))

        predecessors = {label: set() for label in self.labels}
        for source, target in self.edges:
            predecessors[target].add(source)
        for block, incoming in self.phis:
            if set(incoming) != predecessors[block]:
                raise ValueError(
                    "PHI predecessors for {0}: expected {1}, got {2}".format(
                        block, sorted(predecessors[block]), sorted(incoming)
                    )
                )

        expected_safepoints = 0
        if self.args.safepoint_mode in ("then", "else"):
            expected_safepoints = self.args.nested_loops * self.args.diamonds
        elif self.args.safepoint_mode == "both":
            expected_safepoints = 2 * self.args.nested_loops * self.args.diamonds
        elif self.args.safepoint_mode == "loop-latch":
            expected_safepoints = self.args.nested_loops
        if self.safepoints != expected_safepoints:
            raise ValueError(
                "expected {0} safepoints, generated {1}".format(
                    expected_safepoints, self.safepoints
                )
            )

    def generate(self):
        allocation = "alloc"
        self.add_block(
            "entry",
            [
                "  %obj = call i8 addrspace(1)* @CJ_MCC_NewObject(i8* null, i32 16)",
                "  %dst.record = bitcast i8 addrspace(1)* %obj to %record addrspace(1)*",
                "  br label %{0}".format(allocation),
            ],
            [allocation],
        )

        self.emit_loop(0, allocation, "final.dispatch")

        if self.args.merged_preds > 0:
            cases = []
            merge_blocks = []
            for index in range(self.args.merged_preds):
                target = "final.pred{0}".format(index)
                merge_blocks.append(target)
                case_value = 200000 + index
                cases.append("    i32 {0}, label %{1}".format(case_value, target))
            switch = [
                "  switch i32 %exit.selector, label %final.merge ["
            ] + cases + ["  ]"]
            self.add_block(
                "final.dispatch", switch, ["final.merge"] + merge_blocks
            )
            for target in merge_blocks:
                self.add_block(
                    target, ["  br label %final.merge"], ["final.merge"]
                )
        else:
            self.add_block(
                "final.dispatch", ["  br label %final.merge"], ["final.merge"]
            )

        self.add_block(
            "final.merge",
            [
                "  %dst = bitcast %record addrspace(1)* %dst.record to i8 addrspace(1)*",
                "  %src.bytes = bitcast %record addrspace(1)* %src to i8 addrspace(1)*",
                "  call void @llvm.memcpy.p1i8.p1i8.i64(i8 addrspace(1)* align 8 %dst, i8 addrspace(1)* align 8 %src.bytes, i64 16, i1 false)",
                "  ret void",
            ],
        )
        self.validate()

        lines = [
            "define void @barrier_need_stress_{0}(%record addrspace(1)* %src, i1 %diamond.cond, i32 %exit.selector, i32 %trip.count) gc \"cangjie\" {{".format(
                self.function_id
            )
        ]
        for label, instructions in self.blocks:
            lines.append("{0}:".format(label))
            lines.extend(instructions)
            lines.append("")
        lines.append("}")
        return "\n".join(lines)


def generate_module(args):
    functions = []
    total_safepoints = 0
    for function_id in range(args.functions):
        builder = FunctionBuilder(function_id, args)
        functions.append(builder.generate())
        total_safepoints += builder.safepoints

    preamble = """; Auto-generated by utils/cjbarrieropt/generate_barrier_need_stress.py
; nested-loops={nested_loops} exits={exits} diamonds={diamonds} merged-preds={merged_preds}
; functions={functions} safepoint-mode={safepoint_mode} safepoints={safepoints}
; This module stresses BarrierNeed CFG reachability and fixed-point propagation.

target datalayout = \"e-m:o-i64:64-i128:128-n32:64-S128-p1:64:64\"

%record = type {{ i32, i8 addrspace(1)* }}

declare i8 addrspace(1)* @CJ_MCC_NewObject(i8*, i32)
declare cangjiegccc void @MCC_SafepointStub()
declare void @llvm.memcpy.p1i8.p1i8.i64(i8 addrspace(1)* noalias nocapture writeonly, i8 addrspace(1)* noalias nocapture readonly, i64, i1 immarg)
""".format(
        nested_loops=args.nested_loops,
        exits=args.exits,
        diamonds=args.diamonds,
        merged_preds=args.merged_preds,
        functions=args.functions,
        safepoint_mode=args.safepoint_mode,
        safepoints=total_safepoints,
    )
    return preamble + "\n\n".join(functions) + "\n"


def verify_ir(ir, tool):
    with tempfile.NamedTemporaryFile(mode="w", suffix=".ll", delete=False) as output:
        output.write(ir)
        path = output.name
    try:
        command = [tool, "-verify", "-disable-output", path]
        subprocess.run(command, check=True)
    finally:
        Path(path).unlink()


def parse_args():
    parser = argparse.ArgumentParser(
        description="Generate nested-loop and diamond stress IR for CJBarrierOpt"
    )
    parser.add_argument("--nested-loops", "--depth", type=int, default=3)
    parser.add_argument("--exits", type=int, default=3)
    parser.add_argument("--diamonds", type=int, default=2)
    parser.add_argument("--merged-preds", type=int, default=4)
    parser.add_argument("--functions", type=int, default=1)
    parser.add_argument(
        "--safepoint-mode",
        choices=("none", "then", "else", "both", "loop-latch"),
        default="none",
    )
    parser.add_argument("-o", "--output", type=Path)
    parser.add_argument(
        "--verify-with-opt",
        metavar="PATH",
        help="run PATH -verify -disable-output on the generated IR",
    )
    args = parser.parse_args()
    for name in ("nested_loops", "exits", "functions"):
        if getattr(args, name) < 1:
            parser.error("--{0} must be at least 1".format(name.replace("_", "-")))
    for name in ("diamonds", "merged_preds"):
        if getattr(args, name) < 0:
            parser.error("--{0} must not be negative".format(name.replace("_", "-")))
    if args.safepoint_mode in ("then", "else", "both") and args.diamonds == 0:
        parser.error("diamond safepoint modes require --diamonds greater than 0")
    return args


def main():
    args = parse_args()
    ir = generate_module(args)
    if args.verify_with_opt:
        verify_ir(ir, args.verify_with_opt)
    if args.output:
        args.output.write_text(ir)
    else:
        sys.stdout.write(ir)


if __name__ == "__main__":
    main()

新需求涉及的分支版本 | Branch versions

likedislike
睡觉对我很重要睡觉对我很重要仓颉Developer
27 天前 添加了label:enhancement
睡觉对我很重要睡觉对我很重要仓颉Developer
27 天前 关联了pull request:perf(cj-barrier-opt): 重构 BarrierNeed 算法并移除 LoopInfo 依赖
睡觉对我很重要睡觉对我很重要仓颉Developer
27 天前 修改了issue 的描述
睡觉对我很重要睡觉对我很重要仓颉Developer
27 天前 修改了issue 的描述