课题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
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 < b,mask = b - a,dst = a + (b - a) = b。如果 a >= b,mask = 0,dst = 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()
详细任务
- 分析当前
instruction_select.py中已有的算子映射(如add → add,mul → mul等)。
- 识别缺失的常用算子:如
div(除法)、mod(取模)、sqrt、min/max等。
- 为缺失算子实现RISC-V指令映射(注意RISC-V整数除法需要
div/rem,浮点需要扩展指令集)。
- 添加
float64(双精度浮点)支持:增加新的寄存器类、加载存储指令(fld/fsd)、算术指令(fadd.d等)。
- 更新类型系统,在IR中区分
f64和f32。
- 编写测试用例验证新算子和新类型。
交付产物
- 更新后的
instruction_select.py和type_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 支持
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()
- 分析当前
instruction_select.py中已有的算子映射(如add→add,mul→mul等)。 - 识别缺失的常用算子:如
div(除法)、mod(取模)、sqrt、min/max等。 - 为缺失算子实现RISC-V指令映射(注意RISC-V整数除法需要
div/rem,浮点需要扩展指令集)。 - 添加
float64(双精度浮点)支持:增加新的寄存器类、加载存储指令(fld/fsd)、算术指令(fadd.d等)。 - 更新类型系统,在IR中区分
f64和f32。 - 编写测试用例验证新算子和新类型。
交付产物
- 更新后的
instruction_select.py和type_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 支持
instruction_select.py和type_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 支持
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 支持
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 属性来判断用单精度还是双精度指令 |
进阶阅读
- RISC-V F/D 扩展规范:RISC-V ISA Manual Vol. 1, Ch. 11-12
- 分支无编程技巧:Bit Twiddling Hacks
- 相关课题: 课题8 — 指令选择 | 课题17 — 寄存器分配
12周每周目标
- W1:学习项目当前指令选择模块,列出已支持的算子和类型。
- W2:识别缺失的常用整数算子(除法、取模),查阅RISC-V手册中
div/rem指令。
- W3:实现整数除法和取模的指令选择,编写简单DSL测试(
a / b)。
- W4:测试除法和取模的正确性,处理除零错误(可忽略或插入陷阱)。
- W5:学习RISC-V浮点扩展(F/D扩展),了解
fld/fsd和fadd.d等指令。
- W6:在IR中添加
float64类型,修改类型解析器。
- W7:实现
float64的加载和存储指令选择。
- W8:实现
float64算术指令(加、减、乘、除)。
- W9:实现
float64比较指令(feq.d, flt.d等)和条件分支。
- W10:编写测试用例:双精度浮点求和、点积等。验证模拟器支持。
- W11:为
sqrt、min/max等添加指令选择(可使用库调用或硬件指令)。
- W12:更新文档,撰写新算子、新类型的使用指南。
- W1:学习项目当前指令选择模块,列出已支持的算子和类型。
- W2:识别缺失的常用整数算子(除法、取模),查阅RISC-V手册中
div/rem指令。 - W3:实现整数除法和取模的指令选择,编写简单DSL测试(
a / b)。 - W4:测试除法和取模的正确性,处理除零错误(可忽略或插入陷阱)。
- W5:学习RISC-V浮点扩展(F/D扩展),了解
fld/fsd和fadd.d等指令。 - W6:在IR中添加
float64类型,修改类型解析器。 - W7:实现
float64的加载和存储指令选择。 - W8:实现
float64算术指令(加、减、乘、除)。 - W9:实现
float64比较指令(feq.d,flt.d等)和条件分支。 - W10:编写测试用例:双精度浮点求和、点积等。验证模拟器支持。
- W11:为
sqrt、min/max等添加指令选择(可使用库调用或硬件指令)。 - W12:更新文档,撰写新算子、新类型的使用指南。