Skip to content

Commit 149ae47

Browse files
jizhenjunclaude
andauthored
perf(kernels): edge blocks so matmul blocking covers any N; shared unroll skeleton (#103)
Two changes, one optimization and one refactor. 1. Edge blocks (bodies/matmul._blocked) The blocked path used to fall back to the generic triple loop whenever N was not a multiple of the block size. That was invisible while the data points were dense ladders, but once they became arbitrary integers it cost most of the benefit: on a declared sample (range:4:64:8) six of eight points took the fallback and scored 1.00x against it. Now the blocked path is three segments, with mainI = N rounded down to mr and mainJ likewise: main blocks [0, mainI) x [0, mainJ) full mr x nr right band [0, mainI) x [mainJ, N) one mr x 1 block per column bottom band [mainI, N) x [0, N) scalar triple loop Measured on that sample: N before after 4 520 536 0.97x (was already blocked) 13 20151 9101 2.21x 21 80460 33434 2.41x 30 228741 98059 2.33x 38 459271 190905 2.41x 47 861499 365786 2.36x 55 1373241 572307 2.40x 64 793767 793768 1.00x (was already blocked) total 3817650 -> 2063896 1.85x Against the -O0 reference the affected points went from about 3.4x to 5.4-9.0x, i.e. into the same band as the points that were already blocked. A defect was found and fixed while doing this: all three segments test at the bottom, so when a band is empty (mainJ == N, or mainI == N) the body ran once anyway and wrote past the end. N = 4, 52 and 64 segfaulted until each band got an explicit skip. Same class of mistake as the remainder tail loop in 5.1: the "remainder is zero" path has to be handled explicitly. 2. Shared unroll skeleton (loopgen.unrolled_loop) add and reducesum each carried their own copy of the unroll-plus-remainder-tail scaffolding (power-of-two check, round count, unrolled-segment end address, tail loop, pointer advance). It is now one function parameterised by the three things that actually differed: the per-round pointer list, the single-element body, and the exit label. Verified equivalent rather than assumed: normalised output is identical for all 18 generated artifacts (6 problems x unroll 1/8/32), the token sequences match, and clang's .text for the changed pair is byte-identical. The only source-level difference was whitespace -- the original used one space after the store mnemonic in the unrolled body and two in the tail, and the refactor unifies them. Also: the CLI now defaults -o to ./<problem>.s instead of printing to stdout, which is what scattered 190 scratch .s files through /tmp; `-o -` still prints. Docs updated with both findings (DEVELOPMENT.md 6.3, OPTIMIZATION.md 5 and a new 5.1 on why "fall back the whole problem" is the most expensive kind of conservative, and pipeline.py's table comments). All six problems still pass 10/10 on the platform evaluator. Co-authored-by: Claude <noreply@anthropic.com>
1 parent 8ea5d24 commit 149ae47

8 files changed

Lines changed: 313 additions & 131 deletions

File tree

‎scratchv/backend/kernels/DEVELOPMENT.md‎

Lines changed: 74 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1725,30 +1725,83 @@ cost = 执行到的指令数 + 15 × (执行期间遇到的未命中)
17251725

17261726
**在 N=64 上是 2.71×**,而且 `d_miss` 完全没变(773 → 773)——省下的全是指令。
17271727

1728-
#### ⚠️ 但分块只在 `N % 4 == 0` 时生效
1728+
#### 边角块:让分块对**任意 N** 生效
17291729

1730-
看同一条采样上的逐点:
1730+
**先看返工前的样子。** 最初的 `_blocked()` 开头有两条回退:
17311731

