@@ -9,8 +9,7 @@ use field::PrimeCharacteristicRing;
99use field:: integers:: QuotientMap ;
1010use field:: { ExtensionField , PrimeField64 } ;
1111use koala_bear:: symmetric:: Permutation ;
12- use rayon:: prelude:: * ;
13- use std:: sync:: atomic:: { AtomicU64 , Ordering } ;
12+ use std:: sync:: atomic:: { AtomicBool , AtomicU64 , Ordering } ;
1413use std:: time:: Duration ;
1514use std:: { fmt:: Debug , sync:: Mutex , time:: Instant } ;
1615
@@ -132,9 +131,16 @@ where
132131 let witness_found = Mutex :: < Option < PF < EF > > > :: new ( None ) ;
133132 // each batch tests lanes witnesses simultaneously
134133 let num_batches = PF :: < EF > :: ORDER_U64 . div_ceil ( lanes as u64 ) ;
135- ( 0 ..num_batches)
136- . into_par_iter ( )
137- . find_any ( |& batch| {
134+ // Work-stealing parallel search: each worker pulls batches from a shared counter and
135+ // stops once any worker has found a witness (`found`).
136+ let next_batch = AtomicU64 :: new ( 0 ) ;
137+ let found = AtomicBool :: new ( false ) ;
138+ parallel:: for_each_index ( parallel:: num_threads ( ) , |_| {
139+ while !found. load ( Ordering :: Relaxed ) {
140+ let batch = next_batch. fetch_add ( 1 , Ordering :: Relaxed ) ;
141+ if batch >= num_batches {
142+ break ;
143+ }
138144 let base = batch * lanes as u64 ;
139145
140146 let packed_witnesses = Packed :: < EF > :: from_fn ( |lane| {
@@ -159,14 +165,14 @@ where
159165 let rand_usize = sample. as_canonical_u64 ( ) as usize ;
160166 if ( rand_usize & ( ( 1 << bits) - 1 ) ) == 0 {
161167 * witness_found. lock ( ) . unwrap ( ) = Some ( * witness) ;
162- return true ;
168+ found. store ( true , Ordering :: Relaxed ) ;
169+ break ;
163170 }
164171 }
165- false
166- } )
167- . expect ( "failed to find witness" ) ;
172+ }
173+ } ) ;
168174
169- let witness = witness_found. lock ( ) . unwrap ( ) . unwrap ( ) ;
175+ let witness = witness_found. lock ( ) . unwrap ( ) . expect ( "failed to find witness" ) ;
170176
171177 self . challenger . observe_many ( & [ witness] ) ;
172178 assert ! ( self . challenger. state[ CAPACITY ] . as_canonical_u64( ) & ( ( 1 << bits) - 1 ) == 0 ) ;
0 commit comments