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}