1732-
| N | N%4 | 分块 cost | 通用 cost | 快多少 |
1732+
```
1733+
bltu a2, unit, .Lgeneric # N 比粒度小 → 回退(这条留着)
1734+
andi TMP, a2, unit-1
1735+
bnez TMP, .Lgeneric # N 不是粒度的倍数 → 整题回退(这条是问题)
1736+
```
1737+
1738+
第二条在"数据点都是 4 的倍数"时看不出来。换成任意整数后,同一条采样上:
1739+
1740+
| N | N%4 | 返工前 | |
1741+
|---:|---:|---:|---|
1742+
| 4 | 0 | 520 | 走上分块 |
1743+
| 13 | 1 | 20,151 | **1.00×** ← 整题退给通用 |
1744+
| 21 | 1 | 80,460 | **1.00×** |
1745+
| 30 | 2 | 228,741 | **1.00×** |
1746+
| 38 | 2 | 459,271 | **1.00×** |
1747+
| 47 | 3 | 861,499 | **1.00×** |
1748+
| 55 | 3 | 1,373,241 | **1.00×** |
1749+
| 64 | 0 | 793,767 | 走上分块 |
1750+
1751+
**8 个点里 6 个一步没走到。**
1752+
1753+
> **"数据点可以是任意整数"的第一个后果。** 当数据点全都规整时,分块几乎总是
1754+
> 生效;一旦允许任意整数,它就成了少数派。**当初写"不是倍数就整题回退"时没想到
1755+
> 这一点 —— 因为那时看得到那些规整的数据点。**
1756+
1757+
**改法:把"整题回退"换成三段**(`mainI = N` 向下取整到 `mr`,`mainJ` 同理):
1758+
1759+
| 段 | 范围 | 用什么块 |
1760+
|---|---|---|
1761+
| **① 主块** | `i ∈ [0, mainI)`,`j ∈ [0, mainJ)` | 完整的 `mr×nr` |
1762+
| **② 右带** | `i ∈ [0, mainI)`,`j ∈ [mainJ, N)` | 每列一个 `mr×1` 块 |
1763+
| **③ 底带** | `i ∈ [mainI, N)`,`j ∈ [0, N)` | 标量三重循环 |
1764+
1765+
**同一个采样上,返工后:**
1766+
1767+
| N | N%4 | 返工前 | **返工后** | 快多少 |
17331768
|---:|---:|---:|---:|---:|
1734-
| 4 | 0 | 520 | 857 | 1.65× |
1735-
| 13 | 1 | 20,151 | 20,116 | **1.00×** ← 通用回退 |
1736-
| 21 | 1 | 80,460 | 80,425 | **1.00×** |
1737-
| 30 | 2 | 228,741 | 228,721 | **1.00×** |
1738-
| 38 | 2 | 459,271 | 459,251 | **1.00×** |
1739-
| 47 | 3 | 861,499 | 861,479 | **1.00×** |
1740-
| 55 | 3 | 1,373,241 | 1,373,221 | **1.00×** |
1741-
| 64 | 0 | 793,767 | 2,154,214 | 2.71× |
1742-
1743-
**8 个采样点里只有 2 个(4 和 64)走上分块**,其余 6 个是 `1.00×` ——
1744-
**完全没走到**,因为 `_blocked()` 开头就把 `N % 4 != 0` 的整题退给了通用实现。
1745-
1746-
**这不是 bug**(正确、且那些点仍比 `-O0` 快约 3.4×),但它是**一块明显的优化空间**:
1747-
给分块路径加**边角块处理**(部分 tile),那 6 个点也能往 9× 靠。
1748-
1749-
> **这正是"数据点可以是任意整数"的第一个后果。**
1750-
> 当数据点全都规整时,分块几乎总是生效;一旦允许任意整数,它就成了少数派。
1751-
> **当初写"`N % 4 != 0` 就整题回退"时没想到这一点 —— 因为那时看得到那些规整的数据点。**
1769+
| 4 | 0 | 520 | 536 | 0.97× |
1770+
| 13 | 1 | 20,151 | **9,101** | **2.21×** |
1771+
| 21 | 1 | 80,460 | **33,434** | **2.41×** |
1772+
| 30 | 2 | 228,741 | **98,059** | **2.33×** |
1773+
| 38 | 2 | 459,271 | **190,905** | **2.41×** |
1774+
| 47 | 3 | 861,499 | **365,786** | **2.36×** |
1775+
| 55 | 3 | 1,373,241 | **572,307** | **2.40×** |
1776+
| 64 | 0 | 793,767 | 793,768 | 1.00× |
1777+
| | | **3,817,650** | **2,063,896** | **1.85×** |
1778+
1779+
**非 4 倍数的点 2.2~2.4×;4 的倍数的点一点没变**(N=4 那个 0.97× 是多了两条守卫指令,
1780+
在那个极小规模上摊销不开)。
1781+
1782+
对 `-O0` 参考解的倍率也跟着上来了:那些点从 ~3.4× 提到 **5.4~9.0×**,
1783+
和 4 的倍数的点落在同一档。
1784+
1785+
#### ⚠️ 写边角块时踩的坑:循环是「先做后判」
1786+
1787+
返工后的第一次跑:**非 4 倍数的点全过,N=4/52/64 段错误。**
1788+
1789+
原因:这三段的循环都是**先执行后判断**(`do { 体 } while (J != N)`)。当边角为**空**时
1790+
——比如 `mainJ == N`(右边没有余列)——循环体**照样跑一次**,于是从第 N 列写出去。
1791+
1792+
**修法是进段之前先判空:**
1793+
1794+
```asm
1795+
beq MAINJ, a2, .Brs # 右边没有余列 → 整段跳过
1796+
... # ② 右带
1797+
.Brs:
1798+
beq MAINI, a2, .Lret # 下边没有余行 → 直接返回
1799+
... # ③ 底带
1800+
```
1801+
1802+
> **这不是分块特有的坑。** 任何"主循环 + 余数尾循环"的写法都有它:
1803+
> §5.1 那个展开循环里,`beqz t2, .Ltail` 和 `beqz a2, .Lret` 就是同一件事——
1804+
> **"余数为 0"这条路径必须显式处理,不能指望循环体自己不跑。**
17521805
17531806
**分块多大**:受**寄存器数量**限制。
17541807

