课题28:完善后端指令选择(扩展)

难度:高 | 类型:项目实战 | 源文件scratchv/backend/instruction_select.py | 行数:~600 状态:✅ 已完成


概述

为RISC-V后端增加对更多ONNX/DSL算子的支持,并添加新数据类型(如float64),扩展编译器的适用场景。


理解背景

是什么?

扩展指令选择器(inst_select_ext.py)在基础 InstructionSelector 之上增加了更多 RISC-V 操作的支持:平方根、绝对值、分支无 min/max、float64(D 扩展)等。

为什么?

基础指令选择器覆盖了 RV32IM 的核心指令。但实际应用需要: - 数学函数: sqrt, abs(某些激活函数需要) - float64: 双精度浮点(TensorFlow 模型可能用 float64) - 更多优化: min/max 的分支无实现

核心概念

新增操作
操作 实现 技巧
sqrt fsqrt.s 或库调用 sqrtf 可选硬件或软件
min (整数) slt + sub + and + add 四指令序列 分支无
max (整数) 复用 MAX 伪指令
abs (整数) srai 31 + xor + sub 三指令 位操作技巧
div/rem 原生 div/rem(M 扩展)

分支无 min 实现
slt tmp, a, b       # tmp = (a < b) ? 1 : 0
sub diff, b, a      # diff = b - a
and tmp, tmp, diff  # mask = tmp & diff
add dst, a, tmp     # dst = a + mask

原理:如果 a < bmask = b - adst = a + (b - a) = b。如果 a >= bmask = 0dst = a

Float64(D 扩展)

enable_fp64=True,自动对 float64 类型使用 D 扩展指令:

IR Op RISC-V 指令
add on f64 fadd.d
mul on f64 fmul.d
load of f64 fld
f64 compare flt.d / feq.d
f64↔f32 convert fcvt.s.d / fcvt.d.s

新增 MachineOp(20+ 个)

SQRT_S, SQRT_D, FMIN_D, FMAX_D, FABS_D, FNEG_D, FADD_D, FSUB_D, FMUL_D, FDIV_D, FLT_D, FEQ_D, FCVT_S_D, FCVT_D_S, FLD, FSD, SRAI, XOR, AND, SLT, REM

使用方法
from scratchv.backend.inst_select_ext import ExtendedInstructionSelector

# 启用 float64
selector = ExtendedInstructionSelector(program, enable_fp64=True)

# 启用硬件 sqrt
selector = ExtendedInstructionSelector(program, use_hardware_sqrt=True)

machine_instrs = selector.run()

详细任务
  1. 分析当前instruction_select.py中已有的算子映射(如addaddmulmul等)。
  2. 识别缺失的常用算子:如div(除法)、mod(取模)、sqrtmin/max等。
  3. 为缺失算子实现RISC-V指令映射(注意RISC-V整数除法需要div/rem,浮点需要扩展指令集)。
  4. 添加float64(双精度浮点)支持:增加新的寄存器类、加载存储指令(fld/fsd)、算术指令(fadd.d等)。
  5. 更新类型系统,在IR中区分f64f32
  6. 编写测试用例验证新算子和新类型。

交付产物
  • 更新后的instruction_select.pytype_system.py
  • 新增的测试程序(使用除法和双精度浮点)
  • 文档:支持的操作列表、数据类型说明

代码走读

分支无 abs
def _select_abs(self, instr):
    """abs(x) = (x ^ (x >> 31)) - (x >> 31)"""
    # srai 31: 提取符号位(全0=正,全1=负)
    # xor: 如果是负数,翻转所有位
    # sub: 如果是负数,+1(补码转换)
    return [
        MachineInstr(MachineOp.SRAI, tmp, src, 31),
        MachineInstr(MachineOp.XOR, dst1, src, tmp),
        MachineInstr(MachineOp.SUB, dst, dst1, tmp),
    ]

类型驱动的指令选择
def _select_add(self, instr):
    dtype = instr.dest.dtype if instr.dest else DataType.FLOAT32
    if dtype == DataType.FLOAT64 and self.enable_fp64:
        return MachineInstr(MachineOp.FADD_D, ...)
    elif dtype == DataType.FLOAT32:
        return MachineInstr(MachineOp.FADD_S, ...)
    else:
        return MachineInstr(MachineOp.ADD, ...)  # 整数

动手练习

练习 1: 添加 fneg 支持

fneg.s rd, rs = fsgnjn.s rd, rs, rs(RISC-V 无单独的 fneg 指令,用符号注入模拟)。

练习 2: 测试 float64 路径

写一个含有 float64 运算的 IR,用 ExtendedInstructionSelector 生成代码,对比和基础选择器的差异。


常见坑
说明
硬件 sqrt RV32IM 标准不包含 fsqrt,需要 F/D 扩展或 Zfa 扩展
float64 ABI 前 8 个浮点参数用 f0-f7,返回值用 f0,和整数 ABI 不同
类型检测 指令选择依赖 dtype 属性来判断用单精度还是双精度指令

进阶阅读

12周每周目标
  • W1:学习项目当前指令选择模块,列出已支持的算子和类型。
  • W2:识别缺失的常用整数算子(除法、取模),查阅RISC-V手册中div/rem指令。
  • W3:实现整数除法和取模的指令选择,编写简单DSL测试(a / b)。
  • W4:测试除法和取模的正确性,处理除零错误(可忽略或插入陷阱)。
  • W5:学习RISC-V浮点扩展(F/D扩展),了解fld/fsdfadd.d等指令。
  • W6:在IR中添加float64类型,修改类型解析器。
  • W7:实现float64的加载和存储指令选择。
  • W8:实现float64算术指令(加、减、乘、除)。
  • W9:实现float64比较指令(feq.d, flt.d等)和条件分支。
  • W10:编写测试用例:双精度浮点求和、点积等。验证模拟器支持。
  • W11:为sqrtmin/max等添加指令选择(可使用库调用或硬件指令)。
  • W12:更新文档,撰写新算子、新类型的使用指南。