code wiki / _hdl_build / nx_sass_sched.nx
nx_sass_sched.nx source
↩ module page · 41 lines · 2114 B
1// nx_sass_sched.nx -- conservative SASS instruction SCHEDULER (the scheduling-POLICY layer toward a runnable
2// kernel; pairs with the nx_mma_asm encoder + nv_control). Given an instruction sequence (per-instr: written
3// register, read registers, variable-latency flag), it assigns a scoreboard BARRIER to each variable-latency
4// write (loads/MMA) and computes each instruction's WAIT mask so EVERY data dependency (RAW) is covered.
5// CONSERVATIVE = provably CORRECT (serialize on every dependency), not yet optimal -- the "make it run" pass;
6// minimizing stalls/barriers = a later optimization pass. Emits the control word per instr via nv_control.
7// Pure funcs, no main. license_tier: ORIGINAL
8import "nx_syscalls.nx"
9import "nx_mma_asm.nx"
10
11// assign a write-barrier (0..5, round-robin) to each variable-latency instruction; -1 for fixed-latency.
12func sched_assign_barriers(varlat: *i64, bar: *i64, n: i64) -> i64 {
13 var b: i64 = 0
14 var i: i64 = 0
15 while i < n { if varlat[i] == 1 { bar[i] = b; b = (b + 1) % 6 } else { bar[i] = 0 - 1 } i = i + 1 }
16 return 0
17}
18// wait mask for instruction j = OR over prior var-latency instrs i<j whose written reg is READ by j (RAW dep).
19func sched_wait_mask(wreg: *i64, r0: *i64, r1: *i64, r2: *i64, varlat: *i64, bar: *i64, j: i64) -> i64 {
20 var mask: i64 = 0
21 var i: i64 = 0
22 while i < j {
23 if varlat[i] == 1 {
24 let w: i64 = wreg[i]
25 if w >= 0 {
26 if w == r0[j] { mask = mask | (1 << bar[i]) }
27 if w == r1[j] { mask = mask | (1 << bar[i]) }
28 if w == r2[j] { mask = mask | (1 << bar[i]) }
29 }
30 }
31 i = i + 1
32 }
33 return mask
34}
35// the full control word for instruction j: write-barrier = its own bar (if var-lat, else none=7), wait = its dep mask.
36func sched_control(wreg: *i64, r0: *i64, r1: *i64, r2: *i64, varlat: *i64, bar: *i64, j: i64, stall: i64) -> i64 {
37 var wbar: i64 = 7
38 if varlat[j] == 1 { wbar = bar[j] }
39 let wait: i64 = sched_wait_mask(wreg, r0, r1, r2, varlat, bar, j)
40 return nv_control(stall, 0, wbar, 7, wait, 0)
41}