# [allow (non_snake_case , unused_imports , dead_code)] mod snapshot_reduce { fn build_module (sm : & str) -> kaio :: core :: ir :: PtxModule { use kaio :: core :: emit :: { Emit , PtxWriter } ; use kaio :: core :: instr :: ArithOp ; use kaio :: core :: instr :: control :: { CmpOp , ControlOp } ; use kaio :: core :: instr :: memory :: MemoryOp ; use kaio :: core :: instr :: special ; use kaio :: core :: ir :: { Operand , PtxInstruction , PtxKernel , PtxModule , PtxParam , RegisterAllocator , SharedDecl , } ; use kaio :: core :: types :: PtxType ; let _kaio_annotate = std :: env :: var ("KAIO_PTX_ANNOTATE") . is_ok () ; let mut alloc = RegisterAllocator :: new () ; let mut kernel = PtxKernel :: new ("snapshot_reduce") ; kernel . add_param (PtxParam :: pointer ("out_ptr" , PtxType :: F32)) ; let _kaio_r0 = alloc . alloc (PtxType :: U64) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: LdParam { dst : _kaio_r0 , param_name : "out_ptr" . to_string () , ty : PtxType :: U64 , })) ; let _kaio_r1 = alloc . alloc (PtxType :: U64) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: CvtaToGlobal { dst : _kaio_r1 , src : _kaio_r0 , })) ; kernel . add_param (PtxParam :: scalar ("n" , PtxType :: U32)) ; let _kaio_r2 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: LdParam { dst : _kaio_r2 , param_name : "n" . to_string () , ty : PtxType :: U32 , })) ; if _kaio_annotate { kernel . push (PtxInstruction :: Comment ("let x" . to_string ())) ; } let _kaio_r3 = alloc . alloc (PtxType :: F32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r3 , src : Operand :: ImmF32 (1f32) , ty : PtxType :: F32 , }) ; if _kaio_annotate { kernel . push (PtxInstruction :: Comment ("let s" . to_string ())) ; } kernel . add_shared_decl (SharedDecl { name : "_kaio_reduce_smem" . to_string () , align : 4 , size_bytes : 32u32 , }) ; let _kaio_r4 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r4 , src : Operand :: SharedAddr ("_kaio_reduce_smem" . to_string ()) , ty : PtxType :: U32 , }) ; let _kaio_r5 = alloc . alloc (PtxType :: F32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r5 , src : Operand :: Reg (_kaio_r3) , ty : PtxType :: F32 , }) ; let _kaio_r6 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r6 , src : Operand :: SpecialReg (kaio :: core :: ir :: SpecialReg :: TidX) , ty : PtxType :: U32 , }) ; let _kaio_r7 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r7 , src : _kaio_r5 , delta : Operand :: ImmU32 (16) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r5 , lhs : Operand :: Reg (_kaio_r5) , rhs : Operand :: Reg (_kaio_r7) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r7 , src : _kaio_r5 , delta : Operand :: ImmU32 (8) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r5 , lhs : Operand :: Reg (_kaio_r5) , rhs : Operand :: Reg (_kaio_r7) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r7 , src : _kaio_r5 , delta : Operand :: ImmU32 (4) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r5 , lhs : Operand :: Reg (_kaio_r5) , rhs : Operand :: Reg (_kaio_r7) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r7 , src : _kaio_r5 , delta : Operand :: ImmU32 (2) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r5 , lhs : Operand :: Reg (_kaio_r5) , rhs : Operand :: Reg (_kaio_r7) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r7 , src : _kaio_r5 , delta : Operand :: ImmU32 (1) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r5 , lhs : Operand :: Reg (_kaio_r5) , rhs : Operand :: Reg (_kaio_r7) , ty : PtxType :: F32 , })) ; let _kaio_r8 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Div { dst : _kaio_r8 , lhs : Operand :: Reg (_kaio_r6) , rhs : Operand :: ImmU32 (32) , ty : PtxType :: U32 , })) ; let _kaio_r9 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Mul { dst : _kaio_r9 , lhs : Operand :: Reg (_kaio_r8) , rhs : Operand :: ImmU32 (32) , ty : PtxType :: U32 , })) ; let _kaio_r10 = alloc . alloc (PtxType :: Pred) ; kernel . push (PtxInstruction :: Control (ControlOp :: SetP { dst : _kaio_r10 , cmp_op : CmpOp :: Eq , lhs : Operand :: Reg (_kaio_r9) , rhs : Operand :: Reg (_kaio_r6) , ty : PtxType :: U32 , })) ; let _kaio_r11 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Mul { dst : _kaio_r11 , lhs : Operand :: Reg (_kaio_r8) , rhs : Operand :: ImmU32 (4) , ty : PtxType :: U32 , })) ; let _kaio_r12 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r12 , lhs : Operand :: Reg (_kaio_r4) , rhs : Operand :: Reg (_kaio_r11) , ty : PtxType :: U32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: BraPred { pred : _kaio_r10 , target : "REDUCE_WRITE_DONE_0" . to_string () , negate : true , })) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: StShared { addr : _kaio_r12 , src : _kaio_r5 , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Label ("REDUCE_WRITE_DONE_0" . to_string ())) ; kernel . push (PtxInstruction :: Control (ControlOp :: BarSync { barrier_id : 0 })) ; let _kaio_r13 = alloc . alloc (PtxType :: Pred) ; kernel . push (PtxInstruction :: Control (ControlOp :: SetP { dst : _kaio_r13 , cmp_op : CmpOp :: Lt , lhs : Operand :: Reg (_kaio_r6) , rhs : Operand :: ImmU32 (32) , ty : PtxType :: U32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: BraPred { pred : _kaio_r13 , target : "REDUCE_BROADCAST_1" . to_string () , negate : true , })) ; let _kaio_r14 = alloc . alloc (PtxType :: Pred) ; kernel . push (PtxInstruction :: Control (ControlOp :: SetP { dst : _kaio_r14 , cmp_op : CmpOp :: Lt , lhs : Operand :: Reg (_kaio_r6) , rhs : Operand :: ImmU32 (8u32) , ty : PtxType :: U32 , })) ; let _kaio_r17 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Mul { dst : _kaio_r17 , lhs : Operand :: Reg (_kaio_r6) , rhs : Operand :: ImmU32 (4) , ty : PtxType :: U32 , })) ; let _kaio_r18 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r18 , lhs : Operand :: Reg (_kaio_r4) , rhs : Operand :: Reg (_kaio_r17) , ty : PtxType :: U32 , })) ; let _kaio_r15 = alloc . alloc (PtxType :: F32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r15 , src : Operand :: ImmF32 (0.0f32) , ty : PtxType :: F32 , }) ; kernel . push (PtxInstruction :: Control (ControlOp :: BraPred { pred : _kaio_r14 , target : "REDUCE_LOAD_DONE_2" . to_string () , negate : true , })) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: LdShared { dst : _kaio_r15 , addr : _kaio_r18 , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Label ("REDUCE_LOAD_DONE_2" . to_string ())) ; let _kaio_r16 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r16 , src : _kaio_r15 , delta : Operand :: ImmU32 (16) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r15 , lhs : Operand :: Reg (_kaio_r15) , rhs : Operand :: Reg (_kaio_r16) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r16 , src : _kaio_r15 , delta : Operand :: ImmU32 (8) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r15 , lhs : Operand :: Reg (_kaio_r15) , rhs : Operand :: Reg (_kaio_r16) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r16 , src : _kaio_r15 , delta : Operand :: ImmU32 (4) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r15 , lhs : Operand :: Reg (_kaio_r15) , rhs : Operand :: Reg (_kaio_r16) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r16 , src : _kaio_r15 , delta : Operand :: ImmU32 (2) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r15 , lhs : Operand :: Reg (_kaio_r15) , rhs : Operand :: Reg (_kaio_r16) , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: ShflSyncDown { dst : _kaio_r16 , src : _kaio_r15 , delta : Operand :: ImmU32 (1) , c : 31 , mask : 0xFFFFFFFF_u32 , })) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r15 , lhs : Operand :: Reg (_kaio_r15) , rhs : Operand :: Reg (_kaio_r16) , ty : PtxType :: F32 , })) ; let _kaio_r19 = alloc . alloc (PtxType :: Pred) ; kernel . push (PtxInstruction :: Control (ControlOp :: SetP { dst : _kaio_r19 , cmp_op : CmpOp :: Eq , lhs : Operand :: Reg (_kaio_r6) , rhs : Operand :: ImmU32 (0) , ty : PtxType :: U32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: BraPred { pred : _kaio_r19 , target : "REDUCE_T0_DONE_3" . to_string () , negate : true , })) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: StShared { addr : _kaio_r18 , src : _kaio_r15 , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Label ("REDUCE_T0_DONE_3" . to_string ())) ; kernel . push (PtxInstruction :: Label ("REDUCE_BROADCAST_1" . to_string ())) ; kernel . push (PtxInstruction :: Control (ControlOp :: BarSync { barrier_id : 0 })) ; let _kaio_r20 = alloc . alloc (PtxType :: U32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r20 , src : Operand :: Reg (_kaio_r4) , ty : PtxType :: U32 , }) ; let _kaio_r21 = alloc . alloc (PtxType :: F32) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: LdShared { dst : _kaio_r21 , addr : _kaio_r20 , ty : PtxType :: F32 , })) ; if _kaio_annotate { kernel . push (PtxInstruction :: Comment ("out[...] = ..." . to_string ())) ; } let _kaio_r22 = alloc . alloc (PtxType :: S32) ; kernel . push (PtxInstruction :: Mov { dst : _kaio_r22 , src : Operand :: ImmI32 (0i32) , ty : PtxType :: S32 , }) ; let _kaio_r23 = alloc . alloc (PtxType :: U64) ; kernel . push (PtxInstruction :: Arith (ArithOp :: MulWide { dst : _kaio_r23 , lhs : Operand :: Reg (_kaio_r22) , rhs : Operand :: ImmU32 (4u32) , src_ty : PtxType :: U32 , })) ; let _kaio_r24 = alloc . alloc (PtxType :: S64) ; kernel . push (PtxInstruction :: Arith (ArithOp :: Add { dst : _kaio_r24 , lhs : Operand :: Reg (_kaio_r1) , rhs : Operand :: Reg (_kaio_r23) , ty : PtxType :: S64 , })) ; kernel . push (PtxInstruction :: Memory (MemoryOp :: StGlobal { addr : _kaio_r24 , src : _kaio_r21 , ty : PtxType :: F32 , })) ; kernel . push (PtxInstruction :: Control (ControlOp :: Ret)) ; kernel . set_registers (alloc . into_allocated ()) ; if std :: env :: var ("KAIO_PTX_STATS") . is_ok () { let _s = kernel . stats () ; eprintln ! ("KAIO stats: kernel '{}' (PTX structure, not runtime profile)" , "snapshot_reduce") ; eprintln ! (" Instructions: {} total" , _s . total_instructions) ; eprintln ! (" Arithmetic: {} fma, {} other" , _s . fma , _s . arith_other) ; eprintln ! (" Memory: {} ld.global, {} st.global, {} ld.shared, {} st.shared" , _s . ld_global , _s . st_global , _s . ld_shared , _s . st_shared) ; eprintln ! (" Control: {} bar.sync, {} branches, {} setp, {} mov, {} cvt" , _s . bar_sync , _s . branches , _s . setp , _s . mov , _s . cvt) ; eprintln ! (" Registers: {} r32, {} r64, {} f32, {} f64, {} pred, {} f16, {} bf16 (PTX-level, not final HW allocation)" , _s . registers_r , _s . registers_rd , _s . registers_f , _s . registers_fd , _s . registers_p , _s . registers_h , _s . registers_hb) ; eprintln ! (" Shared mem: {} bytes" , _s . shared_bytes) ; } let sm_target = std :: env :: var ("KAIO_SM_TARGET") . unwrap_or_else (| _ | sm . to_string ()) ; let mut module = PtxModule :: new (& sm_target) ; module . add_kernel (kernel) ; if std :: env :: var ("KAIO_DUMP_PTX") . is_ok () { let mut w = PtxWriter :: new () ; module . emit (& mut w) . unwrap () ; let ptx = w . finish () ; let dump_dir = std :: env :: var ("OUT_DIR") . unwrap_or_else (| _ | "." . to_string ()) ; let dump_path = format ! ("{}/{}.ptx" , dump_dir , "snapshot_reduce") ; match std :: fs :: write (& dump_path , & ptx) { Ok (()) => eprintln ! ("KAIO: wrote {}" , dump_path) , Err (e) => eprintln ! ("KAIO: failed to write {}: {}" , dump_path , e) , } } module } # [doc = r" Launch this GPU kernel on the given device."] pub fn launch (device : & kaio :: runtime :: KaioDevice , out : & mut kaio :: runtime :: GpuBuffer < f32 > , n : u32) -> Result < () , kaio :: runtime :: KaioError > { use kaio :: runtime :: PushKernelArg ; let info = device . info () ? ; let (major , minor) = info . compute_capability ; let sm = format ! ("sm_{major}{minor}") ; let ptx_module = build_module (& sm) ; let module = device . load_module (& ptx_module) ? ; let func = module . function ("snapshot_reduce") ? ; let cfg = kaio :: runtime :: LaunchConfig { grid_dim : (n . div_ceil (256u32) , 1 , 1) , block_dim : (256u32 , 1 , 1) , shared_mem_bytes : 0 , } ; unsafe { device . stream () . launch_builder (func . inner ()) . arg (out . inner_mut ()) . arg (& n) . launch (cfg) ? ; } Ok (()) } }