nx_f32_lora_linear.nx source
↩ module page · 98 lines · 3801 B
1// nx_f32_lora_linear.nx -- sovereign LoRA-adapted linear (the Aligned-LoRA / fine-tune-fewer-images lever).
2//
3// From DiT-IC's "Aligned LoRA Adaptation" (ingested 2026-07-01): fine-tune a frozen weight W [out,in] with a
4// low-rank adapter B [out,r], A [r,in] (r << in,out). Forward: out = W x + B (A x). Only r*(in+out) params
5// train instead of in*out -> huge param/VRAM/data saving (the "fewer images, faster training" lever), and
6// LoRA merges into W for zero-overhead inference. Gate: LoRA forward == (W + B@A) x within f32 tolerance.
7// license_tier: ORIGINAL
8import "nx_syscalls.nx"
9import "nx_f32.nx"
10import "nx_f32_div.nx"
11import "nx_f32_cvt.nx"
12const K_MAGIC_100000: i64 = 100000
13
14// out[out_dim] = W x + B (A x). W:[out,in], B:[out,r], A:[r,in], x:[in].
15func nx_f32_lora_forward(W: *i64, B: *i64, A: *i64, x: *i64, in_dim: i64, out_dim: i64, rank: i64,
16 out: *i64, ax: *i64) -> i64 {
17 // Ax[r] = A @ x
18 var k: i64 = 0
19 while k < rank {
20 var acc: i64 = 0
21 var i: i64 = 0
22 while i < in_dim { acc = nx_f32_add(acc, nx_f32_mul(A[k * in_dim + i], x[i])); i = i + 1 }
23 ax[k] = acc
24 k = k + 1
25 }
26 // out = W x + B (Ax)
27 var o: i64 = 0
28 while o < out_dim {
29 var wx: i64 = 0
30 var i: i64 = 0
31 while i < in_dim { wx = nx_f32_add(wx, nx_f32_mul(W[o * in_dim + i], x[i])); i = i + 1 }
32 var bax: i64 = 0
33 k = 0
34 while k < rank { bax = nx_f32_add(bax, nx_f32_mul(B[o * rank + k], ax[k])); k = k + 1 }
35 out[o] = nx_f32_add(wx, bax)
36 o = o + 1
37 }
38 return 0
39}
40
41func lo_close(x: i64, e: i64, tol: i64) -> i64 {
42 let ax: i64 = e & 0x7FFFFFFF
43 var thr: i64 = tol
44 if (nx_f32_mul(tol, ax) & 0x7FFFFFFF) > tol { thr = nx_f32_mul(tol, ax) } // relative-or-abs
45 if (nx_f32_sub(x, e) & 0x7FFFFFFF) < thr { return 1 }
46 return 0
47}
48
49func main() -> i64 {
50 let IN: i64 = 4
51 let OUT: i64 = 4
52 let R: i64 = 2
53 let W: *i64 = sys_mmap(OUT * IN * 8) as *i64
54 let B: *i64 = sys_mmap(OUT * R * 8) as *i64
55 let A: *i64 = sys_mmap(R * IN * 8) as *i64
56 let x: *i64 = sys_mmap(IN * 8) as *i64
57 let out: *i64 = sys_mmap(OUT * 8) as *i64
58 let ax: *i64 = sys_mmap(R * 8) as *i64
59 let ten: i64 = nx_i32_to_f32(10)
60
61 var i: i64 = 0
62 while i < OUT * IN { W[i] = nx_f32_div(nx_i32_to_f32((i - (i / 5) * 5) + 1), ten); i = i + 1 }
63 i = 0
64 while i < OUT * R { B[i] = nx_f32_div(nx_i32_to_f32((i - (i / 3) * 3) + 1), ten); i = i + 1 }
65 i = 0
66 while i < R * IN { A[i] = nx_f32_div(nx_i32_to_f32((i - (i / 7) * 7) + 1), ten); i = i + 1 }
67 i = 0
68 while i < IN { x[i] = nx_f32_div(nx_i32_to_f32((i - (i / 4) * 4) + 1), ten); i = i + 1 }
69
70 nx_f32_lora_forward(W, B, A, x, IN, OUT, R, out, ax)
71
72 // reference: full merged W' = W + B@A, then W' x
73 let tol: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000)) // 0.001
74 var o: i64 = 0
75 while o < OUT {
76 var ref: i64 = 0
77 i = 0
78 while i < IN {
79 // W'[o][i] = W[o][i] + Σ_k B[o][k]*A[k][i]
80 var ba: i64 = 0
81 var k: i64 = 0
82 while k < R { ba = nx_f32_add(ba, nx_f32_mul(B[o * R + k], A[k * IN + i])); k = k + 1 }
83 let wp: i64 = nx_f32_add(W[o * IN + i], ba)
84 ref = nx_f32_add(ref, nx_f32_mul(wp, x[i]))
85 i = i + 1
86 }
87 if lo_close(out[o], ref, tol) != 1 { return 10 + o }
88 o = o + 1
89 }
90
91 // adapter must actually contribute (out != W x alone) -> proves B(Ax) is applied
92 var wxonly: i64 = 0
93 i = 0
94 while i < IN { wxonly = nx_f32_add(wxonly, nx_f32_mul(W[0 * IN + i], x[i])); i = i + 1 }
95 if (nx_f32_sub(out[0], wxonly) & 0x7FFFFFFF) < nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_100000)) { return 20 }
96
97 return 0
98}