code wiki / (root) / nx_f32_maxpool2d_gate.nx

nx_f32_maxpool2d_gate.nx source

↩ module page · 52 lines · 2817 B

1// nx_f32_maxpool2d_gate.nx -- exact proof of nx_f32_maxpool2d on an integer-f32 4x4 map. expect_exit: 0 2import "nx_syscalls.nx" 3import "nx_f32_cvt.nx" 4import "nx_f32_maxpool2d.nx" 5 6func gp(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } return sys_write(1, s, n) } 7func gn(v: i64) -> i64 { 8 let bb: *u8 = sys_mmap(28); var m: i64 = v 9 if m < 0 { sys_write(1, "-" as *u8, 1); m = 0 - m } 10 let t: *u8 = sys_mmap(28); var k: i64 = 0 11 if m == 0 { t[0] = 48 as u8; k = 1 } 12 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 13 var i: i64 = 0; while i < k { bb[i] = t[k - 1 - i]; i = i + 1 } 14 return sys_write(1, bb, k) 15} 16 17func main(argc: i64, argv: *i64) -> i64 { 18 var pass: i64 = 0 19 let inp: *i64 = sys_mmap(8 * 16) as *i64 20 let out: *i64 = sys_mmap(8 * 16) as *i64 21 // 4x4 = 1..16 row-major 22 var i: i64 = 0; while i < 16 { inp[i] = nx_i32_to_f32(i + 1); i = i + 1 } 23 24 // MP1: 2x2 stride 2 pad 0 -> 2x2 = [max(1,2,5,6)=6, max(3,4,7,8)=8, max(9,10,13,14)=14, max(11,12,15,16)=16] 25 var r1: i64 = nx_f32_maxpool2d(inp, 1, 1, 4, 4, 2, 2, 2, 0, out) 26 if out[0] == nx_i32_to_f32(6) { if out[1] == nx_i32_to_f32(8) { if out[2] == nx_i32_to_f32(14) { if out[3] == nx_i32_to_f32(16) { 27 pass = pass + 1; gp("MP1 2x2/s2 -> [6,8,14,16] exact OK\n" as *u8) 28 } } } } 29 if out[0] != nx_i32_to_f32(6) { gp("MP1 FAIL o0=" as *u8); gn(out[0]); gp("\n" as *u8) } 30 31 // MP2: 3x3 stride 1 pad 0 -> 2x2 ; out(0,0)=max top-left 3x3=11 ; out(1,1)=max bottom-right 3x3=16 32 var r2: i64 = nx_f32_maxpool2d(inp, 1, 1, 4, 4, 3, 3, 1, 0, out) 33 if out[0] == nx_i32_to_f32(11) { if out[3] == nx_i32_to_f32(16) { 34 pass = pass + 1; gp("MP2 3x3/s1 -> corner 11, 16 exact OK\n" as *u8) 35 } } 36 if out[0] != nx_i32_to_f32(11) { gp("MP2 FAIL o0=" as *u8); gn(out[0]); gp(" o3=" as *u8); gn(out[3]); gp("\n" as *u8) } 37 38 // MP3: padding excluded -- 2x2 stride 2 pad 1 on 4x4 -> OH=OW=3 ; out(0,0) sees only in(0,0)=1 (rest padding) -> 1 39 var r3: i64 = nx_f32_maxpool2d(inp, 1, 1, 4, 4, 2, 2, 2, 1, out) 40 if out[0] == nx_i32_to_f32(1) { pass = pass + 1; gp("MP3 pad-excluded corner = in(0,0)=1 OK\n" as *u8) } 41 if out[0] != nx_i32_to_f32(1) { gp("MP3 FAIL o0=" as *u8); gn(out[0]); gp("\n" as *u8) } 42 43 // MP4 guard: stride 0 -> ERR 44 var r4: i64 = nx_f32_maxpool2d(inp, 1, 1, 4, 4, 2, 2, 0, 0, out) 45 if r1 == 0 { if r2 == 0 { if r3 == 0 { if r4 == 4 { pass = pass + 1; gp("MP4 rc ok (valid=0, stride0=ERR) OK\n" as *u8) } } } } 46 if r4 != 4 { gp("MP4 FAIL r4=" as *u8); gn(r4); gp("\n" as *u8) } 47 48 gp("MAXPOOL2D-GATE pass=" as *u8); gn(pass); gp("/4\n" as *u8) 49 if pass == 4 { gp("MAXPOOL2D-GATE verdict=GREEN 4/4 (windowed max + stride + pad-excluded + guard)\n" as *u8); sys_exit(0) } 50 sys_exit(1) 51 return 0 52}