课题3:中间表示系统 (IR System)

难度:中 | 类型:项目实战 | 源文件scratchv/ir/types.py, scratchv/ir/builder.py, scratchv/ir/printer.py 状态:✅ 已完成


概述

IR(Intermediate Representation,中间表示)是 ScratchV 编译器的核心数据结构。采用三地址码(Three-Address Code)格式和 SSA(Static Single Assignment)设计。实现它需要定义 Program→Function→BasicBlock→Instruction 的层次结构、36 种 OpCode 枚举、Value 的 SSA 类型系统,以及 IRBuilder 和 IRPrinter 两个核心工具。


理解背景

是什么?

IR 是 ScratchV 编译器的核心数据结构。它夹在前端和后端之间,是所有模块沟通的"通用语言"。

ONNX/DSL → [前端] → IR → [优化器] → IR' → [后端] → RISC-V/LLVM IR
                      ↑        ↑        ↑
                   所有模块都读写 IR

IR 采用三地址码格式:每条指令最多三个操作数(一个输出 + 两个输入),像 %3 = add %1, %2

为什么?

为什么编译器需要一个"中间层"而不是直接把 ONNX 翻译成 RISC-V?

  1. 解耦:前端只需要产出 IR,后端只需要消费 IR。换一个前端(比如支持 PyTorch)不需要改后端
  2. 优化:所有优化 pass 在 IR 层面做,不用关心原始格式或目标架构
  3. 双路径支持:同一个 IR 既可以生成 RISC-V 也可以生成 LLVM IR
  4. 可调试:IR 是人类可读的文本,可以打印出来检查每一步

核心概念

1. 层次结构
Program                          ← 整个程序(一个 ONNX 模型 = 一个 Program)
  └── Function[]                 ← 函数列表(每个算子对应一个 Function)
        └── BasicBlock[]         ← 基本块列表
              └── Instruction[]  ← 指令列表(三地址码)

2. OpCode(36 种操作码)
类别 操作码 示例
算术 ADD, SUB, MUL, DIV, NEG, EXP %3 = add %1, %2
访存 LOAD, STORE, LOAD_CONST, ALLOCA %2 = load_const 3.14
控制流 BR, BR_IF, LABEL, RETURN, FOR, ENDFOR br_if %cmp, %L1, %L2
神经网络 CONV, GEMM, MATMUL, MAXPOOL, RELU, SIGMOID, SOFTMAX, GELU, DOT %out = conv %in, %w, %b

3. Value(SSA 值)
@dataclass
class Value:
    name: str           # SSA 名称: "%1", "%conv_result"
    dtype: DataType     # FLOAT32, INT32, FLOAT64, INT64
    constant: bool      # 是编译时常量吗?
    const_value: Any    # 常量值
    shape: tuple        # Tensor 形状: (1, 32, 64, 64)

4. SSA(Static Single Assignment)

每个变量只被赋值一次。如果需要修改,创建一个新版本:

%1 = add a, b       # 正确
%1 = add %1, c      # 错误!%1 被赋值了两次
%2 = add %1, c      # 正确:%2 是新变量

详细任务
  1. 定义 IR 核心数据结构:Program, Function, BasicBlock, Instruction 类。
  2. 设计并实现 36 种 OpCode 的完整枚举,按类别组织(算术/访存/控制流/神经网络)。
  3. 实现 Value 的 SSA 类型系统(name, dtype, constant, const_value, shape)。
  4. 实现 IRBuilder:new_value(), add_instruction(), create_function() 等核心方法。
  5. 实现 IRPrinter:将 IR 程序输出为人类可读的文本格式。
  6. 实现 IR 的序列化(IRSerializer):支持 JSON/二进制序列化和反序列化。
  7. 为 IR 添加基础验证逻辑(def-use 检查、SSA 唯一性检查、类型一致性检查)。
  8. 测试 IR 构建的完整性:从 DSL 程序和 ONNX 模型两个前端验证。
  9. 确保 IR 数据结构支持原地修改(优化 pass 的需求)。
  10. 编写完整测试,覆盖所有 OpCode 和边界情况。