‎scratchv/backend/kernels/OPTIMIZATION.md‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,15 +109,30 @@ N=4 时 cost=466,其中指令 256 条、`15×miss = 210`(**占 45%**,其
109109

110110
| # | 入口 | 现象/机制 | 怎么测 | 框架落点 | 状态 |
111111
|---|---|---|---|---|---|
112+
| **0** | **分块的边角块** ★已做 | 分块要求 `N % 4 == 0`,否则**整题退回通用**。数据点改成任意整数后,采样上 8 点里 **6 点一步没走到** | `probe.py` 看那些点是否从 `1.00x` 变成 ~2.4× | `bodies/matmul._blocked` 的三段 | ✅ **已完成**(采样合计 **1.85×**) |
112113
| **1** | **matmul N=64 的 L1 cache blocking** | §4.2 的容量门槛;可达 gap 的 ~79% | `measure.py` 看 `d_miss` 是否 3634 → ~1000 | `CacheBlocking` | `[P1]` 接口 |
113114
| **2** | **通用回退提速(GenericPlus)** | 通用版 9 条/元素 vs 特化 4.5 ⇒ **2.0×**;matmul 9 条/MAC vs 2.75 ⇒ **3.3×** | 强制 `ShapeKnowledge=Unknown` 跑对照 | `LoopFormLower` + `AddrStrengthReduce` | `[P1]` 接口 |
114115
| **3** | **小 N 的固定开销** | §4.3;N=4 单点 466 → 目标 ~440 | 把 256 条拆成 wrapper/分发/序言/主体分别量 | `DispatchLower` | `[P0]` 部分 |
115116
| **4** | **k 循环加「部分展开」档** | 现在只有「全展开」与「真循环」二选一;`i_miss` 在 N≥24 是 29,N=64 因转真循环降到 18 | 扫 unroll ∈ {2,4,8},看 instr 与 i_miss 的**联合** cost | `LoopFormLower` | `[P1]` 接口 |
116117
| **5** | **per-kernel 起始对齐** | 现在只有段级 `RODATA_PAD`,每个 kernel 的起始地址被动地由前面所有 kernel 的长度决定 | 给每个 kernel 加独立 `.balign` 扫 | `SectionLayout` | `[P1]` 接口 |
117118
| **6** | **`RODATA_PAD` 重扫** | 现值是对 10 个已知 N 扫出来的;**换 wrapper / 工具链 / N 集合都必须重扫** | `padsweep.py` | `SectionLayout` | `[P1]` 接口 |
118119
| **7** | **fp32 的低指令地板**(内测比赛2) | f32 matmul = `fmul.s`+`fadd.s` = **2 条/MAC 且无需 `slli`**;Q16 是 `mulh`+`add`+摊薄的 `slli` | 同 #1 的方法 | `DtypePolicy` | `[P0]` 本阶段 |
120+
| **8** | **q16 也能分块**(内测比赛1) | 现在 q16 的 matmul 完全没有分块(整数组 23 个,装不下 (4,4) 的 25 个)。**但要注意**:`blocked_regs_needed()` 只数数据寄存器、不算机制那 13~17 个,所以 (2,2)/(3,3) 能过容量检查却会**撞车** | 先用 `_blocked()` 单独试 (2,2),看生成的寄存器有没有重叠 | `blocked_regs_needed` + `_blocked` 的机制寄存器 | `[P1]` 未做 |
119121
| — | ~~指令调度~~ / ~~周期估算~~ | **成本口径不含周期**,收益恒为 0 | — | — | `[P2]` 不做 |
120122

123+
### 5.1 入口 0 的经验:**"整题回退"是最贵的一种保守**
124+
125+
写 `N % 4 != 0 → 整题回退` 时,代价只有"那些点慢一点"。**但"那些点"的数量取决于
126+
数据点怎么取** —— 全都是 4 的倍数时是 0 个,任意整数时是 70%。
127+
128+
**判断这类取舍,不要问"这个约束合不合理",要问"它排除掉的数据点占多少"。**
129+
130+
改法(三段:主块 + 右带 + 底带)见 `DEVELOPMENT.md` §6.3。
131+
132+
> 顺带一个**通用的坑**:三段循环都是"先做后判",所以**边角为空时整段必须显式跳过**
133+
> (`beq MAINJ, a2, .Brs`)。第一版漏了这个 → N=4/52/64 段错误。
134+
> **"余数为 0"这条路径永远要显式处理** —— 展开循环的 `beqz` 守卫是同一个道理。
135+
121136
**按「是否把某点从 <1.0 推到 ≥1.0」判据**:入口 1 修的是 N=64 单点(差 7.3%,最大单点);
122137
入口 3+4 修的是其余 9 个点(各差 2~4%,是「面」)。
123138

‎scratchv/backend/kernels/__main__.py‎

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
"""
77

88
import argparse
9+
import os
910

1011
from scratchv.backend.kernels.bodies import PROBLEMS
1112
from scratchv.backend.kernels.pipeline import build_program
@@ -14,7 +15,8 @@
1415
def main(argv=None) -> int:
1516
ap = argparse.ArgumentParser(description='ScratchV 算子内核生成器')
1617
ap.add_argument('--problem', help='题名,见 --list')
17-
ap.add_argument('-o', '--output', help='输出 .s 路径(默认打印到屏幕)')
18+
ap.add_argument('-o', '--output', default=None,
19+
help='输出路径;省略则写到 ./<题名>.s;给 "-" 打到屏幕')
1820
ap.add_argument('--unroll', type=int, default=None,
1921
help='展开因子(2 的幂);省略则用实测表')
2022
ap.add_argument('--list', action='store_true', help='列出所有题名')
@@ -29,12 +31,16 @@ def main(argv=None) -> int:
2931
ap.error('要么给 --problem,要么给 --list')
3032

3133
text = build_program(args.problem, args.unroll)
32-
if args.output:
33-
with open(args.output, 'w', encoding='utf-8') as f:
34-
f.write(text)
35-
print(f'已写出 {args.output}')
36-
else:
34+
35+
# 省略 -o 时**落到当前目录**(与 gcc/clang 同一约定),而不是打到屏幕。
36+
# 打屏幕那条路(`-o -`)留着,但必须是显式的 —— 否则产物会散在终端里没有人接。
37+
if args.output == '-':
3738
print(text)
39+
return 0
40+
path = args.output or f'{args.problem}.s'
41+
with open(path, 'w', encoding='utf-8') as f:
42+
f.write(text)
43+
print(f'已写出 {os.path.abspath(path)}')
3844
return 0
3945

4046

‎scratchv/backend/kernels/bodies/add.py‎

Lines changed: 29 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,11 @@
77
unroll == 1 通用循环,每元素 9 条指令
88
unroll > 1 展开 u 次 + 余数尾循环;偏移变立即数,循环开销摊到 u 个元素上
99
每元素 ≈ 4 + 5/u 条指令。**不需要知道 N** —— 尾循环处理 N % u。
10+
11+
骨架(幂校验、轮数、尾循环)在 `loopgen.unrolled_loop`,与 `reducesum` 共用。
1012
"""
1113

12-
from scratchv.backend.kernels.loopgen import prologue, epilogue, unrolled_body
14+
from scratchv.backend.kernels.loopgen import prologue, epilogue, unrolled_loop
1315

1416

1517
def build(target, dtype, unroll: int = 1) -> list[str]:
@@ -18,17 +20,22 @@ def build(target, dtype, unroll: int = 1) -> list[str]:
1820
ADD = dtype.add
1921
T1, T2 = dtype.tmp1, dtype.tmp2
2022

23+
# 单元素体:读 A、读 B、加、写 C。展开体与尾体共用这一段。
24+
single = [
25+
f' {L} {T1}, 0(a0)',
26+
f' {L} {T2}, 0(t1)',
27+
f' {ADD} {T1}, {T1}, {T2}',
28+
f' {S} {T1}, 0(a1)',
29+
]
30+
2131
if unroll == 1:
2232
return [
2333
*prologue(),
2434
' blez a2, .Lret', # N <= 0 直接返回
2535
' slli t0, a2, 2', # t0 = N * 4
2636
' add t1, a0, t0', # t1 = &B[0]
2737
'.Lloop:',
28-
f' {L} {T1}, 0(a0)', # 读 A[i]
29-
f' {L} {T2}, 0(t1)', # 读 B[i]
30-
f' {ADD} {T1}, {T1}, {T2}', # 算
31-
f' {S} {T1}, 0(a1)', # 写 C[i]
38+
*single,
3239
' addi a0, a0, 4', # 三个指针各前进一个元素
3340
' addi t1, t1, 4',
3441
' addi a1, a1, 4',
@@ -37,39 +44,25 @@ def build(target, dtype, unroll: int = 1) -> list[str]:
3744
*epilogue(),
3845
]
3946

40-
if unroll & (unroll - 1):
41-
raise ValueError(f'unroll 必须是 2 的幂,收到 {unroll}')
42-
43-
step = unroll * 4 # 每轮前进的字节数
44-
sh = unroll.bit_length() - 1 # log2(unroll)
47+
body: list[str] = []
48+
for k in range(unroll):
49+
off = k * 4
50+
body += [
51+
f' {L} {T1}, {off}(a0)',
52+
f' {L} {T2}, {off}(t1)',
53+
f' {ADD} {T1}, {T1}, {T2}',
54+
f' {S} {T1}, {off}(a1)',
55+
]
4556

4657
return [
4758
*prologue(),
48-
' blez a2, .Lret',
49-
' slli t0, a2, 2', # t0 = N * 4
50-
' add t1, a0, t0', # t1 = &B[0]
51-
f' srli t2, a2, {sh}', # t2 = N / unroll(完整轮数)
52-
' beqz t2, .Ltail', # N < unroll:直接走尾循环
53-
f' slli t3, t2, {sh + 2}', # t3 = 轮数 * unroll * 4
54-
' add t3, a0, t3', # t3 = 展开段结束的地址
55-
'.Lloop:',
56-
*unrolled_body(dtype, unroll), # 偏移全是立即数
57-
f' addi a0, a0, {step}', # 三个指针每轮只前进一次
58-
f' addi t1, t1, {step}',
59-
f' addi a1, a1, {step}',
60-
' bne a0, t3, .Lloop',
61-
'.Ltail:',
62-
f' andi a2, a2, {unroll - 1}', # 余数 = N % unroll
63-
' beqz a2, .Lret',
64-
'.Ltail_loop:',
65-
f' {L} {T1}, 0(a0)',
66-
f' {L} {T2}, 0(t1)',
67-
f' {ADD} {T1}, {T1}, {T2}',
68-
f' {S} {T1}, 0(a1)',
69-
' addi a0, a0, 4',
70-
' addi t1, t1, 4',
71-
' addi a1, a1, 4',
72-
' addi a2, a2, -1',
73-
' bnez a2, .Ltail_loop',
59+
*unrolled_loop(
60+
unroll=unroll,
61+
prologue=[' slli t0, a2, 2', # t0 = N * 4
62+
' add t1, a0, t0'], # t1 = &B[0]
63+
body=body,
64+
tail=single,
65+
ptrs=['a0', 't1', 'a1'], # 三个指针每轮各前进 u 个元素
66+
),
7467
*epilogue(),
7568
]

0 commit comments

Comments
 (0)