原文:https://malisper.me/how-ai-changes-the-economics-of-jit-compilers/
AI 如何改变 JIT 编译器的成本效益
从历史上看,JIT 编译一直是一门黑魔法。要编写一个快速的 JIT 编译器,你需要知道如何编写汇编代码。一个典型的例子是:当今没有一个生产就绪的数据库拥有自己的 JIT 编译器。它们要么使用 LLVM,要么生成 C/C++ 代码。这两种方案都存在编译时间长的问题,这限制了它们的适用性。现在,借助 AI 的使用,通过直接生成汇编代码来编写一个编译速度快的 JIT 编译器比以往任何时候都更容易。这也是新数据库能够在旧数据库基础上进行改进的领域之一。在构建 pgrust 时,我最初认为实现一个 JIT 编译器会非常困难。最后,我发现由于 AI 的帮助,它比我预期的要容易得多,这也是 pgrust 速度如此之快的原因之一。在这篇文章中,我将带你了解如何构建你自己的 JIT 编译器。我们将构建一个使用 JIT 编译的简单正则表达式引擎作为例子。
为什么使用 JIT 编译
JIT 编译是在运行时“即时”生成编译代码的做法。如果做得好,它可以带来巨大的性能提升,通常是 2-5 倍,有时甚至更多。JIT 编译的主要用例是当你在运行时获得的信息会极大地改变程序行为时。这在编程语言解释器中尤其常见;它们在运行时接收要执行的代码。JIT 编译器在编程语言之外的领域也很有用,例如解析数据。有时你在运行时才知道要解析的数据的模式(schema),JIT 可以帮助解决这个问题。
首先,让我们实现一个玩具正则表达式引擎。为了保持简单,我们只支持两个特性:字面量字符串和重复(即正则表达式*)。我们还将跳过解析器,将正则表达式表示为已经解析好的 Rust 结构。这意味着我们将能够支持如下字符串:
applesb(an)*
但不支持选择、后向查找或类似功能。
在代码中这相当简单。我们将有 3 种类型的节点:一个字面量字符串节点、一个重复节点和一个连接节点(它是两个节点的组合)。最终看起来像这样:
enumNode{Literal(&'staticstr),Concatenation(Box<Node>,Box<Node>),Repetition(Box<Node>),}fnliteral(text:&'staticstr)->Node{Node::Literal(text)}fnconcatenation(left:Node,right:Node)->Node{Node::Concatenation(Box::new(left),Box::new(right))}fnrepetition(body:Node)->Node{Node::Repetition(Box::new(body))}为我们的正则表达式引擎编写一个解释器也很直接:
fnmatch_node(node:&Node,input:&[u8],pos:usize,next:&dynFn(usize)->bool)->bool{matchnode{Node::Literal(text)=>{letliteral=text.as_bytes();input[pos..].starts_with(literal)&&next(pos+literal.len())}Node::Concatenation(left,right)=>{match_node(left,input,pos,&|left_end|{match_node(right,input,left_end,next)})}Node::Repetition(body)=>{match_node(body,input,pos,&|body_end|{match_node(node,input,body_end,next)})||next(pos)}}}fninterp_match(regex:&Node,input:&str)->bool{letbytes=input.as_bytes();match_node(regex,bytes,0,&|pos|pos==bytes.len())}这个正则表达式引擎非常简单。它不到 20 行代码,但让我们看看它在性能方面的表现。作为比较,我们将此代码与专门为正则表达式实现的手写代码进行对比。我们的示例将使用正则表达式b(an)*。手写代码最终看起来像:
fnhandwritten_b_an_star(input:&str)->bool{letbytes=input.as_bytes();letmutpos=0;ifpos==bytes.len()||bytes[pos]!=b'b'{returnfalse;}pos+=1;whilepos<bytes.len(){ifbytes[pos]!=b'a'{returnfalse;}pos+=1;ifpos==bytes.len()||bytes[pos]!=b'n'{returnfalse;}pos+=1;}true}(还有一些方法可以优化这段代码并使其更快,但出于我们的目的,它作为一个很好的比较基准)
当我对几个例子进行基准测试时,我得到手写版本比解释器快 10-20 倍。显然有很大的改进空间。
现在让我们看看如何使用 JIT 编译来获得一个性能与手写版本相当的正则表达式引擎。
如何进行 JIT 编译
JIT 编译代码有两个步骤。首先,为你想要运行的代码生成汇编代码。一旦你有了代码,然后将汇编代码打包成一个函数,你可以像调用程序中的任何其他代码一样调用它。
为了生成汇编代码,我们将使用一种称为“复制和修补”(copy-and-patch)的方法的变体。其思想是,我们有针对不同操作的汇编模板系列。这些模板被称为“模版”(stencils)。当我们想要 JIT 编译一个操作时,我们取关联的模版,并根据操作的具体细节进行小的调整。非常类似于填充真正的模版。通过将几个这样的填充模版串在一起,我们可以在运行时构建一个程序,其性能与手写版本相似。
以下是我们将采取的路径:首先,我们将查看我们想要为b(an)*生成的 ARM64 代码。然后,我们将重复的指令序列转化为可重用的模版,编写一个发射器(emitter),从正则表达式 AST 填充并组合这些模版,最后将生成的指令复制到可执行内存中,以便 Rust 可以像调用普通函数一样调用它们。
为了让你了解这是如何工作的,最简单的方法是从生成的代码开始,然后反向推导到 JIT 编译器本身。再次强调,我们正在处理正则表达式b(an)*。为了说明一些设计决策:
- 我们将使用一个栈进行回溯(backtracking)。栈将跟踪如果我们在正则表达式中遇到死路时应该转到的状态。
- 我们匹配的字符串将以空字节(NUL byte)结尾。这意味着如果我们到达字符串的末尾,我们的任何字符比较都会自动失败。这意味着我们在任何时候都不需要进行长度比较。
- 对于程序状态,我们将使用以下寄存器:
x0– 字符串中的当前位置和返回值x1– 用于回溯的栈顶x2– 用于回溯的栈底(这是确定栈是否为空所必需的)x9– 用作临时变量
- 对于程序的输入,我们将传入:
x0– 指向字符串开头的指针x1– 指向我们将用于栈的内存的指针
生成的 ARM64 代码
现在我们已经处理好了这些,让我们逐部分地浏览生成的汇编代码。这专门针对 macOS 上的 ARM64。首先,我们有 prologue,它初始化程序。它所做的只是通过将栈顶和栈底设置为传入的值来初始化栈:
0: aa0103e2 mov x2, x1接下来,我们有检查字符b的代码。如果它看到一个不是b的字符,我们跳转到一个处理回退逻辑的代码块。否则,我们推进字符串中的位置:
; CHAR 'b' 4: 39400009 ldrb w9, [x0] ; 加载当前输入字节 8: 7101893f cmp w9, #0x62 ; 是 'b' 吗? c: 54000281 b.ne 0x5c ; 不是 -> 回退块 10: 91000400 add x0, x0, #1 ; 是 -> 推进输入接下来,我们有重复(an)*。对于重复,我们需要进行回溯。如果我们在这里回溯,这意味着我们立即跳转到循环结束。这意味着我们需要将循环后的指令地址和我们在字符串中的位置都存储在栈上。
14: d2800989 movz x9, #0x004c ; 构建恢复地址 18: f2a00009 movk x9, #0x0000, lsl #16 ; = 0x1_0000_004c 1c: f2c00029 movk x9, #0x0001, lsl #32 ; (循环出口) 20: f2e00009 movk x9, #0x0000, lsl #48 ; 24: a8810029 stp x9, x0, [x1], #16 ; 将 (exit, pos) 压入栈有了这些,我们现在可以执行循环体了。这将检查字符a和n,如果看到它们,就回到循环顶部,但是是在一个新的字符串位置。
; CHAR 'a' 28: 39400009 ldrb w9, [x0] 2c: 7101853f cmp w9, #0x61 ; 'a'? 30: 54000161 b.ne 0x5c ; 不是 -> 回退块 34: 91000400 add x0, x0, #1 ; CHAR 'n' 38: 39400009 ldrb w9, [x0] 3c: 7101b93f cmp w9, #0x6e ; 'n'? 40: 540000e1 b.ne 0x5c ; 不是 -> 回退块 44: 91000400 add x0, x0, #1 ; JMP 48: 17fffff3 b 0x14 ; 回到循环顶部现在我们过了循环。这是我们回溯时将跳转到的位置。一旦我们完成了重复,就到了正则表达式的末尾。我们现在要做的只是检查我们是否在字符串的末尾。如果我们在末尾,我们返回 1 表示成功。如果我们不在,这意味着正则表达式匹配失败,我们需要运行失败逻辑进行回退。
4c: 39400009 ldrb w9, [x0] 50: 35000069 cbnz w9, 0x5c ; 不是 NUL -> 回退块 54: d2800020 mov x0, #1 ; 成功 58: d65f03c0 ret最后,我们有回退逻辑。这检查栈是否为空。如果是,我们返回 0。如果不是,我们从栈中弹出回退地址和回退字符串位置,然后跳转到回退地址。
5c: eb02003f cmp x1, x2 ; 还有帧吗? 60: 54000060 b.eq 0x6c ; 没有 -> 放弃 64: a9ff0029 ldp x9, x0, [x1, #-16]! ; 弹出 (resume, pos) 68: d61f0120 br x9 ; 跳转到那里 6c: d2800000 mov x0, #0 ; 不匹配 70: d65f03c0 ret构建模版
现在你已经有机会看到编译后的代码,你应该开始理解复制和修补编译器是如何工作的。我们有共同的指令集,它们之间只有微小的差异。对于这些函数块中的每一个,我们可以编写一个函数来生成相应的代码。每个函数将接收用于修改代码的值。例如,stencil_char的参数之一将是正则表达式中要比较的字符。我们将该字符直接插入到机器码中。
Prologue 很简单,因为它只是一个代码块:
constPROLOGUE_WORDS:usize=1;fnstencil_prologue()->[u32;PROLOGUE_WORDS]{[0xAA0103E2]// mov x2, x1}对于字符比较,我们需要插入我们正在比较的字符以及跳转到回退逻辑的位置:
constCHAR_WORDS:usize=4;fnstencil_char(byte:u8,stencil_pos:usize,fail_pos:usize)->[u32;CHAR_WORDS]{[0x39400009,// ldrb w9, [x0]0x7100013F|((byteasu32)<<10),// cmp w9, #byte0x54000001|cond_branch_offset(stencil_pos+2,fail_pos),// b.ne fail0x91000400,// add x0, x0, #1]}对于重复,我们有循环开始处压入栈的操作和跳转到循环结束的操作:
constSPLIT_WORDS:usize=5;fnstencil_split(resume_addr:u64)->[u32;SPLIT_WORDS]{[0xD2800009|addr_bits(resume_addr,0),// movz x9, #addr[0..16]0xF2A00009|addr_bits(resume_addr,1),// movk x9, #addr[16..32], lsl 160xF2C00009|addr_bits(resume_addr,2),// movk x9, #addr[32..48], lsl 320xF2E00009|addr_bits(resume_addr,3),// movk x9, #addr[48..64], lsl 480xA8810029,// stp x9, x0, [x1], #16]}constJMP_WORDS:usize=1;fnstencil_jmp(stencil_pos:usize,target_pos:usize)->[u32;JMP_WORDS]{[0x14000000|branch_offset(stencil_pos,target_pos)]// b target}然后我们有匹配和失败块,它们相当简洁:
constMATCH_WORDS:usize=4;fnstencil_match(stencil_pos:usize,fail_pos:usize)->[u32;MATCH_WORDS]{[0x39400009,// ldrb w9, [x0]0x35000009|cond_branch_offset(stencil_pos+1,fail_pos),// cbnz w9, fail0xD2800020,// mov x0, #10xD65F03C0,// ret]}constFAIL_WORDS:usize=6;fnstencil_fail()->[u32;FAIL_WORDS]{[0xEB02003F,// cmp x1, x20x54000060,// b.eq +3 (to the mov below)0xA9FF0029,// ldp x9, x0, [x1, #-16]!0xD61F0120,// br x90xD2800000,// mov x0, #00xD65F03C0,// ret]}为了完整性,这里是我们使用的辅助函数,它们只是帮助我们将特定数据插入到指令中:
// 计算条件分支(b.ne / cbnz)的分支偏移字段:// 从分支到目标的指令数,存储在位 5..24 中。fncond_branch_offset(branch_pos:usize,target_pos:usize)->u32{letinstr_count=target_posasi64-branch_posasi64;// 可能为负(((instr_countasu64)&0x7FFFF)<<5)asu32}// 计算无条件分支(b)的分支偏移字段:// 思路相同,但存储在位 0..26 中。fnbranch_offset(branch_pos:usize,target_pos:usize)->u32{letinstr_count=target_posasi64-branch_posasi64;// 可能为负((instr_countasu64)&0x3FF_FFFF)asu32}// 提取绝对地址的 16 位,为 movz/movk 立即数定位。fnaddr_bits(addr:u64,part:usize)->u32{(((addr>>(16*part))&0xFFFF)asu32)<<5}发射代码
现在是驱动它的代码:
// 计算一个节点编译成多少条指令。fnnode_words(node:&Node)->usize{matchnode{Node::Literal(text)=>text.len()*CHAR_WORDS,Node::Concatenation(left,right)=>node_words(left)+node_words(right),Node::Repetition(body)=>SPLIT_WORDS+node_words(body)+JMP_WORDS,}}structEmitter{code:Vec<u32>,fail:usize,// 共享失败块的偏移(以字为单位)base:u64,// code[0] 的运行时地址,用于绝对地址空洞}implEmitter{// 返回下一条指令将被放置的偏移量。fnpos(&self)->usize{self.code.len()}// 将填充的模版追加到代码缓冲区。fnemit(&mutself,stencil:&[u32]){self.code.extend_from_slice(stencil);}// 为一个节点发射代码,递归处理子节点。fnemit_node(&mutself,node:&Node){matchnode{Node::Literal(text)=>{for&byteintext.as_bytes(){self.emit(&stencil_char(byte,self.pos(),self.fail));}}Node::Concatenation(left,right)=>{self.emit_node(left);self.emit_node(right);}Node::Repetition(body)=>{letsplit_at=self.pos();letexit=split_at+SPLIT_WORDS+node_words(body)+JMP_WORDS;self.emit(&stencil_split(self.base+exitasu64*4));self.emit_node(body);self.emit(&stencil_jmp(self.pos(),split_at));}}}}// 生成完整程序:prologue、编译后的 AST、MATCH、失败块。fngenerate_code(regex:&Node,base:u64)->Vec<u32>{letnwords=PROLOGUE_WORDS+node_words(regex)+MATCH_WORDS+FAIL_WORDS;letmutemitter=Emitter{code:Vec::with_capacity(nwords),fail:nwords-FAIL_WORDS,base,};emitter.emit(&stencil_prologue());emitter.emit_node(regex);letmatch_at=emitter.pos();emitter.emit(&stencil_match(match_at,emitter.fail));emitter.emit(&stencil_fail());assert_eq!(emitter.pos(),nwords);emitter.code}这就是困难的部分!就我个人而言,编写汇编是我觉得 AI 最有帮助的地方。我对汇编的主要经验是完成 microcorruption CTF。我自己从未真正编写过汇编。我真的很难弄清楚需要的具体指令以及如何修改它们以获得我想要的输出。有了 AI,我可以给我的编码代理(coding agent)一个 JIT 编译器工作方式的总体框架,它就可以为我处理许多这些细节。
加载机器码
为了完成我们的编译器,我们需要实际加载代码。为此,我们将使用mmap分配一块可读、可写和可执行的内存。然后我们将代码复制到该内存中,并将该内存块转换为一个函数,然后调用它:
constBSTACK_MAX:usize=4096;// 这些函数包含在 mac 系统库中unsafeextern"C"{fnpthread_jit_write_protect_np(enabled:libc::c_int);fnsys_icache_invalidate(start:*mutlibc::c_void,len:libc::size_t);}typeMatchFn=unsafeextern"C"fn(input:*constu8,bstack:*mutu64)->u64;structJit{buf:*mutu32,nbytes:usize,bstack:Vec<u64>,}implJit{fncompile(regex:&Node)->Jit{letnwords=PROLOGUE_WORDS+node_words(regex)+MATCH_WORDS+FAIL_WORDS;letnbytes=nwords*4;unsafe{letbuf=libc::mmap(std::ptr::null_mut(),nbytes,libc::PROT_READ|libc::PROT_WRITE|libc::PROT_EXEC,libc::MAP_PRIVATE|libc::MAP_ANON|libc::MAP_JIT,-1,0,)as*mutu32;assert!(bufas*mutlibc::c_void!=libc::MAP_FAILED,"mmap failed");letcode=generate_code(regex,bufasu64);pthread_jit_write_protect_np(0);// 使区域可写(Apple W^X)std::slice::from_raw_parts_mut(buf,code.len()).copy_from_slice(&code);pthread_jit_write_protect_np(1);// 恢复为可执行sys_icache_invalidate(bufas*mutlibc::c_void,nbytes);Jit{buf,nbytes,bstack:vec![0;BSTACK_MAX*2]}}}// 运行生成的代码。输入必须以 NUL 字节结尾。fnis_match(&mutself,nul_terminated:&[u8])->bool{debug_assert_eq!(nul_terminated.last(),Some(&0));unsafe{letmatcher:MatchFn=std::mem::transmute(self.buf);matcher(nul_terminated.as_ptr(),self.bstack.as_mut_ptr())!=0}}}implDropforJit{fndrop(&mutself){unsafe{libc::munmap(self.bufas*mutlibc::c_void,self.nbytes);}}}结果
完成所有这些后,让我们比较我们构建的不同实现的性能:
| 输入长度 | 解释器 | JIT | 手写 | JIT 加速比 | 手写加速比 |
|---|---|---|---|---|---|
| 9 | 45 纳秒 | 3.8 纳秒 | 3.8 纳秒 | 11.7 倍 | 11.9 倍 |
| 33 | 103 纳秒 | 7.9 纳秒 | 10.5 纳秒 | 13.0 倍 | 9.8 倍 |
| 129 | 597 纳秒 | 30 纳秒 | 32 纳秒 | 19.7 倍 | 18.6 倍 |
| 513 | 1,955 纳秒 | 126 纳秒 | 120 纳秒 | 15.5 倍 | 16.2 倍 |
| 2,049 | 8,301 纳秒 | 470 纳秒 | 393 纳秒 | 17.7 倍 | 21.1 倍 |
所以 JIT 和手写实现基本上并驾齐驱。有时 JIT 版本更快,有时手写版本更快。
网上流传着一个梗,说 AI 没什么用,因为“代码从来都不是困难的部分”。我认为这在某些领域是对的,但在其他领域,编写代码绝对是困难的部分。JIT 编译器就是一个很好的例子。对于许多软件来说,JIT 编译器会大大有助于加速代码。JIT 编译器的稀有性让我相信,从历史上看,实现 JIT 编译器太难了,以至于不值得。LLM 降低了准入门槛,使得编写 JIT 编译器变得容易得多。这就是 pgrust 背后的理念。数据库历来是最难构建的软件,并因此受到限制。现在,有了 AI,我们可以对我们构建的软件类型更加雄心勃勃。
感谢阅读,如果你想支持这个项目,支持 pgrust 的最好方式是在 GitHub 上给我们一个 star。如果你想持续关注:
- GitHub
- Discord
- 邮件列表
- pgrust.com
每周更新 pgrust,包括关于 JIT 编译的后续内容。