code wiki / (root) / nx_f32_rectflow_denoise.nx

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}