code wiki / (root) / nx_linear_attn_bench.nx

nx_linear_attn_bench.nx source

↩ module page · 115 lines · 4195 B

1// nx_linear_attn_bench.nx -- measure SANA linear attention O(n·d^2) vs softmax O(n^2·d) at scale. 2// 3// Confirms the DiT-IC/SANA lever concretely: at n=512 tokens, head_dim=64, linear attention should be 4// materially faster than softmax (which builds the 512x512 score matrix + 512x512 exps). Linear attention 5// copied inline (nx_f32_linear_attention has a main); softmax via the nx_f32_attention lib. 6// license_tier: ORIGINAL 7import "nx_syscalls.nx" 8import "nx_f32.nx" 9import "nx_f32_div.nx" 10import "nx_f32_cvt.nx" 11import "nx_f32_attention.nx" 12import "nx_clock.nx" 13import "nx_strconv.nx" 14 15func lab_phi(x: i64) -> i64 { 16 var r: i64 = x 17 if (x & 0x80000000) != 0 { r = 0 } 18 return nx_f32_add(r, nx_i32_to_f32(1)) 19} 20 21func lab_linear(Q: *i64, K: *i64, V: *i64, n_tokens: i64, head_dim: i64, out: *i64, S: *i64, z: *i64) -> i64 { 22 var a: i64 = 0 23 while a < head_dim * head_dim { S[a] = 0; a = a + 1 } 24 a = 0 25 while a < head_dim { z[a] = 0; a = a + 1 } 26 var j: i64 = 0 27 while j < n_tokens { 28 a = 0 29 while a < head_dim { 30 let pk: i64 = lab_phi(K[j * head_dim + a]) 31 z[a] = nx_f32_add(z[a], pk) 32 var b: i64 = 0 33 while b < head_dim { S[a * head_dim + b] = nx_f32_add(S[a * head_dim + b], nx_f32_mul(pk, V[j * head_dim + b])); b = b + 1 } 34 a = a + 1 35 } 36 j = j + 1 37 } 38 var i: i64 = 0 39 while i < n_tokens { 40 var den: i64 = 0 41 a = 0 42 while a < head_dim { den = nx_f32_add(den, nx_f32_mul(lab_phi(Q[i * head_dim + a]), z[a])); a = a + 1 } 43 var b: i64 = 0 44 while b < head_dim { 45 var num: i64 = 0 46 a = 0 47 while a < head_dim { num = nx_f32_add(num, nx_f32_mul(lab_phi(Q[i * head_dim + a]), S[a * head_dim + b])); a = a + 1 } 48 out[i * head_dim + b] = nx_f32_div(num, den) 49 b = b + 1 50 } 51 i = i + 1 52 } 53 return 0 54} 55 56func lab_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 { 57 let line: *u8 = sys_mmap(64) 58 var lo: i64 = 0 59 var i: i64 = 0 60 while i < kl { line[lo] = key[i]; lo = lo + 1; i = i + 1 } 61 line[lo] = 0x3D; lo = lo + 1 62 let dec: *u8 = sys_mmap(32) 63 let nd: i64 = nx_strconv_format_i64(v, dec) 64 var k: i64 = 0 65 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 } 66 line[lo] = 0x0A; lo = lo + 1 67 return sys_write(fd, line, lo) 68} 69 70func main() -> i64 { 71 let N: i64 = 512 72 let D: i64 = 64 73 let Q: *i64 = sys_mmap(N * D * 8) as *i64 74 let K: *i64 = sys_mmap(N * D * 8) as *i64 75 let V: *i64 = sys_mmap(N * D * 8) as *i64 76 let out: *i64 = sys_mmap(N * D * 8) as *i64 77 let seven: i64 = nx_i32_to_f32(7) 78 var i: i64 = 0 79 while i < N * D { 80 Q[i] = nx_f32_div(nx_i32_to_f32((i - (i / 7) * 7) + 1), seven) 81 K[i] = nx_f32_div(nx_i32_to_f32((i - (i / 11) * 11) + 1), nx_i32_to_f32(11)) 82 V[i] = nx_f32_div(nx_i32_to_f32((i - (i / 13) * 13) + 1), nx_i32_to_f32(13)) 83 i = i + 1 84 } 85 let S: *i64 = sys_mmap(D * D * 8) as *i64 86 let z: *i64 = sys_mmap(D * 8) as *i64 87 let sr: *i64 = sys_mmap(N * 8) as *i64 88 let pr: *i64 = sys_mmap(N * 8) as *i64 89 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(8)) // 1/sqrt(64) 90 91 let IT: i64 = 5 92 let t0: i64 = nx_clock_monotonic_ns() 93 var k: i64 = 0 94 while k < IT { lab_linear(Q, K, V, N, D, out, S, z); k = k + 1 } 95 let t1: i64 = nx_clock_monotonic_ns() 96 let lin_ns: i64 = t1 - t0 97 98 let t2: i64 = nx_clock_monotonic_ns() 99 k = 0 100 while k < IT { nx_f32_attention(Q, K, V, N, D, scale, out, sr, pr); k = k + 1 } 101 let t3: i64 = nx_clock_monotonic_ns() 102 let sm_ns: i64 = t3 - t2 103 104 let ofd: i64 = sys_openat_wr("/tmp/linattn_bench.txt" as *u8, 0x1a4) 105 if ofd >= 0 { 106 lab_emit(ofd, "n_tokens" as *u8, 8, N) 107 lab_emit(ofd, "head_dim" as *u8, 8, D) 108 lab_emit(ofd, "linear_ns" as *u8, 9, lin_ns) 109 lab_emit(ofd, "softmax_ns" as *u8, 10, sm_ns) 110 if lin_ns > 0 { lab_emit(ofd, "softmax_over_linear_x100" as *u8, 24, sm_ns * 100 / lin_ns) } 111 sys_close(ofd) 112 } 113 if lin_ns >= sm_ns { return 80 } // linear must be faster at this scale 114 return 0 115}