nx_f32_rectflow_denoise.nx source
↩ module page · 111 lines · 5027 B
1// nx_f32_rectflow_denoise.nx -- software-f32 rectified-flow (flow-matching) Euler denoise loop.
2//
3// sd-server -> Nishi migration: this REPLACES Z-Image Turbo's exact sampler -- FLOW_PRED (rectified flow)
4// + euler method + discrete scheduler. Z-Image uses it at 6 steps. The loop runs the DiT block as the
5// velocity predictor each step and integrates along the schedule:
6//
7// for s in 0..n_steps:
8// v = DiT_block(x ; weights, conditioning_s) # predicted velocity (flow)
9// dt = sigma[s+1] - sigma[s] # schedule step (rectified-flow: linear in sigma)
10// x = x + dt * v # Euler update
11//
12// Composes the gated `nx_f32_dit_block_tiny` + f32 add/sub/mul. x: flat *i64 f32 bits [n_tokens, D]
13// (mutated in place: noise -> clean latent). sigmas: [n_steps+1] f32 (the noise schedule). Weights are the
14// DiT block's (shared across steps here; real inference varies the adaLN conditioning per timestep).
15// license_tier: ORIGINAL
16import "nx_syscalls.nx"
17import "nx_f32.nx"
18import "nx_f32_div.nx"
19import "nx_f32_cvt.nx"
20import "nx_f32_dit_block_tiny.nx"
21const NX_MAGIC_100000: i64 = 100000
22
23const NX_F32RF_OK: i64 = 0
24const NX_F32RF_ERR: i64 = 1
25
26func nx_f32_rectflow_denoise(x: *i64, n_tokens: i64, D: i64, d_ff: i64, gamma: *i64, eps: i64,
27 Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64,
28 W1: *i64, W3: *i64, W2: *i64,
29 sc: *i64, sh: *i64, gt: *i64,
30 sigmas: *i64, n_steps: i64) -> i64 {
31 if n_steps < 0 { return NX_F32RF_ERR }
32 let nel: i64 = n_tokens * D
33 let v: *i64 = sys_mmap(nel * 8) as *i64
34 var s: i64 = 0
35 while s < n_steps {
36 // v = DiT_block(x) -- the predicted flow/velocity (adaLN params reused for both sublayers here)
37 nx_f32_dit_block_tiny(x, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, sc, sh, gt, W1, W3, W2, sc, sh, gt, v)
38 // dt = sigma[s+1] - sigma[s]
39 let dt: i64 = nx_f32_sub(sigmas[s + 1], sigmas[s])
40 // x = x + dt * v (Euler step)
41 var i: i64 = 0
42 while i < nel { x[i] = nx_f32_add(x[i], nx_f32_mul(dt, v[i])); i = i + 1 }
43 s = s + 1
44 }
45 return NX_F32RF_OK
46}
47
48// ===== Self-test (inline gate) ====================================
49// (a) constant schedule (all sigmas equal) -> dt == 0 every step -> x UNCHANGED (bit-exact), any weights.
50// (b) varying schedule -> x changes AND is deterministic (two identical runs -> bit-identical).
51func main() -> i64 {
52 let n_tokens: i64 = 2
53 let D: i64 = 2
54 let d_ff: i64 = 2
55 let nel: i64 = n_tokens * D
56 let x0: *i64 = sys_mmap(nel * 8) as *i64 // pristine copy
57 let xa: *i64 = sys_mmap(nel * 8) as *i64
58 let xb: *i64 = sys_mmap(nel * 8) as *i64
59 let gamma: *i64 = sys_mmap(D * 8) as *i64
60 let Wq: *i64 = sys_mmap(D * D * 8) as *i64
61 let Wk: *i64 = sys_mmap(D * D * 8) as *i64
62 let Wv: *i64 = sys_mmap(D * D * 8) as *i64
63 let Wo: *i64 = sys_mmap(D * D * 8) as *i64
64 let W1: *i64 = sys_mmap(d_ff * D * 8) as *i64
65 let W3: *i64 = sys_mmap(d_ff * D * 8) as *i64
66 let W2: *i64 = sys_mmap(D * d_ff * 8) as *i64
67 let sc: *i64 = sys_mmap(D * 8) as *i64
68 let sh: *i64 = sys_mmap(D * 8) as *i64
69 let gt: *i64 = sys_mmap(D * 8) as *i64
70 let sig: *i64 = sys_mmap(8 * 8) as *i64
71 let one: i64 = nx_i32_to_f32(1)
72 let p1: i64 = nx_f32_div(one, nx_i32_to_f32(10))
73 let eps: i64 = nx_f32_div(one, nx_i32_to_f32(NX_MAGIC_100000))
74
75 var i: i64 = 0
76 while i < D * D { Wq[i] = p1; Wk[i] = p1; Wv[i] = p1; Wo[i] = p1; i = i + 1 }
77 i = 0
78 while i < d_ff * D { W1[i] = p1; W3[i] = p1; i = i + 1 }
79 i = 0
80 while i < D * d_ff { W2[i] = p1; i = i + 1 }
81 i = 0
82 while i < D { gamma[i] = one; sc[i] = 0; sh[i] = 0; gt[i] = one; i = i + 1 } // gate=1 -> block contributes
83 i = 0
84 while i < nel { x0[i] = nx_i32_to_f32(i + 1); i = i + 1 }
85
86 // (a) constant schedule -> unchanged
87 i = 0
88 while i < nel { xa[i] = x0[i]; i = i + 1 }
89 sig[0] = one; sig[1] = one; sig[2] = one // dt = 0 each step
90 let va: i64 = nx_f32_rectflow_denoise(xa, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, W1, W3, W2, sc, sh, gt, sig, 2)
91 if va != NX_F32RF_OK { return 10 }
92 i = 0
93 while i < nel { if xa[i] != x0[i] { return 20 } i = i + 1 }
94
95 // (b) varying schedule: sigmas [0, 1, 2] -> deterministic + changed
96 sig[0] = 0; sig[1] = one; sig[2] = nx_i32_to_f32(2)
97 i = 0
98 while i < nel { xa[i] = x0[i]; xb[i] = x0[i]; i = i + 1 }
99 nx_f32_rectflow_denoise(xa, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, W1, W3, W2, sc, sh, gt, sig, 2)
100 nx_f32_rectflow_denoise(xb, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, W1, W3, W2, sc, sh, gt, sig, 2)
101 var changed: i64 = 0
102 i = 0
103 while i < nel {
104 if xa[i] != xb[i] { return 30 } // deterministic
105 if xa[i] != x0[i] { changed = 1 }
106 i = i + 1
107 }
108 if changed == 0 { return 40 }
109
110 return 0
111}