diff --git a/benchmarks/bench_const_merge.py b/benchmarks/bench_const_merge.py index 21d837b..9a5e043 100644 --- a/benchmarks/bench_const_merge.py +++ b/benchmarks/bench_const_merge.py @@ -12,15 +12,22 @@ from __future__ import annotations import argparse +import os import statistics +import sys import time -from typing import Optional -from scratchv.backend.const_merge import merge_constants +BENCH_DIR = os.path.dirname(__file__) +PROJ_DIR = os.path.dirname(BENCH_DIR) +sys.path.insert(0, PROJ_DIR) + +from scratchv.backend._asm_parser import parse_asm +from scratchv.backend.const_merge import merge_constants_detailed def _gen_synthetic_asm(num_instrs: int, seed: int = 42, - lui_ratio: float = 0.3) -> str: + lui_ratio: float = 0.3, + redundant_lui_ratio: float = 0.1) -> str: """Generate synthetic assembly with lui+addi patterns. Parameters @@ -31,16 +38,34 @@ def _gen_synthetic_asm(num_instrs: int, seed: int = 42, Random seed for reproducibility. lui_ratio: Fraction of instructions that form lui+addi pairs. + redundant_lui_ratio: + Fraction of generated groups that contain a redundant LUI pattern. """ + if num_instrs < 0: + raise ValueError("num_instrs must be non-negative") + if not 0.0 <= lui_ratio <= 1.0: + raise ValueError("lui_ratio must be between 0 and 1") + if not 0.0 <= redundant_lui_ratio <= 1.0: + raise ValueError("redundant_lui_ratio must be between 0 and 1") + if lui_ratio + redundant_lui_ratio > 1.0: + raise ValueError("lui_ratio + redundant_lui_ratio must not exceed 1") import random random.seed(seed) lines = [".text", "synthetic_func:"] i = 0 while i < num_instrs: - use_lui = random.random() < lui_ratio + choice = random.random() - if use_lui and i + 1 < num_instrs: + if choice < redundant_lui_ratio and i + 2 < num_instrs: + regs = ["t0", "t1", "t2", "s0", "s1", "a0", "a1"] + r = random.choice(regs) + imm_hi = random.choice([0x10000, 0x20000, 0x12345]) + lines.append(f" lui {r}, {hex(imm_hi)}") + lines.append(f" add a4, a5, a6") + lines.append(f" lui {r}, {hex(imm_hi)}") + i += 3 + elif choice < redundant_lui_ratio + lui_ratio and i + 1 < num_instrs: regs = ["t0", "t1", "t2", "s0", "s1", "a0", "a1", "a2", "a3"] r = random.choice(regs) imm_hi = random.choice([0x10000, 0x20000, 0x12345, 0xABCDE, 0xFFFFF]) @@ -78,24 +103,39 @@ def _gen_synthetic_asm(num_instrs: int, seed: int = 42, def bench_merge(asm_text: str, repeats: int = 50) -> dict: """Benchmark the constant merge optimizer.""" + if repeats < 1: + raise ValueError("repeats must be at least 1") times = [] results = [] for _ in range(repeats): t0 = time.perf_counter() - result, changes = merge_constants(asm_text) + result, stats = merge_constants_detailed(asm_text) t1 = time.perf_counter() times.append(t1 - t0) - results.append((result, changes)) - - changes_list = [r[1] for r in results] - input_lines = asm_text.count("\n") - output_lines = results[0][0].count("\n") if results else 0 + results.append((result, stats)) + + changes_list = [r[1].total_changes for r in results] + first_stats = results[0][1] + parsed_input = parse_asm(asm_text) + parsed_output = parse_asm(results[0][0]) if results else [] + input_instructions = sum( + line.opcode is not None and not line.is_directive + for line in parsed_input + ) + output_instructions = sum( + line.opcode is not None and not line.is_directive + for line in parsed_output + ) return { - "input_lines": input_lines, - "output_lines": output_lines, - "line_reduction": input_lines - output_lines, + "benchmark_type": "synthetic", + "input_instructions": input_instructions, + "output_instructions": output_instructions, + "instruction_reduction": input_instructions - output_instructions, + "candidate_pairs": first_stats.candidate_pairs, + "merged_pairs": first_stats.merged_pairs, + "redundant_lui_removed": first_stats.redundant_lui_removed, "changes_mean": statistics.mean(changes_list), "changes_stdev": statistics.stdev(changes_list) if len(changes_list) > 1 else 0, "repeats": repeats, @@ -111,35 +151,46 @@ def main(): parser = argparse.ArgumentParser(description="Constant Merge Benchmark") parser.add_argument("--repeats", type=int, default=50, help="Number of repeat measurements") + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--pair-density", type=float, default=0.3) + parser.add_argument("--redundant-lui-density", type=float, default=0.1) args = parser.parse_args() sizes = [100, 500, 1000, 2000, 5000] print("=" * 80) print("RISC-V Constant Load Merge Optimizer Benchmark") + print("benchmark_type=synthetic") print("=" * 80) print(f"\n{'Size':>8} {'Mean(ms)':>10} {'Stdev(ms)':>10} " - f"{'Changes':>8} {'InpLines':>10} {'OutLines':>10} {'Reduc':>8}") + f"{'Pairs':>8} {'RedLUI':>8} {'InpInst':>10} {'OutInst':>10}") print("-" * 80) for size in sizes: - asm = _gen_synthetic_asm(size, lui_ratio=0.3) + asm = _gen_synthetic_asm( + size, seed=args.seed, lui_ratio=args.pair_density, + redundant_lui_ratio=args.redundant_lui_density, + ) stats = bench_merge(asm, repeats=args.repeats) print(f"{size:>8} {stats['mean_s'] * 1000:>10.3f} " f"{stats['stdev_s'] * 1000:>10.3f} " - f"{stats['changes_mean']:>8.1f} " - f"{stats['input_lines']:>10} {stats['output_lines']:>10} " - f"{stats['line_reduction']:>8}") + f"{stats['merged_pairs']:>8} " + f"{stats['redundant_lui_removed']:>8} " + f"{stats['input_instructions']:>10} " + f"{stats['output_instructions']:>10}") # Test different lui densities print(f"\nLUI Density Impact (2000 instructions):") print("-" * 60) for ratio in [0.0, 0.1, 0.3, 0.5]: - asm = _gen_synthetic_asm(2000, lui_ratio=ratio) + asm = _gen_synthetic_asm( + 2000, seed=args.seed, lui_ratio=ratio, + redundant_lui_ratio=args.redundant_lui_density, + ) stats = bench_merge(asm, repeats=args.repeats) print(f" ratio={ratio:.1f} {stats['mean_s'] * 1000:.3f} ms " f"changes: {stats['changes_mean']:.1f} " - f"reduction: {stats['line_reduction']}") + f"reduction: {stats['instruction_reduction']}") if __name__ == "__main__": diff --git a/benchmarks/run_benchmark.py b/benchmarks/run_benchmark.py index f473214..75a80bf 100644 --- a/benchmarks/run_benchmark.py +++ b/benchmarks/run_benchmark.py @@ -36,6 +36,7 @@ from scratchv.frontend.dsl_parser import DSLParser from scratchv.ir.builder import IRBuilder from scratchv.ir.types import Program +from scratchv.backend._asm_parser import ParsedAsmLine, parse_asm # --------------------------------------------------------------------------- @@ -55,6 +56,18 @@ class BenchResult: ir_opt_inst_count: int = 0 codegen_time_s: float = 0.0 asm_line_count: int = 0 + lui_count_before: int = 0 + candidate_pairs: int = 0 + merged_pairs: int = 0 + redundant_lui_removed: int = 0 + asm_instructions_before: int = 0 + asm_instructions_after: int = 0 + machine_instructions_before: Optional[int] = None + machine_instructions_after: Optional[int] = None + code_size_before: Optional[int] = None + code_size_after: Optional[int] = None + const_merge_time_ms: float = 0.0 + output_equal: Optional[bool] = None total_time_s: float = 0.0 verified: bool = False error: Optional[str] = None @@ -70,6 +83,18 @@ def _count_ir(program: Program) -> tuple[int, int]: return inst, bb +def _count_asm_instructions(asm_text: str) -> int: + """Count assembly instructions, excluding labels, blanks and directives.""" + return _count_parsed_asm_instructions(parse_asm(asm_text)) + + +def _count_parsed_asm_instructions(lines: list[ParsedAsmLine]) -> int: + return sum( + 1 for line in lines + if line.opcode is not None and not line.is_directive + ) + + def _parse_onnx(path: str) -> Program: parser = ONNXParser() return parser.parse(path) @@ -173,6 +198,23 @@ def run_benchmark(model_name: str, model_path: str, *, asm_str, result.codegen_time_s = _codegen_llvm(program) else: asm_str, result.codegen_time_s = _codegen_riscv(program) + from scratchv.backend.const_merge import merge_constants_detailed + parsed_before = parse_asm(asm_str) + result.lui_count_before = sum( + line.opcode == "lui" for line in parsed_before + ) + result.asm_instructions_before = _count_parsed_asm_instructions( + parsed_before, + ) + t0 = time.perf_counter() + asm_after, merge_stats = merge_constants_detailed(asm_str) + result.const_merge_time_ms = (time.perf_counter() - t0) * 1000 + result.candidate_pairs = merge_stats.candidate_pairs + result.merged_pairs = merge_stats.merged_pairs + result.redundant_lui_removed = merge_stats.redundant_lui_removed + result.asm_instructions_after = _count_parsed_asm_instructions( + parse_asm(asm_after), + ) result.asm_line_count = len(asm_str.splitlines()) # 4. Verify diff --git a/benchmarks/test_benchmark.py b/benchmarks/test_benchmark.py index 74e08c2..43da522 100644 --- a/benchmarks/test_benchmark.py +++ b/benchmarks/test_benchmark.py @@ -23,6 +23,7 @@ from benchmarks.generate_models import ensure_all_models from benchmarks.run_benchmark import run_benchmark +from benchmarks.bench_const_merge import _gen_synthetic_asm, bench_merge # --------------------------------------------------------------------------- @@ -38,6 +39,23 @@ def benchmark_models() -> dict[str, str]: BACKEND_PARAMS = ["riscv"] +@pytest.mark.parametrize( + "pair_density,redundant_density", + [(-0.1, 0.1), (0.1, -0.1), (1.1, 0.0), (0.6, 0.5)], +) +def test_synthetic_density_validation(pair_density, redundant_density): + with pytest.raises(ValueError): + _gen_synthetic_asm( + 10, lui_ratio=pair_density, + redundant_lui_ratio=redundant_density, + ) + + +def test_synthetic_repeats_validation(): + with pytest.raises(ValueError, match="repeats"): + bench_merge(" nop\n", repeats=0) + + def _model_id(name: str) -> str: return name @@ -151,6 +169,15 @@ def test_perf_pipeline(model_name: str, benchmark_models: dict[str, str]): assert result.error is None, f"Benchmark failed: {result.error}" assert result.ir_inst_count > 0 + assert isinstance(result.asm_instructions_before, int) + assert isinstance(result.asm_instructions_after, int) + reduction = ( + result.asm_instructions_before - result.asm_instructions_after + ) + tracked_changes = result.merged_pairs + result.redundant_lui_removed + assert ( + reduction == tracked_changes + ), f"instruction reduction {reduction} != tracked changes {tracked_changes}" print(f"\n {model_name}:") print(f" parse: {result.parse_time_s:.4f}s") diff --git "a/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\345\274\200\345\217\221\346\226\207\346\241\243\345\210\235\347\250\277.md" "b/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\345\274\200\345\217\221\346\226\207\346\241\243\345\210\235\347\250\277.md" new file mode 100644 index 0000000..95e3e14 --- /dev/null +++ "b/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\345\274\200\345\217\221\346\226\207\346\241\243\345\210\235\347\250\277.md" @@ -0,0 +1,579 @@ +# ScratchV 课题 14:常量加载合并优化开发文档 + +> **文档版本**:v0.3(核心实现与验证状态回填) +> **创建日期**:2026-07-28 +> **更新日期**:2026-08-08 +> **作者**:[yuki] +> **关联 Issue**:[#待补充] +> **涉及模块**:`scratchv/backend/`、`scratchv/compiler.py`、`scratchv/main.py`、`tests/`、`docs/` + +--- + +## 0. 开发前说明 + +ScratchV 主分支已经存在课题 14 的参考模块、测试和命令行集成。本开发计划将任务定义为: + +> 对现有常量加载合并优化进行行为梳理、安全性加固、共享解析器重构、边界测试补全和集成验证。 + +### 0.1 2026-08-08 实现状态 + +已完成共享解析器迁移、RV32 常量计算、寄存器别名规范化、基本块边界保护、未知指令保守处理、固定点迭代、分类统计、CLI 统计、CompilerDriver 集成、真实 case A/B 指标和 synthetic benchmark 改造。旧接口 `merge_constants(...)->(str, int)` 保持兼容。 + +验证结果:课题相关测试 62 项通过;排除与本课题无关的 TinyFive stub 环境问题后,回归测试 377 项通过、4 项跳过。抽查 `add`、`mixed_ops`、`deep_relu` 三个真实 case 均为零 `lui`、零候选和零转换,符合当前后端直接生成 `li` 的预期。当前环境没有 GNU RISC-V 工具链或 Spike,机器指令数、代码大小和执行等价字段保留为 `N/A`。 + + +--- + +## 1. 功能概述与目标 + +### 1.1 背景与动机 + +- **现状问题**:RISC-V 加载一般 32 位常量常出现 `lui + addi` 序列;重复加载相同高 20 位还可能产生冗余 `lui`。`li` 是汇编器伪指令,不是 RV32I 的单条真实机器指令。现有参考实现具备基本功能,但仍需要补充控制流安全、寄存器别名、共享解析器和更严格的测试。 +- **应用场景**:处理 ScratchV 后端生成的汇编、外部 RISC-V 汇编文件、课程测试 fixture,以及汇编级窥孔优化和统计流程。 +- **现实限制(已按 commit `109a6a2` 核实)**:`instruction_select.py` 的常量、循环初值等路径直接生成 `MachineOp.LI`。因此现有 ONNX/DSL case 可能不产生 `lui+addi`,主流水线命中率可能为 0。零命中是需要记录的实验结果,不是需要通过改造真实输入掩盖的问题。 + +### 1.2 功能描述 + +- **一句话定义**:在汇编生成后识别安全的常量加载模式,将 `lui+addi` 规范化为 `li`,并删除同一基本块内可证明冗余的 `lui`。 +- **核心价值**:提高汇编可读性,消除真实冗余指令,建立可复用的汇编级优化开发流程。 + +### 1.3 目标与非目标 + +| 类型 | 内容 | +|---|---| +| ✅ 包含范围 | RV32 数值立即数;12 位符号扩展;相邻 `lui+addi`;基本块内冗余 `lui`;寄存器别名;迭代扫描;统计;独立 CLI;主编译器开关;单元与集成测试 | +| ❌ 不包含范围 | RV64 完整常量构造;符号重定位;跨基本块数据流;跨函数复用;IR 级常量折叠;完整汇编器实现 | + +--- + +## 2. 设计与规格说明 + +### 2.1 外部接口 + +#### Python API + +保持: + +```python +from scratchv.backend.const_merge import merge_constants + +optimized_asm, changes = merge_constants(asm_text) +``` + +可选新增: + +```python +optimized_asm, stats = merge_constants_detailed(asm_text) +``` + +#### 独立 CLI + +```bash +python -m scratchv.backend.const_merge input.s -o output.s -v +``` + +#### ScratchV 主 CLI + +```bash +scratchv input.dsl -o output.s --const-merge +``` + +注意:旧版归档文档写的是 `--merge-constants`,但当前主分支真实参数为 `--const-merge`。本开发默认服从当前代码;是否增加旧名称别名由评审决定。 + +### 2.2 内部设计 + +#### 数据结构 + +- 复用 `scratchv/backend/_asm_parser.py` 中的 `ParsedAsmLine`; +- 新增或内部使用 `ConstantMergeStats`; +- 使用 `dict[str, int]` 跟踪基本块内各物理寄存器的最后一次 `lui` 值; +- 使用 ABI 名到 `xN` 的映射做寄存器规范化。 + +#### 核心处理流程 + +```text +读取汇编文本 + -> parse_asm + -> 固定点循环 + -> 删除基本块内冗余 lui + -> 合并安全的 lui+addi + -> lines_to_asm + -> 返回输出与统计 +``` + +#### 状态清空条件 + +遇到以下情况清空或失效化寄存器状态: + +- 标签; +- 条件分支; +- 无条件跳转; +- `call`、`jal`、`jalr`; +- `ret`、`jr`; +- 对已跟踪寄存器产生定义的普通指令使该寄存器状态失效;这不妨碍规则 A 对相邻 `lui+addi` 进行整体匹配; +- 无法可靠分析的未知指令一律清空全部状态。 + +### 2.3 模块间交互 + +- **上游**:接收 `AsmEmitter` 或线性扫描寄存器分配器生成的汇编文本,也可直接接收用户输入 `.s`; +- **下游**:输出给调度器、汇编美化器、指令计数器和最终文件写入; +- **对 IR 无影响**:不修改 AST、IR 或寄存器分配结果; +- **Pass 顺序**:位于汇编窥孔之后、调度器之前,使窥孔优化暴露的模式可以被合并,并避免调度器打散待匹配的相邻序列。 + +--- + +## 3. 开发环境与基线 + +### 3.1 Ubuntu 24.04 环境初始化 + +```bash +git clone https://github.com/ScratchV-Compiler/ScratchV.git +cd ScratchV + +python3 -m venv .venv +source .venv/bin/activate +python -m pip install --upgrade pip +pip install -e . +``` + +可选汇编验证工具: + +```bash +sudo apt update +sudo apt install binutils-riscv64-linux-gnu +``` + +### 3.2 建立功能分支 + +```bash +git switch main +git pull +git switch -c feat/topic14-const-merge +``` + +### 3.3 运行基线测试 + +```bash +pytest tests/test_const_merge.py -v +pytest tests/ -q +``` + +记录: + +- 当前 commit hash; +- Python 版本; +- 测试总数、通过数、失败数; +- 当前 `const_merge` 测试输出; +- 一份未修改前的示例汇编结果。 + +```bash +git rev-parse HEAD +python --version +``` + +--- + +## 4. 涉及文件清单 + +| 文件路径 | 修改类型 | 修改内容概述 | +|---|---|---| +| `scratchv/backend/const_merge.py` | 重点修改 | 复用共享解析器;实现安全匹配、寄存器规范化、块边界处理、固定点迭代和详细统计 | +| `scratchv/backend/_asm_parser.py` | 复用/小改 | 必要时补充公开的寄存器定义/使用判断;避免在 const-merge 中再写一套解析器 | +| `scratchv/compiler.py` | 检查/小改 | 确认 `const_merge` 在 post-codegen 中的顺序和统计输出;必要时接入详细统计 | +| `scratchv/main.py` | 检查/小改 | 确认 `--const-merge` 参数及配置传递;可选增加旧参数别名 | +| `tests/test_const_merge.py` | 重点修改 | 补充边界、安全、别名、迭代、幂等性和 CLI 测试 | +| `tests/test_backend.py` | 修改 | 增加 CompilerConfig/CompilerDriver 集成验证 | +| `tests/fixtures/const_merge/*.s` | 新增 | 存放可读的汇编输入、预期输出和控制流反例 | +| `docs/topics/14-常量加载合并优化.md` 或对应文档源 | 修改 | 更新算法原理、限制、命令和测试结果 | + +实际提交前先用以下命令确认路径: + +```bash +find scratchv/backend -maxdepth 1 -type f | sort +grep -R "const_merge" -n scratchv tests docs | head -100 +``` + +--- + +## 5. 分步实现计划 + +### 步骤 0:确认课题边界 + +**任务**:向维护者确认 RV32/RV64、跨基本块、API 和 CLI 命名。 +**产出**:评审结论写入设计文档“工作假设”章节。 +**验证**:所有待评审问题有明确答案或标记为非目标。 + +### 步骤 1:补 characterization tests + +**任务**:先把当前实现行为固定下来,不立即重构。 +**产出**:现有正向样例、无匹配样例、CLI 可导入测试。 +**验证**:修改前新增测试通过,或明确暴露当前 bug。 + +### 步骤 2:迁移到共享汇编解析器 + +**任务**:用 `parse_asm`、`lines_to_asm` 和 `classify_def_use` 替代重复解析逻辑。 +**产出**:`const_merge.py` 不再维护完整的独立 `AsmInst` 解析器,或仅保留兼容别名。 +**验证**:标签、注释、空行、内存操作数的解析测试通过。 + +### 步骤 3:实现寄存器规范化 + +**任务**:建立 ABI 名到 `xN` 的映射。 +**产出**:`canonical_reg("t0") == "x5"`。 +**验证**:别名混用测试通过。 + +### 步骤 4:实现 `lui+addi` 安全合并 + +**任务**: + +- 检查操作码和操作数数量; +- 检查 `rd == addi.rd == addi.rs1`; +- 只接受纯数值立即数; +- 处理 12 位符号扩展; +- 按 RV32 截断; +- 不跨标签; +- 保留注释。 + +**产出**:`merge_lui_addi_once()`。 +**验证**:正数、负数和边界立即数测试通过。 + +### 步骤 5:实现块内冗余 `lui` 消除 + +**任务**: + +- 维护 `lui_state`; +- 遇到定义寄存器的指令使对应状态失效; +- 遇到标签、分支、跳转、调用和返回清空状态; +- 使用寄存器规范化; +- 未知指令保守处理。 + +**产出**:`remove_redundant_lui_once()`。 +**验证**:安全删除和控制流反例测试通过。 + +### 步骤 6:实现固定点迭代与统计 + +**任务**:迭代运行两个规则直到无变化,记录分类统计。 +**产出**:`ConstantMergeStats`、迭代次数和兼容包装函数。 +**验证**:需要两轮才能完成的样例通过;第二次运行零变化。 + +### 步骤 7:集成主编译器 + +**任务**:检查: + +```text +CompilerConfig.const_merge +main.py --const-merge +CompilerDriver._run_asm_passes +``` + +**产出**:主 CLI 和 Python 配置均可启用优化。 +**验证**:集成测试确认参数被传递并调用优化器。 + +### 步骤 8:工具链语义验证 + +**任务**:把优化前后汇编分别交给 GNU assembler。 +**产出**:对象文件和反汇编对比记录。 +**验证**:两者能汇编,并在测试程序中产生相同返回值/寄存器值。 + +示例: + +```bash +riscv64-linux-gnu-as -march=rv32im before.s -o before.o +riscv64-linux-gnu-as -march=rv32im after.s -o after.o +riscv64-linux-gnu-objdump -d -M no-aliases before.o > before.dump +riscv64-linux-gnu-objdump -d -M no-aliases after.o > after.dump +diff -u before.dump after.dump +``` + +`li` 可能重新展开为 `lui+addi`,因此反汇编差异不应只看源文件行数。 +若工具链不支持 `-march=rv32im` 或命令不可用,应记录完整错误并把机器指令数、代码大小和执行等价指标标为 `N/A`,不得用汇编文本变化代替。 + +### 步骤 9:改造 Benchmark 并建立真实 case A/B + +**当前基线(commit `109a6a2`)**: + +- `benchmarks/bench_runner.py` 测 DSL/IR 执行,不进入 RISC-V 汇编 post-pass,不能证明 const-merge 有效; +- `benchmarks/run_benchmark.py` 能生成 RISC-V 汇编,但当前不调用 `merge_constants`,并且只记录 `asm_line_count`; +- `benchmarks/bench_const_merge.py` 是 synthetic microbenchmark,当前只生成 `lui+addi` 模式、只返回总变化数,并用换行数表示规模与收益。 + +**任务 A:真实 case 命中率。** 在 `run_benchmark.py` 的 RISC-V 路径中,对同一份原始汇编做单变量 A/B: + +```python +asm_before = asm_str +t0 = time.perf_counter() +asm_after, stats = merge_constants_detailed(asm_before) +pass_time_ms = (time.perf_counter() - t0) * 1000 +``` + +不得通过分别执行两次完整编译来取得 before/after,否则会把前端、优化、指令选择和寄存器分配差异混入结果。也不得修改现有 ONNX/DSL case 来制造目标模式。 + +每个真实 case 至少输出: + +| 字段 | 说明 | +|---|---| +| `lui_count_before` | 优化前有效 `lui` 数量 | +| `candidate_pairs` | 同块内可检查的 `lui+addi` 候选数 | +| `merged_pairs` | 规则 A 实际转换数 | +| `redundant_lui_removed` | 规则 B 实际删除数 | +| `asm_instructions_before/after` | 排除标签、伪操作、空行和纯注释后的汇编指令数 | +| `machine_instructions_before/after` | 汇编并反汇编后的真实机器指令数;工具链不可用时记为 `N/A` | +| `code_size_before/after` | `.text` 大小;工具链不可用时记为 `N/A` | +| `pass_time_ms` | 仅 const-merge 的耗时 | +| `output_equal` | 在相同初始寄存器和内存输入下,比较退出码、约定输出寄存器及输出内存;全部一致才为 `true`,无法执行时记为 `N/A` | + +**任务 B:人工 microbenchmark。** 保留 `bench_const_merge.py`,但明确标记为 synthetic,并增加可调规模、目标模式密度和重复 `lui` 密度。输出 `merged_pairs` 与 `redundant_lui_removed`,使用有效指令计数替代 `asm_text.count("\n")`。人工输入的下降量只能说明算法对目标模式有效,不能表述为真实项目端到端收益。 + +**结果判定:** + +- 真实 case 有命中:报告命中 case、两类转换数量、机器指令/代码大小变化及语义验证; +- 真实 case 零命中:明确写出“当前代码生成形式与 const-merge 目标输入不匹配”,同时用 microbenchmark 证明规则正确性和扫描开销; +- 规则 A 仅减少源汇编行数、机器指令不变:按事实分别报告,不宣称运行性能提升; +- 只有规则 B 的删除可直接计入真实机器指令减少,仍需工具链及语义验证支持。 + +**验证**:固定 seed 的 benchmark 可重复;统计满足 `total_changes == merged_pairs + redundant_lui_removed`;有效指令差值与转换分类一致;JSON/表格中保留零值与 `N/A`。 + +### 步骤 10:回归、文档与 PR + +**任务**: + +```bash +pytest tests/test_const_merge.py -v +pytest tests/ -q +make check # 若当前仓库提供 +``` + +**产出**:代码、测试、设计文档、开发文档、使用说明和 PR 描述。 +**验证**:所有验收标准完成。 + +--- + +## 6. 异常处理与边界条件 + +- [ ] 空字符串输入返回空结果和零变化; +- [ ] 只有注释或标签时不报错; +- [ ] 操作数缺失时跳过,不抛出未处理异常; +- [ ] `0x7FF`、`0x800`、`0xFFF` 正确符号扩展; +- [ ] `0x80000000` 等 32 位边界按 RV32 规范化; +- [ ] `-0x1` 可解析; +- [ ] `%hi(symbol)`、`%lo(symbol)` 保持原样; +- [ ] `t0` 和 `x5` 被视为同一寄存器; +- [ ] 标签或分支不会导致错误删除; +- [ ] `call` 后不复用 caller-saved 寄存器状态; +- [ ] 注释和标签不会静默丢失; +- [ ] 未知操作码采用保守策略; +- [ ] 达到迭代上限时能正常停止; +- [ ] 优化结果具有幂等性。 + +--- + +## 7. 测试与验证方案 + +### 7.1 单元测试矩阵 + +| 编号 | 场景 | 输入关键点 | 预期 | +|---|---|---|---| +| T01 | 基本合并 | `lui t0,0x12345` + `addi t0,t0,0x678` | 一条 `li` | +| T02 | 低位最大正数 | `0x7FF` | +2047 | +| T03 | 低位最小负数 | `0x800` | -2048 | +| T04 | 低位 -1 | `0xFFF` | -1 | +| T05 | 目标不同 | `addi t1,t0,...` | 不合并 | +| T06 | 源不同 | `addi t0,t1,...` | 不合并 | +| T07 | 中间注释 | 注释/空行 | 可按设计合并并保留注释 | +| T08 | 中间标签 | `L1:` | 不合并 | +| T09 | 重定位 | `%hi/%lo` | 不合并 | +| T10 | 安全冗余 LUI | 同块、同寄存器、同立即数 | 删除后一个 | +| T11 | 中间 clobber | `addi t0,t0,1` | 不删除 | +| T12 | 寄存器别名 | `t0` 与 `x5` | 正确识别修改 | +| T13 | 跨标签 | 标签两侧相同 `lui` | 不删除 | +| T14 | 跨分支 | 分支可能绕过前一 `lui` | 不删除 | +| T15 | 调用边界 | `call foo` | 清空状态 | +| T16 | 迭代暴露模式 | 删除冗余后出现相邻对 | 完成二次优化 | +| T17 | 幂等性 | 对结果再次优化 | 0 变化 | +| T18 | 空输入 | `""` | 空输出、0 变化 | + +### 7.2 集成测试 + +- **独立模块**:输入 `.s`,检查 `-o` 文件和 `-v`; +- **主 CLI**:检查 `--const-merge` 参数解析和配置映射; +- **CompilerDriver**:检查 post-pass 顺序; +- **汇编器**:优化前后均可汇编; +- **模拟器**:至少 3 个可执行用例结果一致; +- **回归**:全量测试通过。 + +### 7.3 建议 fixture + +```text +tests/fixtures/const_merge/ +├── basic_pair.s +├── sign_extension.s +├── redundant_lui.s +├── control_flow_guard.s +├── register_alias.s +└── relocation_noop.s +``` + +### 7.4 Benchmark 验证矩阵 + +| 类型 | 输入来源 | 目的 | 可以得出的结论 | +|---|---|---|---| +| 真实 case | `run_benchmark.py` 现有 ONNX/DSL case 生成的一份原始汇编 | 测主流水线实际模式覆盖与收益 | 当前后端是否产生目标模式,以及真实汇编/机器码是否变化 | +| Synthetic microbenchmark | `bench_const_merge.py` 固定 seed 人工汇编 | 测规则命中正确性、密度影响和 pass 开销 | 目标模式存在时的算法效果;不能外推为端到端收益 | +| 工具链等价用例 | 至少 3 个可执行 `.s` fixture | 验证汇编、反汇编和执行语义 | before/after 在限定输入上的语义一致性 | + +真实 case 报告必须逐 case 保留零命中结果,并提供汇总行;不得只展示发生变化的 case。`bench_runner.py` 可继续作为 DSL/IR 性能基线,但应在报告中注明它不覆盖 const-merge。 + +--- + +## 8. 验收标准(Definition of Done) + +- [ ] 设计文档已评审,范围和非目标明确; +- [ ] `merge_constants` 现有公共接口保持兼容; +- [ ] 使用共享汇编解析器,或对不迁移给出明确理由; +- [ ] `lui+addi` 正确处理 12 位符号扩展和 RV32 截断; +- [ ] 冗余 `lui` 只在可证明安全的基本块内删除; +- [ ] 正确处理 ABI/数字寄存器别名; +- [ ] 不优化重定位表达式; +- [ ] 实现固定点迭代和统计; +- [ ] 新增正向、负向、边界和反例测试; +- [ ] `pytest tests/test_const_merge.py -v` 通过; +- [ ] `pytest tests/ -q` 全量通过; +- [ ] 至少 3 个汇编样例通过工具链或模拟器等价验证; +- [ ] `run_benchmark.py` 基于同一份原始 RISC-V 汇编完成 const-merge A/B; +- [ ] 真实 case 逐项报告候选数、分类转换数、有效汇编指令数和 pass 耗时; +- [ ] 工具链可用时报告机器指令数和 `.text` 大小,不可用时明确标记 `N/A`; +- [ ] `bench_const_merge.py` 明确标记 synthetic,并覆盖 `lui+addi` 与重复 `lui` 两类密度; +- [ ] Benchmark 不再使用换行数作为指令数,且零命中 case 不被过滤; +- [ ] 文档明确说明 `li` 是伪指令; +- [ ] 主 CLI `--const-merge` 可用; +- [ ] PR 描述包含前后示例、测试结果和已知限制。 + +--- + +## 9. 风险评估与依赖 + +| 风险项 | 影响程度 | 缓解措施 | +|---|---|---| +| 错误理解 `li` 的性能收益 | 高 | 用 objdump 展开验证;统计分类 | +| 符号扩展或位宽错误 | 高 | 数学公式 + 边界测试 | +| 跨块错误删除 | 高 | 在标签和控制流边界清空状态 | +| 寄存器别名错误 | 高 | 统一映射为 `xN` | +| 解析器重构导致格式回归 | 中 | 先写 characterization tests | +| 现有后端不产生 `lui+addi` | 中 | 用独立 `.s` fixture;记录命中率 | +| GNU 汇编器版本差异 | 低/中 | 限定测试命令和 `-march=rv32im` | +| 与调度器顺序冲突 | 中 | const-merge 固定在 scheduler 前 | + +- **Python 依赖**:项目现有依赖,原则上不新增第三方 Python 库; +- **可选工具链**:`binutils-riscv64-linux-gnu`、TinyFive、Spike; +- **兼容性**:默认不改变未开启 `--const-merge` 时的编译结果。 + +--- + +## 10. 开发进度跟踪 + +以下日期为建议示例,可按课程安排调整: + +| 阶段 | 计划完成日期 | 状态 | +|---|---|---| +| 仓库走读与基线记录 | 2026-07-29 | ✅ 已完成 | +| 设计文档评审 | 2026-07-31 | ✅ 已完成 | +| Characterization tests | 2026-08-02 | ✅ 已完成 | +| 核心编码与重构 | 2026-08-07 | ✅ 已完成 | +| 边界测试与调试 | 2026-08-08 | ✅ 已完成 | +| 工具链等价验证 | 待工具链可用 | ⏸ N/A | +| 文档完善与 PR | 2026-08-08 | ✅ 已完成 | +| 代码审查与修订 | 2026-08-08 | 🔄 进行中 | + +--- + +## 11. 第一天实际工作清单(历史执行记录) + +以下内容保留为开发过程记录,相关步骤现已执行完毕,不代表当前仍需新建分支或保持核心代码未修改: + +```bash +# 1. 进入仓库和环境 +cd ScratchV +source .venv/bin/activate + +# 2. 确认分支与基线 +git status +git rev-parse HEAD +pytest tests/test_const_merge.py -v + +# 3. 找到所有相关代码 +grep -R "const_merge\|merge_constants\|--const-merge" -n scratchv tests docs + +# 4. 阅读顺序 +# scratchv/backend/const_merge.py +# scratchv/backend/_asm_parser.py +# scratchv/compiler.py +# scratchv/main.py +# tests/test_const_merge.py +# scratchv/backend/instruction_select.py +# scratchv/backend/asm_emit.py + +# 5. 新建分支 +git switch -c feat/topic14-const-merge +``` + +阅读时建立一张表: + +| 问题 | 当前答案 | 证据文件/行 | 是否要修改 | +|---|---|---|---| +| 优化位于哪个阶段? | post-codegen | `compiler.py` | 否/确认 | +| 当前 CLI 名称? | `--const-merge` | `main.py` | 评审 | +| 当前是否迭代? | 待代码确认 | `const_merge.py` | 是 | +| 是否跨标签清空状态? | 待测试确认 | `const_merge.py` | 是 | +| 是否处理寄存器别名? | 待测试确认 | `const_merge.py` | 是 | +| 当前后端是否产生 `lui+addi`? | 常量通常直接输出 `li` | `instruction_select.py` | 记录限制 | + +当天结束时应得到: + +1. 基线测试结果; +2. 相关文件关系图; +3. 5–10 个明确测试用例; +4. 设计文档 v0.1; +5. 需要确认的问题列表; +6. 尚未修改核心代码的干净功能分支。 + +--- + +## 12. PR 描述建议结构 + +```markdown +## What +实现/改进 RV32 常量加载合并优化: +- 安全合并 lui+addi +- 基本块内冗余 lui 消除 +- 寄存器别名处理 +- 固定点迭代与统计 + +## Why +现有实现缺少若干控制流与边界测试,且重复维护汇编解析逻辑。 + +## Safety +- 不跨标签/分支/调用 +- 不处理重定位表达式 +- 按 RV32 语义截断 + +## Tests +- pytest tests/test_const_merge.py -v +- pytest tests/ -q +- GNU assembler + objdump 对比 + +## Limitations +- 不支持 RV64 完整常量构造 +- 不做跨基本块数据流分析 +- li 为伪指令,不保证减少最终机器指令 +``` + +--- + +## 13. 参考资料 + +- ScratchV 课程首页: +- 课题 14 课程页: +- 课题 14 归档说明: +- 当前实现: +- 当前共享解析器: +- 当前测试: +- 贡献指南: +- RISC-V RV32I 规范: diff --git "a/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\346\212\200\346\234\257\350\256\276\350\256\241\346\226\207\346\241\243\345\210\235\347\250\277.md" "b/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\346\212\200\346\234\257\350\256\276\350\256\241\346\226\207\346\241\243\345\210\235\347\250\277.md" new file mode 100644 index 0000000..a2fda95 --- /dev/null +++ "b/docs/\350\257\276\351\242\23014-\345\270\270\351\207\217\345\212\240\350\275\275\345\220\210\345\271\266\344\274\230\345\214\226-\346\212\200\346\234\257\350\256\276\350\256\241\346\226\207\346\241\243\345\210\235\347\250\277.md" @@ -0,0 +1,763 @@ +# ScratchV 课题 14:常量加载合并优化技术设计文档 + +> **文档版本**:v0.3(核心实现与验证状态回填) +> **编写日期**:2026-07-28 +> **更新日期**:2026-08-08 +> **作者**:[yuki] +> **关联 Issue**:[#待补充] +> **涉及模块**:`scratchv/backend/const_merge.py`、`scratchv/backend/_asm_parser.py`、`scratchv/compiler.py`、`scratchv/main.py` +> **目标架构**:RV32I/RV32IM,GNU Assembler(GAS)语法 +> **课题定位**:在现有参考实现基础上,完成安全性分析、重构设计、测试补全和编译器集成验证 + +--- + +## 0. 文档说明与工作假设 + +ScratchV 主分支已经存在 `scratchv/backend/const_merge.py` 参考实现,课程网站也将课题 14 标记为“已完成”。因此,本课题不应简单复制现有代码,而应完成以下工作: + +1. 理解并准确说明 `lui`、`addi` 与 `li` 的语义; +2. 分析当前实现的适用范围和潜在安全问题; +3. 设计一个行为明确、保守安全、可测试的汇编后处理优化; +4. 补充边界测试、集成测试和优化统计; +5. 明确区分“汇编文本行数减少”和“真实机器指令数减少”。 + +本初稿采用以下假设,提交评审前需要与指导教师或项目维护者确认: + +- 第一阶段只保证 **RV32** 语义,不处理 RV64 任意宽常量展开; +- 第一阶段只在 **单基本块内部** 消除冗余 `lui`,不做跨基本块数据流分析; +- 只处理数值立即数,不处理 `%hi(symbol)`、`%lo(symbol)` 等重定位表达式; +- 保持公共函数 `merge_constants(asm_text) -> tuple[str, int]` 兼容; +- ScratchV 主命令行开关使用当前代码中的 `--const-merge`,而不是归档题目中的 `--merge-constants`。 + +### 0.1 实现结论(2026-08-08) + +第一阶段已按 RV32、单基本块、纯数值立即数和兼容旧 API 的假设完成。固定点循环实际采用“先删除冗余 `lui`,再合并 `lui+addi`”的顺序;这是为了让“重复 `lui` 删除后暴露相邻合并机会”的设计用例能够得到最简结果。共享解析器现已提供 `canonical_reg()`,详细接口为 `merge_constants_detailed()`。 + +真实 case A/B 抽查确认当前后端没有产生目标 `lui+addi` 模式,零命中作为正式实验结果保留。当前验证环境未提供 GNU RISC-V assembler 或 Spike,因此机器码、代码大小和执行等价指标为 `N/A`,不据此宣称运行性能提升。 + +--- + +## 一、功能介绍 + +### 1.1 背景 + +RISC-V 基础整数指令为定长编码。`addi` 只能携带 12 位有符号立即数,而 `lui` 将 20 位 U 型立即数放入目标寄存器的高 20 位、低 12 位补零。加载一般 32 位常量时,汇编器或编译器常使用: + +```asm +lui t0, 0x12345 +addi t0, t0, 0x678 +``` + +其结果为: + +```text +value = ((imm_hi & 0xFFFFF) << 12) + sign_extend_12(imm_lo) +``` + +在 RV32 中,最终结果按 32 位截断。 + +### 1.2 功能概述 + +常量加载合并优化是一个 **RISC-V 汇编层 post-pass**,在代码生成之后扫描汇编文本,执行两类转换: + +1. **常量序列规范化**:将符合条件的相邻 `lui + addi` 序列重写为 `li` 伪指令; +2. **冗余真实指令消除**:若同一物理寄存器在同一基本块中重复执行相同 `lui`,且中间没有被修改,则删除后一次 `lui`。 + +### 1.3 价值说明 + +- 统一常量加载的汇编表示,提高可读性; +- 删除可证明无用的 `lui`,减少真实机器指令; +- 为汇编级优化、指令统计和调度提供更规范的输入; +- 训练汇编解析、局部数据流跟踪、编译器 pass 集成和测试验证能力。 + +### 1.4 必须澄清的性能口径 + +`li` 是汇编器伪指令,不是 RV32I 的真实单条机器指令。对于较大的常量,汇编器通常仍会把 `li` 展开为 `lui + addi`。因此: + +- `lui + addi -> li` **必然减少汇编文本中的指令行数**; +- 但它 **不保证减少最终机器指令数或运行周期**; +- 真正可稳定减少机器指令的部分主要是安全的冗余 `lui` 消除; +- 若常量可用单条真实指令编码,例如有符号 12 位常量可使用 `addi rd, x0, imm`,则可以在后续增强版本中实现真实单指令替换。 + +文档、测试报告和优化统计不得把伪指令行数减少直接等同于机器指令数减少。 + +--- + +## 二、设计目标与非目标 + +### 2.1 设计目标 + +1. **语义正确**:正确处理 `addi` 12 位立即数的符号扩展和 RV32 截断; +2. **保守安全**:无法证明安全的序列不优化; +3. **基本块安全**:不跨标签、分支、跳转、调用和返回复用寄存器状态; +4. **别名正确**:把 `t0` 与 `x5` 等 ABI 名和数字寄存器名视为同一物理寄存器; +5. **幂等性**:对已优化结果再次运行不应继续发生变化; +6. **格式兼容**:尽可能保留标签、注释、空行和汇编指令顺序; +7. **可观测性**:统计合并对数、删除的冗余 `lui` 数和文本行数变化; +8. **可集成性**:支持独立模块 CLI 和 ScratchV 主编译流程开关。 + +### 2.2 非目标 + +第一阶段不实现: + +- RV64 任意 64 位常量的完整 `li` 展开与合并; +- `%hi(symbol)`、`%lo(symbol)`、`%pcrel_hi` 等重定位表达式优化; +- 跨基本块、跨循环或跨函数的全局常量传播; +- 基于控制流图的到达定义分析; +- 指令调度、寄存器分配或 IR 级常量折叠; +- 对所有 GNU/LLVM 汇编语法变体的完整支持。 + +--- + +## 三、当前系统分析 + +### 3.1 ScratchV 编译流水线中的位置 + +ScratchV 当前主要流程为: + +```text +DSL/ONNX 解析 + -> IR 优化 + -> RISC-V 指令选择 + -> 寄存器分配 + -> 汇编生成 + -> 汇编级 post-pass + 1. asm peephole + 2. const merge + 3. scheduler + 4. beautifier + 5. instruction counter + -> 输出 .s +``` + +常量加载合并属于 **汇编生成之后** 的局部优化,不修改 AST、IR 或机器指令数据结构。 + +### 3.2 当前代码接口 + +```python +from scratchv.backend.const_merge import merge_constants + +optimized_asm, changes = merge_constants(asm_text) +``` + +独立命令行: + +```bash +python -m scratchv.backend.const_merge input.s -o output.s -v +``` + +ScratchV 主命令行: + +```bash +scratchv input.dsl -o output.s --const-merge +``` + +### 3.3 基线实现问题(现已完成整改) + +以下问题记录的是 commit `109a6a2` 的开发前基线;v0.3 实现已逐项整改: + +1. `const_merge.py` 自己维护一套 `AsmInst` 解析逻辑,而项目已有共享的 `_asm_parser.py`; +2. 当前共享解析器的文档声称供 const-merge 使用,但实际参考实现仍重复解析代码; +3. 当前冗余 `lui` 跟踪没有明确在标签和控制流边界清空,跨基本块可能产生不安全删除; +4. 当前实现按字符串比较寄存器,`t0` 与 `x5` 的别名可能导致错误判断; +5. 现有测试主要覆盖正向样例,缺少标签、分支、寄存器别名、符号立即数、负十六进制和幂等性测试; +6. 归档课题要求“迭代扫描”,参考实现目前主要是一次合并加一次冗余消除; +7. 已在 commit `109a6a2` 核实:ScratchV 当前指令选择的多个常量路径直接生成 `MachineOp.LI`,所以主编译流程中 `lui + addi` 合并规则可能零命中;真实 case 必须测量并保留这一结果,同时用独立汇编 fixture 验证目标模式; +8. 当前 `changes` 只给总数,无法区分文本规范化与真实冗余指令删除。 + +--- + +## 四、语义模型与计算规则 + +### 4.1 12 位符号扩展 + +```python +def sign_extend_12(value: int) -> int: + """把 12 位编码(0..0xFFF)解释为有符号值。""" + value &= 0xFFF + return value - 0x1000 if value & 0x800 else value +``` + +边界: + +| 输入编码 | 符号扩展结果 | +|---|---:| +| `0x000` | 0 | +| `0x7FF` | 2047 | +| `0x800` | -2048 | +| `0xFFF` | -1 | + +### 4.2 RV32 最终常量 + +```python +upper_u32 = ((imm_hi & 0xFFFFF) << 12) & 0xFFFFFFFF +lower_s32 = sign_extend_12(imm_lo) +final_u32 = (upper_u32 + lower_s32) & 0xFFFFFFFF +final_s32 = final_u32 if final_u32 < 0x80000000 else final_u32 - 0x100000000 +``` + +输出 `li` 时建议统一使用有符号十进制 `final_s32`,避免 Python 无界整数与 RV32 位宽语义混淆。若项目更偏好十六进制输出,也必须明确按 32 位规范化。 + +### 4.3 示例 + +```asm +lui t0, 0x12345 +addi t0, t0, 0x678 +``` + +```text +0x12345000 + 0x678 = 0x12345678 +``` + +```asm +lui t0, 0x1 +addi t0, t0, 0x800 +``` + +`0x800` 作为 12 位立即数解释为 `-2048`: + +```text +0x00001000 - 0x800 = 0x00000800 +``` + +--- + +## 五、总体架构 + +### 5.1 模块职责 + +```text +输入汇编文本 + | + v +共享汇编解析器 parse_asm() + | + v +ParsedAsmLine 列表 + | + +--> 规则 A:lui + addi 合并 + | + +--> 规则 B:块内冗余 lui 消除 + | + +--> 固定点迭代与统计 + v +lines_to_asm() + | + v +输出汇编文本 + 统计信息 +``` + +### 5.2 建议数据结构 + +```python +from dataclasses import dataclass + +@dataclass +class ConstantMergeStats: + merged_pairs: int = 0 + redundant_lui_removed: int = 0 + iterations: int = 0 + + @property + def total_changes(self) -> int: + return self.merged_pairs + self.redundant_lui_removed +``` + +为保持兼容: + +```python +def merge_constants(asm_text: str) -> tuple[str, int]: + optimized, stats = ConstantMergeOptimizer().optimize(asm_text) + return optimized, stats.total_changes +``` + +详细统计可通过新增 API 获得,但第一阶段不强制修改已有调用方。 + +--- + +## 六、核心算法设计 + +### 6.1 汇编解析 + +优先复用: + +```python +from scratchv.backend._asm_parser import ( + ParsedAsmLine, + parse_asm, + lines_to_asm, + classify_def_use, +) +``` + +立即数解析建议使用: + +```python +def parse_numeric_imm(text: str) -> int | None: + try: + return int(text.strip(), 0) + except ValueError: + return None +``` + +这可同时处理十进制、十六进制及 `-0x1`。包含符号或重定位表达式时返回 `None`,放弃优化。 +对 `lui` 还必须校验立即数可表示为 20 位字段:接受有符号 20 位写法或 `0..0xFFFFF` 的编码写法,超出范围时拒绝优化,不得用按位与静默截断。 + +### 6.2 寄存器规范化 + +必须把 ABI 名称转换为统一的 `xN`: + +```text +t0 -> x5 +fp -> x8 +s0 -> x8 +a0 -> x10 +``` + +所有定义、使用和状态跟踪都使用共享解析器的 `canonical_reg()` 规范化名称,避免别名漏判。完整 ABI 映射由 `_asm_parser.py` 统一维护,包括 `zero/ra/sp/gp/tp`、`t0-t6`、`s0-s11`、`a0-a7` 和 `fp`。 + +### 6.3 规则 A:`lui + addi -> li` + +#### 匹配条件 + +设第一条有效指令为: + +```asm +lui rd, imm_hi +``` + +第二条有效指令为: + +```asm +addi rd2, rs1, imm_lo +``` + +仅当以下条件全部成立时转换: + +1. 两条指令在同一基本块内; +2. 中间最多只有空行或纯注释,没有标签或其他有效指令; +3. `canonical(rd) == canonical(rd2) == canonical(rs1)`; +4. 两个立即数都是纯数值; +5. 第二条指令不能带可作为跳转目标的标签; +6. 操作数数量合法; +7. 计算结果可按 RV32 语义规范化。 + +#### 输出 + +```asm +li rd, final_s32 # merged lui+addi +``` + +若第一条指令带标签,则标签保留在新 `li` 上;原有注释按约定合并或保留。 + +#### 不转换示例 + +```asm +lui t0, 0x12345 +addi t1, t0, 0x678 # 目标寄存器不同 +``` + +```asm +lui t0, %hi(symbol) +addi t0, t0, %lo(symbol) # 重定位表达式 +``` + +```asm +lui t0, 1 +L1: +addi t0, t0, 2 # addi 是跳转目标 +``` + +### 6.4 规则 B:基本块内冗余 `lui` 消除 + +维护: + +```python +lui_state: dict[str, int] +``` + +表示当前基本块中,某物理寄存器最后一次可确认的 `lui` 高位立即数。 + +处理规则: + +1. 遇到 `lui rd, imm`: + - 若 `lui_state[canonical(rd)] == imm`,当前 `lui` 冗余,可删除;删除后保持原有状态不变; + - 否则更新状态。 +2. 遇到会定义某寄存器的指令:清除该寄存器状态; +3. 遇到标签、条件分支、无条件跳转、函数调用或返回:清空全部状态; +4. 遇到未知指令且无法可靠判断定义集合:一律清空全部状态; +5. 不能跨基本块使用线性扫描状态。 + +#### 安全示例 + +```asm +lui t0, 0x10000 +addi t1, t0, 16 +lui t0, 0x10000 # 可删除,t0 中间未被修改 +addi t2, t0, 32 +``` + +#### 不安全示例 + +```asm +beq a0, zero, L1 +lui t0, 0x10000 +L1: +lui t0, 0x10000 # 不能删除:分支可能跳过前一个 lui +``` + +### 6.5 固定点迭代 + +单次规则应用可能暴露新的匹配: + +```asm +lui t0, 1 +lui t0, 1 # 删除后 +addi t0, t0, 2 # 与前一个 lui 相邻,可继续合并 +``` + +因此采用: + +```text +repeat: + 应用规则 B + 应用规则 A +until 本轮无变化 or 达到最大迭代次数 +``` + +先应用规则 B,才能在删除重复 `lui` 后把前一条 `lui` 与后续 `addi` 暴露为相邻匹配;若先应用规则 A,第二条 `lui` 可能被提前消费,无法得到最简结果。 + +最大迭代次数可设为 `max(1, len(lines))` 或一个保守上限。每次转换都会减少有效指令数,因此算法必然终止。 + +### 6.6 复杂度 + +若每轮线性扫描为 `O(n)`,最坏迭代次数为 `O(n)`,上界 `O(n^2)`。在典型汇编文件中迭代次数通常为 1–2。若需要严格线性复杂度,可调整规则顺序或使用工作队列,但不属于第一阶段目标。 + +--- + +## 七、接口设计 + +### 7.1 Python API + +兼容接口: + +```python +def merge_constants(asm_text: str) -> tuple[str, int]: + """返回优化后的汇编文本和转换总数。""" +``` + +建议新增详细接口: + +```python +def merge_constants_detailed( + asm_text: str, + *, + max_iterations: int | None = None, +) -> tuple[str, ConstantMergeStats]: + """返回优化后的汇编文本和分类统计。""" +``` + +### 7.2 独立 CLI + +```bash +python -m scratchv.backend.const_merge input.s -o output.s -v +``` + +`-v` 输出: + +```text +Constant merge: + merged lui+addi pairs: 3 + redundant lui removed: 2 + total transformations: 5 + iterations: 2 +``` + +### 7.3 ScratchV 主编译器开关 + +当前项目接口: + +```bash +scratchv model.onnx -o output.s --const-merge +scratchv --dsl example.dsl -o output.s --const-merge +``` + +配置字段: + +```python +CompilerConfig(const_merge=True) +``` + +### 7.4 Pass 顺序 + +建议维持: + +```text +asm peephole -> const merge -> scheduler -> beautifier -> instruction counter +``` + +原因: + +- 窥孔优化可能暴露新的相邻常量序列; +- 调度器可能打乱相邻关系,应在常量合并之后运行; +- 指令统计应在所有优化完成后运行。 + +--- + +## 八、正确性与安全约束 + +优化必须满足: + +```text +对任意允许的初始寄存器状态和执行路径, +优化前后程序在可观察行为上等价。 +``` + +第一阶段采用保守策略: + +- 不跨标签; +- 不跨控制流指令; +- 不跨调用; +- 不处理符号立即数; +- 不对未知操作码做激进定义/使用推断; +- 规范化寄存器别名; +- 按 RV32 位宽截断。 + +--- + +## 九、测试设计 + +### 9.1 单元测试分类 + +#### A. 汇编解析 + +- 十进制、十六进制、负十进制、负十六进制; +- 标签、注释、空行; +- `0(sp)` 等带括号操作数; +- 指令重建后内容可用。 + +#### B. `lui + addi` 合并 + +- 正常大常量 `0x12345678`; +- 低 12 位边界 `0x7FF`; +- 符号扩展边界 `0x800`; +- `0xFFF` 对应 `-1`; +- 结果为 `0x80000000`; +- 结果溢出后按 RV32 截断; +- ABI/数字寄存器别名混用; +- 中间有注释或空行; +- 中间有标签时不合并; +- 目标寄存器或源寄存器不同时不合并; +- `%hi/%lo` 不合并。 + +#### C. 冗余 `lui` 消除 + +- 同寄存器同立即数且未修改:删除; +- 同寄存器不同立即数:不删除; +- 中间写目标寄存器:不删除; +- 中间通过别名写目标寄存器:不删除; +- 跨标签:不删除; +- 跨分支、跳转、调用、返回:不删除; +- 未知操作码:保守处理。 + +#### D. 属性测试 + +- 幂等性:`opt(opt(x)) == opt(x)`; +- 无匹配输入:输出语义和有效指令不变; +- 转换总数与分类统计一致; +- 固定点迭代可发现删除后暴露的新模式。 + +### 9.2 集成测试 + +1. 独立 CLI 读取输入文件并生成输出文件; +2. `--verbose` 输出正确统计; +3. ScratchV 主 CLI 的 `--const-merge` 能正确传递到 `CompilerConfig`; +4. 编译器 post-pass 顺序符合设计; +5. 使用 GNU assembler 对优化前后汇编分别汇编; +6. 使用 `objdump -d -M no-aliases` 比较 `li` 的真实展开; +7. 使用 TinyFive、Spike 或项目模拟器执行可运行样例,确认最终寄存器/返回值相同。 + +### 9.3 关键验收用例 + +#### 用例 1:符号扩展 + +```asm +lui t0, 0x1 +addi t0, t0, 0x800 +``` + +期望: + +```asm +li t0, 2048 +``` + +#### 用例 2:安全冗余消除 + +```asm +lui t0, 0x10000 +addi t1, t0, 16 +lui t0, 0x10000 +addi t2, t0, 32 +``` + +期望:删除第二条 `lui`,其他指令不变。 + +#### 用例 3:跨块不得删除 + +```asm +beq a0, zero, L1 +lui t0, 0x10000 +L1: +lui t0, 0x10000 +``` + +期望:两个 `lui` 均保留。 + +#### 用例 4:寄存器别名 + +```asm +lui t0, 0x10000 +addi x5, x5, 1 +lui t0, 0x10000 +``` + +期望:第二条 `lui` 保留,因为 `x5` 与 `t0` 是同一寄存器。 + +--- + +## 十、Benchmark 设计 + +### 10.1 当前覆盖缺口 + +基于 commit `109a6a2`,三条现有 benchmark 路径的职责如下: + +| 路径 | 当前行为 | 对课题 14 的覆盖 | +|---|---|---| +| `benchmarks/bench_runner.py` | 运行 DSL/IR 解释与性能 case | 不生成 RISC-V 汇编,不经过 const-merge | +| `benchmarks/run_benchmark.py` | ONNX 解析、IR 优化、RISC-V/LLVM codegen | RISC-V 路径生成汇编,但未调用 const-merge;仅统计 `splitlines()` 行数 | +| `benchmarks/bench_const_merge.py` | 固定 seed 生成人工 `lui+addi` 并计时 | 能证明人工模式命中,但当前没有重复 `lui` 密度、分类统计或有效指令计数 | + +因此,本课题不能只修改 `bench_const_merge.py` 后声称优化已体现在“项目 benchmark case”中。设计必须同时提供真实 case A/B 和 synthetic microbenchmark,两者结论分开呈现。 + +### 10.2 真实 case 的单变量 A/B + +真实 case 复用 `run_benchmark.py` 的现有 ONNX/DSL 输入。在 RISC-V codegen 得到 `asm_before` 后,只对同一字符串调用一次 const-merge: + +```python +asm_before = asm_str +t0 = time.perf_counter() +asm_after, stats = merge_constants_detailed(asm_before) +pass_time_ms = (time.perf_counter() - t0) * 1000 +``` + +禁止为 A/B 重新运行两次完整编译,避免前端、IR 优化、指令选择或寄存器分配成为额外变量。禁止修改真实 case 以人工插入 `lui+addi`。 + +逐 case 记录: + +```text +case_name +lui_count_before +candidate_pairs +merged_pairs +redundant_lui_removed +asm_instructions_before / asm_instructions_after +machine_instructions_before / machine_instructions_after +code_size_before / code_size_after +pass_time_ms +output_equal +``` + +其中 `candidate_pairs` 是满足操作码、操作数、寄存器关系、数值立即数范围和同块相邻条件、可进入规则 A 转换的数量。基本块边界包括标签,以及分支、跳转、调用和返回之后的位置;纯注释和空行不是边界。该字段与 `merged_pairs` 分开,便于观察固定点转换前的输入模式数量。 + +### 10.3 有效汇编指令计数 + +`asm_text.count("\n")` 和 `len(asm_text.splitlines())` 都不是指令计数。计数器应解析每行并排除: + +- 空行与纯注释; +- 只有标签的行; +- `.text`、`.globl`、`.section` 等汇编伪操作。 + +标签与指令位于同一行时只计该指令。该口径用于 `asm_instructions_before/after`;它仍包含 `li` 等伪指令,因此不能替代机器指令数。 + +### 10.4 机器码与语义指标 + +工具链可用时,before/after 分别经 assembler 与 `objdump -d -M no-aliases`,统计真实机器指令和 `.text` 大小。`lui+addi -> li` 可能重新展开为两条机器指令,所以允许出现: + +```text +asm_instructions_after < asm_instructions_before +machine_instructions_after == machine_instructions_before +``` + +`output_equal` 只能由相同输入下的模拟器/执行结果比较得出;只完成 IR verifier、只成功汇编或只比较文本时记为 `N/A`,不能视为语义等价证明。外部工具不可用时,机器码、代码大小与执行等价字段均保留并标记 `N/A`,不得静默省略。 + +### 10.5 Synthetic microbenchmark + +`bench_const_merge.py` 继续负责可控压力测试,并明确打印 `benchmark_type=synthetic`。输入参数至少包括: + +- `num_instructions`:有效汇编指令规模; +- `pair_density`:`lui+addi` 模式密度; +- `redundant_lui_density`:同块重复 `lui` 模式密度; +- `seed` 与 `repeats`:可重复性及计时稳定性。 + +输出分类统计、有效指令差值及 pass 时间分布。人工构造的收益只能证明目标模式存在时优化器工作正常,并用于观察复杂度;不能外推为 ONNX/DSL 端到端收益。 + +### 10.6 结果解释与验收 + +真实 case 可能全部得到零候选,这是当前后端直接生成 `MachineOp.LI` 时的合理结果。报告必须保留所有 case(包括零值),并使用以下结论: + +> 当前真实代码生成路径没有产生 const-merge 的目标模式,因此本组 case 命中率为 0;synthetic microbenchmark 仅用于验证目标模式存在时的正确性和 pass 开销。 + +Benchmark 验收条件: + +1. 同一份原始汇编完成 A/B,唯一变量是 const-merge; +2. 逐 case 和汇总结果均包含候选、两类转换、有效指令数及耗时; +3. 工具链指标与 `output_equal` 不可测时明确为 `N/A`; +4. synthetic 输入同时覆盖两条优化规则,并固定 seed; +5. 满足 `total_changes == merged_pairs + redundant_lui_removed`; +6. 不将源汇编减少直接表述为机器指令、代码大小或周期减少。 + +--- + +## 十一、风险与权衡 + +| 风险 | 影响 | 缓解措施 | +|---|---|---| +| 把 `li` 当作真实单指令 | 高 | 区分文本行数、伪指令数和机器指令数 | +| `addi` 符号扩展错误 | 高 | 边界测试覆盖 `0x7FF/0x800/0xFFF` | +| 跨基本块错误删除 | 高 | 标签和控制流边界清空状态 | +| 寄存器 ABI 别名漏判 | 高 | 所有寄存器先规范化为 `xN` | +| 重定位表达式被错误解析 | 高 | 只接受纯数值立即数 | +| 注释或标签丢失 | 中 | 复用共享解析器并增加格式测试 | +| 当前后端已直接输出 `li`,命中率低 | 中 | 增加独立汇编 fixture;报告实际命中率 | +| 与调度器 pass 顺序冲突 | 中 | 固定在调度器之前运行 | +| 重构公共解析器引入回归 | 中 | 先补 characterization tests,再替换解析实现 | + +--- + +## 十二、待评审问题 + +提交设计评审时需要明确回答: + +1. 本课题是重现课程参考实现,还是改进主分支现有实现? +2. 第一阶段目标是 RV32 还是同时支持 RV64? +3. `lui+addi -> li` 的主要目标是代码可读性,还是需要证明真实机器指令减少? +4. 是否要求跨基本块优化?若要求,需要引入 CFG 和数据流分析,不应使用简单线性状态; +5. 是否必须沿用 `merge_constants(...)->(str, int)`,还是允许新增统计对象? +6. 是否需要把现有 `const_merge.py` 迁移到共享 `_asm_parser.py`? +7. 最终命令行名称以当前代码的 `--const-merge` 为准,还是兼容归档文档的 `--merge-constants`? +8. 真实 case A/B 是否合入 `run_benchmark.py` 的统一 JSON schema? +9. CI 是否提供 GNU RISC-V 工具链或模拟器;若不提供,机器码和语义指标是否作为可选阶段? +10. 零命中是否接受为当前主流水线的正式实验结论? + +--- + +## 十三、参考资料 + +- ScratchV 课程首页: +- 课题 14 当前课程页: +- 课题 14 归档说明: +- 当前参考实现: +- 当前测试: +- RISC-V ISA Manual: +- ScratchV 贡献指南: diff --git a/scratchv/backend/_asm_parser.py b/scratchv/backend/_asm_parser.py index 3cdef7e..16f5ed1 100644 --- a/scratchv/backend/_asm_parser.py +++ b/scratchv/backend/_asm_parser.py @@ -134,6 +134,32 @@ def is_comment_only(self) -> bool: "j", "jal", "jalr", "ret", "jr", } +# Integer register aliases from the RISC-V psABI. Assembly-level passes must +# key data-flow state by physical register rather than by spelling (``t0`` and +# ``x5`` are the same register). +_INTEGER_REGISTER_ALIASES: dict[str, str] = { + "zero": "x0", "ra": "x1", "sp": "x2", "gp": "x3", "tp": "x4", + "t0": "x5", "t1": "x6", "t2": "x7", "s0": "x8", "fp": "x8", + "s1": "x9", "a0": "x10", "a1": "x11", "a2": "x12", + "a3": "x13", "a4": "x14", "a5": "x15", "a6": "x16", + "a7": "x17", "s2": "x18", "s3": "x19", "s4": "x20", + "s5": "x21", "s6": "x22", "s7": "x23", "s8": "x24", + "s9": "x25", "s10": "x26", "s11": "x27", "t3": "x28", + "t4": "x29", "t5": "x30", "t6": "x31", +} + + +def canonical_reg(reg: str) -> str: + """Return the canonical ``xN`` spelling for an integer register. + + Unknown operands are returned lower-cased so callers can use this helper + without first proving that an operand is a register. + """ + name = reg.strip().lower() + if re.fullmatch(r"x([0-9]|[12][0-9]|3[01])", name): + return name + return _INTEGER_REGISTER_ALIASES.get(name, name) + # ═══════════════════════════════════════════════════════════════════════════════ # Parsing diff --git a/scratchv/backend/const_merge.py b/scratchv/backend/const_merge.py index d0394b0..a68c161 100644 --- a/scratchv/backend/const_merge.py +++ b/scratchv/backend/const_merge.py @@ -1,8 +1,7 @@ """Constant Load Merge Optimizer for RISC-V. -Detects and merges lui+addi instruction pairs into single li -pseudo-instructions, and eliminates redundant lui instructions -across basic blocks. +Detects and merges lui+addi instruction pairs into li pseudo-instructions, +and eliminates redundant lui instructions within basic blocks. Usage:: @@ -13,75 +12,33 @@ from __future__ import annotations import argparse -import re import sys +from dataclasses import dataclass from typing import Optional +from scratchv.backend._asm_parser import ( + ParsedAsmLine, + canonical_reg, + classify_def_use, + lines_to_asm, + parse_asm, +) + # --------------------------------------------------------------------------- # Data types # --------------------------------------------------------------------------- -class AsmInst: +class AsmInst(ParsedAsmLine): """Represents one parsed assembly instruction.""" def __init__(self, raw: str, lineno: int = 0): - self.raw = raw - self.lineno = lineno - self.label: Optional[str] = None - self.opcode: Optional[str] = None - self.operands: list[str] = [] - self.comment: Optional[str] = None - self._parse() - - def _parse(self) -> None: - """Parse the raw line into components.""" - stripped = self.raw.strip() - - # Empty line or pure comment - if not stripped or stripped.startswith("#"): - self.comment = stripped.lstrip("#").strip() - return - - # Separate code from comment - code = stripped - if "#" in stripped: - idx = stripped.find("#") - code = stripped[:idx].strip() - self.comment = stripped[idx + 1:].strip() - - # Check for label - label_match = re.match(r'^([A-Za-z_.][A-Za-z0-9_.]*):\s*(.*)', code) - if label_match: - self.label = label_match.group(1) - code = label_match.group(2).strip() - - if not code: - return - - # Extract opcode and operands - tokens = code.replace(",", " ").split() - if not tokens: - return - - self.opcode = tokens[0].lower().lstrip(".") - self.operands = tokens[1:] if len(tokens) > 1 else [] - - def to_asm(self) -> str: - """Reconstruct the assembly line.""" - parts = [] - if self.label: - parts.append(f"{self.label}:") - if self.opcode: - parts.append(f" {self.opcode}") - if self.operands: - parts.append(" " + ", ".join(self.operands)) - if self.comment: - parts.append(f" # {self.comment}") - result = "".join(parts) - if not result.strip() and self.raw.strip() == "": - return "" - return result + parsed = parse_asm(raw)[0] + super().__init__( + raw=parsed.raw, label=parsed.label, opcode=parsed.opcode, + operands=parsed.operands, comment=parsed.comment, lineno=lineno, + is_directive=parsed.is_directive, + ) def __repr__(self) -> str: return f"AsmInst({self.opcode}, {self.operands})" @@ -92,23 +49,25 @@ def __repr__(self) -> str: # --------------------------------------------------------------------------- def _parse_asm(asm_text: str) -> list[AsmInst]: - """Parse assembly text into AsmInst objects.""" - lines = asm_text.strip().split("\n") + """Compatibility wrapper around the shared assembly parser.""" + lines = asm_text.split("\n") return [AsmInst(line, lineno=i) for i, line in enumerate(lines)] def _insts_to_asm(insts: list[AsmInst]) -> str: - """Convert AsmInst list back to assembly string.""" - return "\n".join(inst.to_asm() for inst in insts) + """Compatibility wrapper around the shared assembly serializer.""" + return lines_to_asm(insts) def _parse_imm(s: str) -> Optional[int]: """Parse an immediate value string to int.""" try: - s = s.strip() - if s.startswith("0x") or s.startswith("0X"): - return int(s, 16) - return int(s) + text = s.strip() + try: + return int(text, 0) + except ValueError: + # Keep accepting ordinary decimal strings with leading zeroes. + return int(text, 10) except ValueError: return None @@ -135,46 +94,193 @@ def _l12(val: int) -> int: # Constant merge optimization # --------------------------------------------------------------------------- -# Standard register names -_STANDARD_REGS = { - "x0", "x1", "x2", "x3", "x4", "x5", "x6", "x7", - "x8", "x9", "x10", "x11", "x12", "x13", "x14", "x15", - "x16", "x17", "x18", "x19", "x20", "x21", "x22", "x23", - "x24", "x25", "x26", "x27", "x28", "x29", "x30", "x31", - "zero", "ra", "sp", "gp", "tp", - "t0", "t1", "t2", "t3", "t4", "t5", "t6", - "s0", "s1", "s2", "s3", "s4", "s5", "s6", "s7", "s8", - "s9", "s10", "s11", - "a0", "a1", "a2", "a3", "a4", "a5", "a6", "a7", - "fp", +@dataclass +class ConstantMergeStats: + """Observable results from one constant-merge optimization run.""" + + candidate_pairs: int = 0 + merged_pairs: int = 0 + redundant_lui_removed: int = 0 + iterations: int = 0 + + @property + def total_changes(self) -> int: + return self.merged_pairs + self.redundant_lui_removed + + +_CONTROL_FLOW = { + "beq", "bne", "blt", "bge", "bltu", "bgeu", "beqz", "bnez", + "blez", "bgtz", "bltz", "bgez", "j", "jr", "jal", "jalr", + "call", "tail", "ret", } +_KNOWN_OPCODES = { + "add", "addi", "sub", "mul", "div", "divu", "rem", "remu", + "sll", "slli", "srl", "srli", "sra", "srai", "xor", "xori", + "or", "ori", "and", "andi", "slt", "slti", "sltu", "sltiu", + "lui", "auipc", "li", "mv", "neg", "not", "seqz", "snez", + "lw", "lh", "lb", "lbu", "lhu", "sw", "sh", "sb", "nop", + "max", "min", "maxu", "minu", +} | _CONTROL_FLOW + + +def _is_separator(line: ParsedAsmLine) -> bool: + return line.opcode is None and line.label is None + + +def _signed_rv32(value: int) -> int: + value &= 0xFFFFFFFF + return value if value < 0x80000000 else value - 0x100000000 + + +def _parse_lui_imm(text: str) -> int | None: + """Parse a representable 20-bit LUI immediate without silent truncation.""" + value = _parse_imm(text) + # Accept either signed 20-bit spelling or the unsigned encoded field. + if value is None or not -(1 << 19) <= value <= 0xFFFFF: + return None + return value & 0xFFFFF + + +def _parse_addi_imm(text: str) -> int | None: + """Parse a 12-bit ADDI immediate in signed or encoded-field spelling.""" + value = _parse_imm(text) + if value is None or not -(1 << 11) <= value <= 0xFFF: + return None + return value + + +def _count_candidates(lines: list[ParsedAsmLine]) -> int: + count = 0 + for i, line in enumerate(lines): + if line.opcode != "lui": + continue + j = i + 1 + while j < len(lines) and _is_separator(lines[j]): + j += 1 + if j >= len(lines): + continue + addi = lines[j] + if ( + addi.opcode != "addi" or addi.label is not None + or len(line.operands) != 2 or len(addi.operands) != 3 + ): + continue + rd = canonical_reg(line.operands[0]) + if ( + rd == canonical_reg(addi.operands[0]) + and rd == canonical_reg(addi.operands[1]) + and _parse_lui_imm(line.operands[1]) is not None + and _parse_addi_imm(addi.operands[2]) is not None + ): + count += 1 + return count + + +def _remove_redundant_lui_once( + lines: list[ParsedAsmLine], +) -> tuple[list[ParsedAsmLine], int]: + result: list[ParsedAsmLine] = [] + state: dict[str, int] = {} + changes = 0 + + for line in lines: + # A label starts a new basic block, including ``label: instruction``. + if line.label is not None: + state.clear() + + if line.opcode is None: + result.append(line) + continue + if line.is_directive: + state.clear() + result.append(line) + continue + + opcode = line.opcode + if opcode == "lui" and len(line.operands) == 2: + reg = canonical_reg(line.operands[0]) + imm = _parse_lui_imm(line.operands[1]) + if imm is not None: + if state.get(reg) == imm: + # Preserve source comments even when the instruction goes. + comment = f"removed redundant lui {line.operands[0]}, {line.operands[1]}" + if line.comment: + comment += f"; {line.comment}" + result.append(ParsedAsmLine( + raw=f" # {comment}", comment=comment, + lineno=line.lineno, + )) + changes += 1 + continue + state[reg] = imm + result.append(line) + continue + + if opcode in _CONTROL_FLOW: + state.clear() + elif opcode not in _KNOWN_OPCODES: + state.clear() + else: + defines, _ = classify_def_use(line) + for reg in defines: + state.pop(canonical_reg(reg), None) + + result.append(line) + + return result, changes + + +def _merge_lui_addi_once( + lines: list[ParsedAsmLine], +) -> tuple[list[ParsedAsmLine], int]: + result: list[ParsedAsmLine] = [] + consumed: set[int] = set() + changes = 0 + + for i, line in enumerate(lines): + if i in consumed: + continue + if line.opcode != "lui" or len(line.operands) != 2: + result.append(line) + continue + + j = i + 1 + while j < len(lines) and _is_separator(lines[j]): + j += 1 + if j >= len(lines): + result.append(line) + continue + addi = lines[j] + if addi.opcode != "addi" or addi.label is not None or len(addi.operands) != 3: + result.append(line) + continue + + rd = canonical_reg(line.operands[0]) + if rd != canonical_reg(addi.operands[0]) or rd != canonical_reg(addi.operands[1]): + result.append(line) + continue + imm_hi = _parse_lui_imm(line.operands[1]) + imm_lo = _parse_addi_imm(addi.operands[2]) + if imm_hi is None or imm_lo is None: + result.append(line) + continue + + value = _signed_rv32((imm_hi << 12) + _sign_extend_12(imm_lo)) + comments = [c for c in (line.comment, addi.comment) if c] + comments.append(f"merged lui+addi -> {value}") + merged = ParsedAsmLine( + raw="", label=line.label, opcode="li", + operands=[line.operands[0], str(value)], + comment="; ".join(comments), lineno=line.lineno, + ) + result.append(merged) + # Preserve blank/comment lines between the pair in their original order. + result.extend(lines[i + 1:j]) + consumed.update(range(i + 1, j + 1)) + changes += 1 -def _is_reg(s: str) -> bool: - """Check if a string is a known register name.""" - return s.strip() in _STANDARD_REGS - - -def _is_clobbered(inst: AsmInst, reg: str) -> bool: - """Check if an instruction writes to the given register.""" - if inst.opcode is None: - return False - if not inst.operands: - return False - # For most instructions, the first operand is the destination - dst_clobbers = { - "add", "addi", "sub", "mul", "div", "rem", "sll", "srl", "sra", - "xor", "or", "and", "slt", "sltu", - "lui", "li", "mv", "lw", "lh", "lb", "lbu", "lhu", - "auipc", "jal", "jalr", - "xori", "ori", "andi", "slli", "srli", "srai", - "slti", "sltiu", - } - if inst.opcode in dst_clobbers: - return inst.operands[0] == reg - # For stores, the first operand is the value (doesn't clobber dest reg) - # For branches, no destination - return False + return result, changes def merge_constants(asm_text: str) -> tuple[str, int]: @@ -189,90 +295,28 @@ def merge_constants(asm_text: str) -> tuple[str, int]: ------- Tuple of (optimized_assembly_string, number_of_changes_made). """ - insts = _parse_asm(asm_text) - total_changes = 0 - - # --- Pass 1: Merge adjacent lui+addi pairs into li --- - new_insts: list[AsmInst] = [] - i = 0 - while i < len(insts): - inst = insts[i] - - # Check for lui followed by addi - if inst.opcode == "lui" and i + 1 < len(insts): - next_inst = insts[i + 1] - if (next_inst.opcode == "addi" - and inst.operands and next_inst.operands): - # Check: rd of lui == rd of addi, and rd == rs1 of addi - lui_rd = inst.operands[0] - if (len(next_inst.operands) >= 3 - and next_inst.operands[0] == lui_rd - and next_inst.operands[1] == lui_rd): - # Merge - imm_hi = ( - _parse_imm(inst.operands[1]) - if len(inst.operands) > 1 else None - ) - imm_lo = ( - _parse_imm(next_inst.operands[2]) - if len(next_inst.operands) > 2 else None - ) - - if imm_hi is not None and imm_lo is not None: - # Compute final constant - final_val = (imm_hi << 12) + _sign_extend_12(imm_lo) - # Replace with li - new_inst = AsmInst("") - new_inst.opcode = "li" - new_inst.operands = [lui_rd, str(final_val)] - new_inst.comment = ( - f"merged lui+addi -> {final_val}" - ) - new_insts.append(new_inst) - total_changes += 1 - i += 2 - continue - - new_insts.append(inst) - i += 1 - - insts = new_insts - - # --- Pass 2: Eliminate redundant lui --- - # Track the last upper-immediate value loaded into each register - # If a new lui loads the same value into the same register (and the - # register hasn't been clobbered in between), the second lui is redundant. - new_insts = [] - lui_state: dict[str, Optional[int]] = {} # reg -> upper imm value - - for inst in insts: - if inst.opcode == "lui" and inst.operands: - rd = inst.operands[0] - imm = ( - _parse_imm(inst.operands[1]) - if len(inst.operands) > 1 else None - ) - if rd in lui_state and lui_state[rd] == imm: - # Redundant: skip it, add a comment to the next instruction - total_changes += 1 - # Replace with a comment - comment_inst = AsmInst("") - comment_inst.comment = ( - f"peephole: removed redundant lui {rd}, {imm}" - ) - new_insts.append(comment_inst) - continue - else: - lui_state[rd] = imm - else: - # If this instruction writes to a tracked register, clear tracking - for reg in list(lui_state.keys()): - if _is_clobbered(inst, reg): - lui_state[reg] = None + optimized, stats = merge_constants_detailed(asm_text) + return optimized, stats.total_changes + + +def merge_constants_detailed( + asm_text: str, *, max_iterations: int | None = None, +) -> tuple[str, ConstantMergeStats]: + """Optimize assembly and return categorized transformation statistics.""" + lines = parse_asm(asm_text) + stats = ConstantMergeStats(candidate_pairs=_count_candidates(lines)) + limit = max_iterations if max_iterations is not None else max(1, len(lines)) - new_insts.append(inst) + for _ in range(max(0, limit)): + lines, removed = _remove_redundant_lui_once(lines) + lines, merged = _merge_lui_addi_once(lines) + stats.iterations += 1 + stats.redundant_lui_removed += removed + stats.merged_pairs += merged + if removed == 0 and merged == 0: + break - return _insts_to_asm(new_insts), total_changes + return lines_to_asm(lines), stats # --------------------------------------------------------------------------- @@ -302,11 +346,16 @@ def main() -> None: with open(args.input, "r") as f: asm_text = f.read() - result, changes = merge_constants(asm_text) + result, stats = merge_constants_detailed(asm_text) if args.verbose: print( - f"Constant merge: {changes} change(s) applied", + "Constant merge:\n" + f" candidate pairs: {stats.candidate_pairs}\n" + f" merged lui+addi pairs: {stats.merged_pairs}\n" + f" redundant lui removed: {stats.redundant_lui_removed}\n" + f" total transformations: {stats.total_changes}\n" + f" iterations: {stats.iterations}", file=sys.stderr, ) diff --git a/scratchv/compiler.py b/scratchv/compiler.py index a3484d2..bd5648a 100644 --- a/scratchv/compiler.py +++ b/scratchv/compiler.py @@ -450,10 +450,14 @@ def _run_asm_passes(self, asm_text: str, warnings: list[str]) -> str: warnings.append(f"Asm peephole: {changes} changes") if self.config.const_merge: - from scratchv.backend.const_merge import merge_constants - asm_text, changes = merge_constants(asm_text) - if changes: - warnings.append(f"Const merge: {changes} changes") + from scratchv.backend.const_merge import merge_constants_detailed + asm_text, stats = merge_constants_detailed(asm_text) + if stats.total_changes: + warnings.append( + f"Const merge: {stats.total_changes} changes " + f"({stats.merged_pairs} pairs, " + f"{stats.redundant_lui_removed} redundant lui)" + ) if self.config.schedule: from scratchv.backend.inst_scheduler import ( diff --git a/tests/test_backend.py b/tests/test_backend.py index 85f3111..3e7738f 100644 --- a/tests/test_backend.py +++ b/tests/test_backend.py @@ -5,6 +5,7 @@ from scratchv.backend.register_alloc import RegisterAllocator, MachineOp from scratchv.backend.asm_emit import AsmEmitter from scratchv.frontend.dsl_parser import DSLParser +from scratchv.compiler import CompilerConfig, CompilerDriver class TestInstructionSelect: @@ -108,3 +109,16 @@ def test_emit_relu(self): asm = emitter.emit() assert "max" in asm + + +class TestAsmPassIntegration: + def test_const_merge_config_runs_post_codegen_pass(self): + driver = CompilerDriver(CompilerConfig(const_merge=True)) + warnings: list[str] = [] + result = driver._run_asm_passes( + " lui t0, 1\n addi t0, t0, 2\n", warnings, + ) + assert "li t0, 4098" in result + assert len(warnings) == 1 + assert "Const merge: 1 changes" in warnings[0] + assert "1 pairs" in warnings[0] diff --git a/tests/test_const_merge.py b/tests/test_const_merge.py index c47bbb2..abf42c3 100644 --- a/tests/test_const_merge.py +++ b/tests/test_const_merge.py @@ -2,8 +2,10 @@ import pytest from scratchv.backend.const_merge import ( - merge_constants, AsmInst, _parse_asm, _insts_to_asm, + ConstantMergeStats, AsmInst, _insts_to_asm, _parse_asm, + merge_constants, merge_constants_detailed, ) +from scratchv.backend._asm_parser import canonical_reg class TestAsmInst: @@ -114,6 +116,119 @@ def test_sign_extension_correct(self): result, changes = merge_constants(asm) assert changes >= 1 assert "li" in result + assert "2048" in result + + def test_rv32_result_is_normalized(self): + asm = " lui t0, 0x80000\n addi t0, t0, 0\n" + result, changes = merge_constants(asm) + assert changes == 1 + assert "-2147483648" in result + + def test_negative_hex_immediate(self): + asm = " lui t0, -0x1\n addi t0, t0, -0x1\n" + result, changes = merge_constants(asm) + assert changes == 1 + assert "li t0, -4097" in result + + def test_relocation_is_not_optimized(self): + asm = " lui t0, %hi(symbol)\n addi t0, t0, %lo(symbol)\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert "%hi(symbol)" in result + assert "%lo(symbol)" in result + + @pytest.mark.parametrize("imm", ["0x100000", "-0x80001"]) + def test_out_of_range_lui_is_not_truncated(self, imm): + asm = f" lui t0, {imm}\n addi t0, t0, 1\n" + result, stats = merge_constants_detailed(asm) + assert stats.candidate_pairs == 0 + assert stats.total_changes == 0 + assert imm in result + + @pytest.mark.parametrize("imm", ["0x1000", "-2049"]) + def test_out_of_range_addi_is_not_truncated(self, imm): + asm = f" lui t0, 1\n addi t0, t0, {imm}\n" + result, stats = merge_constants_detailed(asm) + assert stats.candidate_pairs == 0 + assert stats.total_changes == 0 + assert imm in result + + def test_candidate_count_requires_mergeable_registers(self): + asm = " lui t0, 1\n addi t1, t0, 2\n" + _, stats = merge_constants_detailed(asm) + assert stats.candidate_pairs == 0 + + def test_comment_and_blank_between_pair_are_preserved(self): + asm = " lui t0, 1 # upper\n# keep me\n\n addi x5, t0, 2 # lower\n" + result, changes = merge_constants(asm) + assert changes == 1 + assert "li t0, 4098" in result + assert "# keep me" in result + assert "upper" in result and "lower" in result + + def test_label_prevents_pair_merge(self): + asm = " lui t0, 1\nL1:\n addi t0, t0, 2\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert "lui" in result and "addi" in result + + def test_label_on_lui_is_preserved_when_pair_merges(self): + asm = "L0: lui t0, 1\n addi t0, t0, 2\n" + result, changes = merge_constants(asm) + assert changes == 1 + assert "L0: li t0, 4098" in result + + def test_redundant_lui_does_not_cross_label(self): + asm = " lui t0, 1\nL1:\n lui x5, 1\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert result.count("lui") == 2 + + def test_redundant_lui_does_not_cross_branch(self): + asm = " lui t0, 1\n beq a0, zero, L1\n lui t0, 1\nL1:\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert result.count("lui") == 2 + + def test_alias_clobber_keeps_later_lui(self): + asm = " lui t0, 1\n add x5, a0, a1\n lui t0, 1\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert result.count("lui") == 2 + + def test_unknown_opcode_is_conservative(self): + asm = " lui t0, 1\n custom.op a0, a1\n lui t0, 1\n" + result, changes = merge_constants(asm) + assert changes == 0 + assert result.count("lui") == 2 + + def test_fixed_point_and_detailed_stats(self): + asm = " lui t0, 1\n lui x5, 1\n addi t0, x5, 2\n" + result, stats = merge_constants_detailed(asm) + assert isinstance(stats, ConstantMergeStats) + assert stats.candidate_pairs == 1 + assert stats.redundant_lui_removed == 1 + assert stats.merged_pairs == 1 + assert stats.total_changes == 2 + assert stats.iterations == 2 + assert "li t0, 4098" in result + + def test_idempotent(self): + asm = " lui t0, 1\n lui x5, 1\n addi t0, x5, 2\n" + once, first = merge_constants_detailed(asm) + twice, second = merge_constants_detailed(once) + assert first.total_changes == 2 + assert twice == once + assert second.total_changes == 0 + + def test_register_canonicalization(self): + assert canonical_reg("t0") == "x5" + assert canonical_reg("fp") == "x8" + assert canonical_reg("s0") == "x8" + assert canonical_reg("x0") == "x0" + assert canonical_reg("x31") == "x31" + assert canonical_reg("X10") == "x10" + assert canonical_reg("not_a_register") == "not_a_register" class TestCli: @@ -123,6 +238,19 @@ def test_main_importable(self): from scratchv.backend.const_merge import main assert callable(main) + def test_cli_writes_output_and_verbose_stats(self, tmp_path, capsys, monkeypatch): + from scratchv.backend.const_merge import main + source = tmp_path / "input.s" + output = tmp_path / "output.s" + source.write_text(" lui t0, 1\n addi t0, t0, 2\n") + monkeypatch.setattr( + "sys.argv", ["const_merge", str(source), "-o", str(output), "-v"], + ) + main() + captured = capsys.readouterr() + assert "merged lui+addi pairs: 1" in captured.err + assert "li t0, 4098" in output.read_text() + if __name__ == "__main__": pytest.main([__file__, "-v"])