课题8:后端指令选择

难度:中 | 类型:参考分析 | 源文件scratchv/backend/instruction_select.py, scratchv/backend/inst_select_ext.py 状态:✅ 已完成


概述

指令选择器(Instruction Selector)是编译器后端的第一个阶段,负责将平台无关的 IR 指令翻译为 RISC-V 特定的机器指令。理解它的映射策略(1:1 映射、展开、降级、内联循环)是理解整个后端管线的基础。


理解背景

是什么?

指令选择器负责将平台无关的 IR 指令翻译为RISC-V 特定的机器指令

IR Program (平台无关)           MachineInstr[] (RISC-V 特定)
  %3 = add %1, %2        →      ADD v3, v1, v2
  %5 = relu %4           →      MAX v5, v4, x0
  br_if %cmp, L1, L2     →      BNEZ vcmp, L1
                                J L2

为什么?

IR 里的操作码是抽象的——CONVRELUGEMM 这些不是 RISC-V 的机器指令。指令选择器要决定:

  1. 哪些 IR 指令能直接映射到 RISC-V 指令(如 ADDadd
  2. 哪些 IR 指令需要"降级"(lowering)——展开为一串 RISC-V 指令(如 RELUMAX rd, rs, x0
  3. 哪些 IR 指令需要展开为循环(如 CONV → 6 层嵌套循环的 load/mul/srai/add/store)

核心概念

1. 映射策略
策略 说明 例子
1:1 映射 IR 指令直接对应一条 RISC-V 指令 ADDadd
展开 IR 指令展开为多条 RISC-V 指令 BR_IFBNEZ + J
降级 高级 IR 操作降到低级指令序列 RELUMAX rd, rs, x0
内联循环 生成完整的嵌套循环 CONV → 6 层循环

2. 指令映射表(关键条目)
IR OpCode RISC-V 指令 策略
ADD, SUB, MUL, DIV add, sub, mul, div 1:1
RELU max rd, rs, x0 降级(1条伪指令)
SIGMOID 查表法 + 线性插值(~12条) 降级
GELU 分段多项式(~8条) 降级
CONV 6层嵌套循环 内联循环
GEMM 3层循环 + transB 内联循环
MAXPOOL 5层循环 + 边界检查 内联循环
BR J target 1:1
BR_IF cond, t, f BNEZ cond, t + J f 展开
LABEL name: 1:1
RETURN JALR x0, ra, 0 降级

3. 管线位置
IR Program
    
[InstructionSelector]       你现在在这里
      MachineInstr[] (虚拟寄存器)
[RegisterAllocator]         课题17
      MachineInstr[] (物理寄存器)
[AsmEmitter]                汇编文本

理解要点
  1. 掌握四种映射策略(1:1、展开、降级、内联循环)及其适用场景
  2. 理解 IR OpCode 到 RISC-V MachineOp 的完整映射表
  3. 能够追踪一个 IR 指令从 _select_instruction 分发到具体 handler 的完整路径
  4. 理解指令选择器输出(虚拟寄存器 MachineInstr)与寄存器分配器输入的关系
  5. 了解扩展指令选择器(inst_select_ext.py)对 F/D 浮点扩展的支持

交付产物
  • 指令映射关系图(IR OpCode → RISC-V 指令/序列)
  • 至少 2 个完整的 IR→MachineInstr 翻译追踪示例
  • 练习:添加一个新的 OpCode 映射

代码走读

核心分发逻辑
class InstructionSelector:
    def select(self, program: Program) -> list[MachineInstr]:
        result = []
        for func in program.functions:
            for block in func.basic_blocks:
                for instr in block.instructions:
                    # 按 OpCode 分发
                    selected = self._select_instruction(instr)
                    result.extend(selected)
        return result

    def _select_instruction(self, instr: Instruction):
        """按 OpCode 类型分发给具体的 _select_* 方法"""
        handlers = {
            OpCode.ADD: self._select_add,
            OpCode.SUB: self._select_sub,
            OpCode.CONV: self._select_conv,
            OpCode.RELU: self._select_relu,
            # ... 36 种 OpCode 各有对应的 handler
        }
        handler = handlers.get(instr.op)
        return handler(instr)

RELU 降级
def _select_relu(self, instr):
    """RELU(x) = max(x, 0)"""
    return [
        MachineInstr(
            MachineOp.MAX,
            dst=MachineOperand.vreg(instr.dest.name),
            src1=MachineOperand.vreg(instr.operands[0].name),
            src2=MachineOperand.reg("x0"),  # zero register
        )
    ]

BR_IF 展开
def _select_br_if(self, instr):
    """br_if cond, true_label, false_label"""
    return [
        MachineInstr(MachineOp.BNEZ, ...),   # if cond != 0 goto true
        MachineInstr(MachineOp.J, ...),       # goto false
    ]

动手练习

练习 1: 添加一个新的指令映射

假设你给 IR 添加了一个 NEG(取反)操作码。在 instruction_select.py 中实现 _select_neg,把它映射为 sub rd, x0, rs(0 - x = -x)。

练习 2: 追踪一个算子的完整翻译

写一个 DSL 程序包含 relu(add(a, b)),追踪从 DSL → IR → MachineInstr 的完整翻译过程。

练习 3: 对比 ScratchV 和 LLVM 的指令选择

对同一个 IR Program,分别用 InstructionSelector(RISC-V 路径)和 LLVMCodeGen(LLVM 路径)生成代码,对比输出。


常见坑
说明
虚拟寄存器 指令选择器的输出使用虚拟寄存器(vreg),还没有分配物理寄存器。如果直接给 AsmEmitter 会报错
CONV/GEMM 的内联 这些算子生成的是完整循环,不是几条指令,代码量很大。如果 IR 里有多个 Conv,注意总代码大小
OpCode 覆盖不全 如果 IR 里出现了指令选择器没处理的 OpCode,通常会静默跳过或报错
扩展指令选择器 inst_select_ext.py 支持更多 RISC-V 扩展(F/D 浮点等),一般场景用基础版就够了

进阶阅读

自学路线
  • 第 1 周:阅读 instruction_select.py 全文,理解 _select_instruction 的分发逻辑。用表格列出所有 36 种 OpCode 及其映射策略(1:1 / 展开 / 降级 / 内联循环)。
  • 第 2 周:选 3 个代表性的 IR 程序(分别含 1:1 映射、降级、内联循环),追踪从 IR 到 MachineInstr 的完整输出。画出每种策略的翻译流程图。
  • 第 3 周:对比 ScratchV 指令选择器与 LLVM 的 SelectionDAG 路径。阅读 LLVM 的 RISCVISelLowering.cpp 中同类算子的 lowering 方式,理解异同。
  • 第 4 周:尝试为指令选择器添加一个新的 OpCode 支持(如 SQRT)。编写完整测试:DSL 程序 → IR → MachineInstr → 汇编 → 仿真验证。