code wiki / (root) / nx_f32_lora_linear.nx

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}