@@ -17,7 +17,7 @@ use std::sync::Arc;
1717
1818use crate :: Result ;
1919use crate :: device:: { DeviceReservation , alloc_or_trim, alloc_zeros_or_trim, backend, htod_or_trim} ;
20- use crate :: sumcheck:: { DeviceFactors , eq_table_ext3, launch_shape} ;
20+ use crate :: sumcheck:: { FactorView , eq_table_ext3, launch_shape} ;
2121
2222const BLOCK_DIM : u32 = 256 ;
2323const MAX_GRID : u32 = 4096 ;
@@ -50,7 +50,7 @@ pub struct FusedBus<'a> {
5050/// One table's fused zerocheck on the card.
5151pub struct FusedZerocheck < ' f > {
5252 stream : Arc < CudaStream > ,
53- source : & ' f DeviceFactors ,
53+ source : FactorView < ' f > ,
5454 width : usize ,
5555 /// The cube the lifted factors span.
5656 rows : usize ,
@@ -151,16 +151,16 @@ impl<'f> FusedZerocheck<'f> {
151151 /// `r_tail` and `rho_tail` are `r[2..]` and `ρ[2..]` (three u64 a
152152 /// coordinate): the grid pass's weights.
153153 pub fn new (
154- source : & ' f DeviceFactors ,
154+ source : FactorView < ' f > ,
155155 program : FusedProgram < ' _ > ,
156156 bus : FusedBus < ' _ > ,
157157 r_tail : & [ u64 ] ,
158158 rho_tail : & [ u64 ] ,
159159 grid_rows : usize ,
160160 gruen_rows : usize ,
161161 ) -> Result < Option < Self > > {
162- let rows = source. len ( ) ;
163- let width = source. width ( ) ;
162+ let rows = source. rows ;
163+ let width = source. width ;
164164 assert ! (
165165 rows >= 8 && rows. is_power_of_two( ) ,
166166 "the grid needs four rows a group"
@@ -194,7 +194,7 @@ impl<'f> FusedZerocheck<'f> {
194194 ) ) else {
195195 return Ok ( None ) ;
196196 } ;
197- let stream = source. stream ( ) . clone ( ) ;
197+ let stream = source. stream . clone ( ) ;
198198 let q = rows / 4 ;
199199 let nonempty = |v : & [ u64 ] | -> Vec < u64 > { if v. is_empty ( ) { vec ! [ 0 ] } else { v. to_vec ( ) } } ;
200200 let nodes = htod_or_trim ( & stream, & nonempty ( program. nodes ) ) ?;
@@ -316,7 +316,7 @@ impl<'f> FusedZerocheck<'f> {
316316 q. div_ceil ( BLOCK_DIM as u64 ) . clamp ( 1 , MAX_GRID as u64 ) as u32 ,
317317 BLOCK_DIM ,
318318 ) ;
319- let factors = self . source . factor_ptrs ( ) ;
319+ let factors = self . source . ptrs ;
320320 unsafe {
321321 self . stream
322322 . launch_builder ( & be. zc_bus_u )
@@ -328,6 +328,7 @@ impl<'f> FusedZerocheck<'f> {
328328 . arg ( & self . constant )
329329 . arg ( & self . eq_rho )
330330 . arg ( & mut self . partials )
331+ . arg ( & self . source . stride )
331332 . launch ( LaunchConfig {
332333 grid_dim : ( grid, 1 , 1 ) ,
333334 block_dim : ( block, 1 , 1 ) ,
@@ -357,7 +358,7 @@ impl<'f> FusedZerocheck<'f> {
357358 }
358359 self . stream . memcpy_htod ( & [ u64:: MAX ] , & mut self . violation ) ?;
359360 let ( grid, block) = self . shape ( q, points. len ( ) , self . num_slots as u64 ) ?;
360- let factors = self . source . factor_ptrs ( ) ;
361+ let factors = self . source . ptrs ;
361362 unsafe {
362363 self . stream
363364 . launch_builder ( & be. zc_grid01 )
@@ -373,6 +374,7 @@ impl<'f> FusedZerocheck<'f> {
373374 . arg ( & mut self . partials )
374375 . arg ( & mut self . violation )
375376 . arg ( & u32:: from ( keep_corners) )
377+ . arg ( & self . source . stride )
376378 . launch ( LaunchConfig {
377379 grid_dim : ( grid, points. len ( ) as u32 , 1 ) ,
378380 block_dim : ( block, 1 , 1 ) ,
@@ -396,7 +398,7 @@ impl<'f> FusedZerocheck<'f> {
396398 let width = self . width as u64 ;
397399 let total = width * q;
398400 let grid = total. div_ceil ( BLOCK_DIM as u64 ) . clamp ( 1 , MAX_GRID as u64 ) as u32 ;
399- let factors = self . source . factor_ptrs ( ) ;
401+ let factors = self . source . ptrs ;
400402 unsafe {
401403 self . stream
402404 . launch_builder ( & be. zc_fold2 )
@@ -405,6 +407,7 @@ impl<'f> FusedZerocheck<'f> {
405407 . arg ( & width)
406408 . arg ( & self . scratch )
407409 . arg ( & mut self . folded )
410+ . arg ( & self . source . stride )
408411 . launch ( LaunchConfig {
409412 grid_dim : ( grid, 1 , 1 ) ,
410413 block_dim : ( BLOCK_DIM , 1 , 1 ) ,
0 commit comments