交付产物
  • scratchv/ir/types.py — IR 核心类型定义
  • scratchv/ir/builder.py — IR 构建器
  • scratchv/ir/printer.py — IR 打印器
  • IR 验证器(def-use / SSA 检查)
  • 完整的单元测试

代码走读

IRBuilder 的关键方法
class IRBuilder:
    def new_value(self, name, dtype):
        """创建一个新的 SSA 值"""
        return Value(name=f"%{name}", dtype=dtype)

    def add_instruction(self, instr):
        """添加指令到当前基本块"""
        self.current_block.instructions.append(instr)

    def create_function(self, name):
        """创建新函数"""
        func = Function(name=name, basic_blocks=[])
        self.program.functions.append(func)
        return func

SSA 命名约定

ScratchV 的 IR 使用 % 前缀表示 SSA 变量,类似 LLVM IR:

%1 = const 3
%2 = const 5
%3 = add %1, %2

动手练习

练习 1: 手写 IR

写一个简单的算术表达式 (3 + 5) * 2,手工构建对应的 IR 指令序列。

练习 2: 阅读 IR 输出

用 DSL 写一个小程序,使用解析器生成 IR,打印出来阅读。

练习 3: 添加新的 OpCode

OpCode 枚举中添加一个新的操作码(比如 SQRT),然后思考: - 前端如何生成这个 OpCode? - 后端如何把它翻译成 RISC-V 指令?


常见坑
说明
SSA 违反 同一个 %name 被赋值两次会导致后续分析和优化错误
类型不匹配 FLOAT32 和 INT32 之间的运算需要显式转换
Shape 信息丢失 优化 pass 可能改变 tensor 形状,需要同步更新 Value.shape
OpCode 过多 神经网络操作码(CONV, GEMM 等)是高级的,后端需要"降级"(lowering)为基本算术指令的循环

进阶阅读

12周每周目标
  • W1:学习编译器 IR 的概念(三地址码、控制流图、基本块)。阅读龙书第 6 章。阅读 types.py 理解现有的 Program/Function/BasicBlock/Instruction 数据结构。
  • W2:设计 IR 的类层次结构。画出 UML 类图:Program 包含 Function[] → Function 包含 BasicBlock[] → BasicBlock 包含 Instruction[]。
  • W3:实现 Value 类和 SSA 命名系统。Value 需要包含 name, dtype, constant flag, const_value, shape。实现 %n 格式的自动命名。
  • W4:设计并实现 36 种 OpCode 枚举。分类:算术(6)、访存(4)、控制流(6)、神经网络(9)、其他。每个 OpCode 定义操作数数量和类型约束。
  • W5:实现 IRBuilder 核心方法:new_value(), add_instruction(), create_function(), new_block()。支持链式调用。
  • W6:实现 IRPrinter。输出格式参考 LLVM IR:层级缩进、SSA 变量、操作码+操作数。支持彩色终端输出。
  • W7:实现 IRSerializer:to_dict() / from_dict() 用于 JSON 序列化。支持跨进程传递(CI dashboard 场景)。
  • W8:实现基础 IR 验证器:def-use 链检查(每个使用的变量必须已定义)、SSA 唯一性检查(每个变量只能定义一次)、类型一致性检查(操作数类型匹配 OpCode 要求)。
  • W9:用 DSL 解析器生成真实的 IR 程序。测试 IRBuilder → IRPrinter → IRSerializer 的完整管线。验证输出正确性。
  • W10:用 ONNX 解析器生成 CNN 模型的 IR。测试大规模 IR(1000+ 条指令)的性能和正确性。
  • W11:编写完整单元测试:覆盖所有 36 种 OpCode、SSA 规则、序列化/反序列化、验证器。边界测试:空程序、单指令程序、循环嵌套。
  • W12:撰写文档(层次结构图、OpCode 参考表、IRBuilder API、使用示例),准备演示。