@@ -7,6 +7,7 @@ use std::time::Instant;
77
88const BATCH_PRUNING_MAX_SIZE : usize = 32 ;
99const BATCH_PRUNING_FIXED_PREFIX : usize = 4 ;
10+ const BATCH_PRUNING_GLOBAL_MAX : usize = 42 ;
1011
1112#[ derive( Debug ) ]
1213#[ cfg_attr( test, derive( Eq , PartialEq ) ) ]
@@ -177,16 +178,46 @@ pub(crate) fn create_task_batches(
177178 ) ;
178179 b. size > 0
179180 } ) ;
181+ apply_global_cut_budget (
182+ & mut batches,
183+ BATCH_PRUNING_FIXED_PREFIX ,
184+ BATCH_PRUNING_GLOBAL_MAX ,
185+ ) ;
180186 batches
181187}
182188
189+ fn apply_global_cut_budget ( batches : & mut [ TaskBatch ] , prefix : usize , budget : usize ) {
190+ let total: usize = batches. iter ( ) . map ( |b| b. cuts . len ( ) ) . sum ( ) ;
191+ if total <= budget {
192+ return ;
193+ }
194+ let n_nonempty = batches. iter ( ) . filter ( |b| !b. cuts . is_empty ( ) ) . count ( ) ;
195+ let prefix = prefix. min ( budget / n_nonempty) . max ( 1 ) ;
196+ let extra_budget = budget. saturating_sub ( prefix * n_nonempty) ;
197+ let extra_total: usize = batches
198+ . iter ( )
199+ . map ( |b| b. cuts . len ( ) . saturating_sub ( prefix) )
200+ . sum ( ) ;
201+ for b in batches. iter_mut ( ) . filter ( |b| !b. cuts . is_empty ( ) ) {
202+ let extra = ( extra_budget * b. cuts . len ( ) . saturating_sub ( prefix) )
203+ . checked_div ( extra_total)
204+ . unwrap_or ( 0 ) ;
205+ prune_progressive ( & mut b. cuts , prefix, prefix + extra) ;
206+ }
207+ }
208+
183209fn prune_progressive < T > ( vec : & mut Vec < T > , prefix_size : usize , size_limit : usize ) {
184210 let original_len = vec. len ( ) ;
185211
186212 if original_len <= size_limit {
187213 return ;
188214 }
189215
216+ if size_limit <= prefix_size + 1 {
217+ vec. truncate ( size_limit) ;
218+ return ;
219+ }
220+
190221 let remaining_slots = size_limit - prefix_size;
191222
192223 let mut indices = Vec :: with_capacity ( size_limit) ;
@@ -248,4 +279,59 @@ mod tests {
248279 ]
249280 ) ;
250281 }
282+
283+ /// A global budget can ask for a limit at or just above the prefix, where the quadratic
284+ /// sampler has no slots left to place (`i / (slots - 1)` would be `0.0 / 0.0`).
285+ #[ test]
286+ fn test_prune_progressive_at_prefix_boundary ( ) {
287+ for size_limit in 0 ..=5 {
288+ let mut vec = ( 0 ..40 ) . collect :: < Vec < _ > > ( ) ;
289+ prune_progressive ( & mut vec, 4 , size_limit) ;
290+ assert_eq ! ( vec, ( 0 ..size_limit as i32 ) . collect:: <Vec <_>>( ) ) ;
291+ }
292+
293+ let mut vec = ( 0 ..40 ) . collect :: < Vec < _ > > ( ) ;
294+ prune_progressive ( & mut vec, 4 , 6 ) ;
295+ assert_eq ! ( vec, vec![ 0 , 1 , 2 , 3 , 4 , 39 ] ) ;
296+ }
297+
298+ #[ test]
299+ fn test_global_cut_budget ( ) {
300+ fn batch ( n_cuts : usize ) -> TaskBatch {
301+ let mut b = TaskBatch :: new ( 0 . into ( ) , 100 , false ) ;
302+ b. cuts = ( 0 ..n_cuts)
303+ . map ( |i| PriorityCut {
304+ size : i as u32 ,
305+ blockers : Vec :: new ( ) ,
306+ } )
307+ . collect ( ) ;
308+ b
309+ }
310+ let total = |bs : & [ TaskBatch ] | bs. iter ( ) . map ( |b| b. cuts . len ( ) ) . sum :: < usize > ( ) ;
311+
312+ // Under budget: untouched.
313+ let mut batches = vec ! [ batch( 3 ) , batch( 3 ) ] ;
314+ apply_global_cut_budget ( & mut batches, 4 , 32 ) ;
315+ assert_eq ! ( total( & batches) , 6 ) ;
316+
317+ // What the per-batch cap cannot reach: 8 batches of 3, each below it, 24 in total.
318+ let mut batches: Vec < _ > = ( 0 ..8 ) . map ( |_| batch ( 3 ) ) . collect ( ) ;
319+ apply_global_cut_budget ( & mut batches, 4 , 8 ) ;
320+ assert_eq ! ( total( & batches) , 8 ) ;
321+
322+ // Proportional: the bigger batch keeps more, and the budget holds.
323+ let mut batches = vec ! [ batch( 60 ) , batch( 10 ) , batch( 10 ) ] ;
324+ apply_global_cut_budget ( & mut batches, 4 , 16 ) ;
325+ assert ! ( batches[ 0 ] . cuts. len( ) > batches[ 1 ] . cuts. len( ) ) ;
326+ assert ! ( total( & batches) <= 16 ) ;
327+
328+ // Below the batch count, each batch still keeps its first cut.
329+ let mut batches: Vec < _ > = ( 0 ..8 ) . map ( |_| batch ( 5 ) ) . collect ( ) ;
330+ apply_global_cut_budget ( & mut batches, 4 , 2 ) ;
331+ assert ! (
332+ batches
333+ . iter( )
334+ . all( |b| b. cuts. len( ) == 1 && b. cuts[ 0 ] . size == 0 )
335+ ) ;
336+ }
251337}
0 commit comments