@@ -52,7 +52,8 @@ use crate::stats::session::StatsSession;
5252
5353/// Register built-in stats rewrite rules.
5454pub ( crate ) fn register_builtins ( session : & StatsSession ) {
55- session. register_rewrite ( BinaryStatsRewrite ) ;
55+ session. register_rewrite ( BinaryStatsRewrite :: legacy_nan_count ( ) ) ;
56+ session. register_rewrite ( BinaryStatsRewrite :: all_non_nan ( ) ) ;
5657 session. register_rewrite ( BetweenStatsRewrite ) ;
5758 session. register_rewrite ( IsNullLegacyStatsRewrite ) ;
5859 session. register_rewrite ( IsNullAllNonNullStatsRewrite ) ;
@@ -61,12 +62,30 @@ pub(crate) fn register_builtins(session: &StatsSession) {
6162 session. register_rewrite ( IsNotNullAllNullStatsRewrite ) ;
6263 session. register_rewrite ( IsNotNullAllNonNullStatsRewrite ) ;
6364 session. register_rewrite ( LikeStatsRewrite ) ;
64- session. register_rewrite ( ListContainsStatsRewrite ) ;
65- session. register_rewrite ( DynamicComparisonStatsRewrite ) ;
65+ session. register_rewrite ( ListContainsStatsRewrite :: legacy_nan_count ( ) ) ;
66+ session. register_rewrite ( ListContainsStatsRewrite :: all_non_nan ( ) ) ;
67+ session. register_rewrite ( DynamicComparisonStatsRewrite :: legacy_nan_count ( ) ) ;
68+ session. register_rewrite ( DynamicComparisonStatsRewrite :: all_non_nan ( ) ) ;
6669}
6770
6871#[ derive( Debug ) ]
69- struct BinaryStatsRewrite ;
72+ struct BinaryStatsRewrite {
73+ nan_proof : NanProof ,
74+ }
75+
76+ impl BinaryStatsRewrite {
77+ fn legacy_nan_count ( ) -> Self {
78+ Self {
79+ nan_proof : NanProof :: LegacyNanCount ,
80+ }
81+ }
82+
83+ fn all_non_nan ( ) -> Self {
84+ Self {
85+ nan_proof : NanProof :: AllNonNan ,
86+ }
87+ }
88+ }
7089
7190impl StatsRewriteRule for BinaryStatsRewrite {
7291 fn scalar_fn_id ( & self ) -> ScalarFnId {
@@ -87,44 +106,59 @@ impl StatsRewriteRule for BinaryStatsRewrite {
87106 let left = min ( lhs, ctx) . zip ( max ( rhs, ctx) ) . map ( |( a, b) | gt ( a, b) ) ;
88107 let right = min ( rhs, ctx) . zip ( max ( lhs, ctx) ) . map ( |( a, b) | gt ( a, b) ) ;
89108 or_collect ( left. into_iter ( ) . chain ( right) )
90- . map ( |value_predicate| with_nan_predicate ( ctx, lhs, rhs, value_predicate) )
109+ . map ( |value_predicate| {
110+ with_nan_predicate ( ctx, self . nan_proof , lhs, rhs, value_predicate)
111+ } )
91112 . transpose ( ) ?
113+ . flatten ( )
92114 }
93115 Operator :: NotEq => min ( lhs, ctx)
94116 . zip ( max ( rhs, ctx) )
95117 . zip ( max ( lhs, ctx) . zip ( min ( rhs, ctx) ) )
96118 . map ( |( ( min_lhs, max_rhs) , ( max_lhs, min_rhs) ) | {
97119 with_nan_predicate (
98120 ctx,
121+ self . nan_proof ,
99122 lhs,
100123 rhs,
101124 and ( eq ( min_lhs, max_rhs) , eq ( max_lhs, min_rhs) ) ,
102125 )
103126 } )
104- . transpose ( ) ?,
127+ . transpose ( ) ?
128+ . flatten ( ) ,
105129 Operator :: Gt => max ( lhs, ctx)
106130 . zip ( min ( rhs, ctx) )
107- . map ( |( a, b) | with_nan_predicate ( ctx, lhs, rhs, lt_eq ( a, b) ) )
108- . transpose ( ) ?,
131+ . map ( |( a, b) | with_nan_predicate ( ctx, self . nan_proof , lhs, rhs, lt_eq ( a, b) ) )
132+ . transpose ( ) ?
133+ . flatten ( ) ,
109134 Operator :: Gte => max ( lhs, ctx)
110135 . zip ( min ( rhs, ctx) )
111- . map ( |( a, b) | with_nan_predicate ( ctx, lhs, rhs, lt ( a, b) ) )
112- . transpose ( ) ?,
136+ . map ( |( a, b) | with_nan_predicate ( ctx, self . nan_proof , lhs, rhs, lt ( a, b) ) )
137+ . transpose ( ) ?
138+ . flatten ( ) ,
113139 Operator :: Lt => min ( lhs, ctx)
114140 . zip ( max ( rhs, ctx) )
115- . map ( |( a, b) | with_nan_predicate ( ctx, lhs, rhs, gt_eq ( a, b) ) )
116- . transpose ( ) ?,
141+ . map ( |( a, b) | with_nan_predicate ( ctx, self . nan_proof , lhs, rhs, gt_eq ( a, b) ) )
142+ . transpose ( ) ?
143+ . flatten ( ) ,
117144 Operator :: Lte => min ( lhs, ctx)
118145 . zip ( max ( rhs, ctx) )
119- . map ( |( a, b) | with_nan_predicate ( ctx, lhs, rhs, gt ( a, b) ) )
120- . transpose ( ) ?,
146+ . map ( |( a, b) | with_nan_predicate ( ctx, self . nan_proof , lhs, rhs, gt ( a, b) ) )
147+ . transpose ( ) ?
148+ . flatten ( ) ,
121149 Operator :: And => {
150+ if !self . nan_proof . emits_unguarded_rewrites ( ) {
151+ return Ok ( None ) ;
152+ }
153+
122154 let lhs_falsifier = ctx. falsify ( lhs) ?;
123155 let rhs_falsifier = ctx. falsify ( rhs) ?;
124156 or_collect ( lhs_falsifier. into_iter ( ) . chain ( rhs_falsifier) )
125157 }
126158 Operator :: Or => match ( ctx. falsify ( lhs) ?, ctx. falsify ( rhs) ?) {
127- ( Some ( lhs) , Some ( rhs) ) => Some ( and ( lhs, rhs) ) ,
159+ ( Some ( lhs) , Some ( rhs) ) if self . nan_proof . emits_unguarded_rewrites ( ) => {
160+ Some ( and ( lhs, rhs) )
161+ }
128162 _ => None ,
129163 } ,
130164 Operator :: Add | Operator :: Sub | Operator :: Mul | Operator :: Div => None ,
@@ -332,7 +366,23 @@ impl StatsRewriteRule for LikeStatsRewrite {
332366}
333367
334368#[ derive( Debug ) ]
335- struct ListContainsStatsRewrite ;
369+ struct ListContainsStatsRewrite {
370+ nan_proof : NanProof ,
371+ }
372+
373+ impl ListContainsStatsRewrite {
374+ fn legacy_nan_count ( ) -> Self {
375+ Self {
376+ nan_proof : NanProof :: LegacyNanCount ,
377+ }
378+ }
379+
380+ fn all_non_nan ( ) -> Self {
381+ Self {
382+ nan_proof : NanProof :: AllNonNan ,
383+ }
384+ }
385+ }
336386
337387impl StatsRewriteRule for ListContainsStatsRewrite {
338388 fn scalar_fn_id ( & self ) -> ScalarFnId {
@@ -375,13 +425,32 @@ impl StatsRewriteRule for ListContainsStatsRewrite {
375425 )
376426 } ) ) ;
377427 value_predicate
378- . map ( |value_predicate| with_all_non_nan_predicate ( ctx, [ needle] , value_predicate) )
428+ . map ( |value_predicate| {
429+ with_all_non_nan_predicate ( ctx, self . nan_proof , [ needle] , value_predicate)
430+ } )
379431 . transpose ( )
432+ . map ( Option :: flatten)
380433 }
381434}
382435
383436#[ derive( Debug ) ]
384- struct DynamicComparisonStatsRewrite ;
437+ struct DynamicComparisonStatsRewrite {
438+ nan_proof : NanProof ,
439+ }
440+
441+ impl DynamicComparisonStatsRewrite {
442+ fn legacy_nan_count ( ) -> Self {
443+ Self {
444+ nan_proof : NanProof :: LegacyNanCount ,
445+ }
446+ }
447+
448+ fn all_non_nan ( ) -> Self {
449+ Self {
450+ nan_proof : NanProof :: AllNonNan ,
451+ }
452+ }
453+ }
385454
386455impl StatsRewriteRule for DynamicComparisonStatsRewrite {
387456 fn scalar_fn_id ( & self ) -> ScalarFnId {
@@ -414,7 +483,7 @@ impl StatsRewriteRule for DynamicComparisonStatsRewrite {
414483 } ,
415484 [ lhs_stat] ,
416485 ) ;
417- with_all_non_nan_predicate ( ctx, [ lhs] , value_predicate) . map ( Some )
486+ with_all_non_nan_predicate ( ctx, self . nan_proof , [ lhs] , value_predicate)
418487 }
419488}
420489
@@ -438,43 +507,64 @@ fn all_non_null(expr: &Expression) -> Expression {
438507 stat_fn ( expr. clone ( ) , AllNonNull . bind ( AggregateEmptyOptions ) )
439508}
440509
510+ #[ derive( Debug , Clone , Copy ) ]
511+ enum NanProof {
512+ LegacyNanCount ,
513+ AllNonNan ,
514+ }
515+
516+ impl NanProof {
517+ fn emits_unguarded_rewrites ( self ) -> bool {
518+ matches ! ( self , Self :: LegacyNanCount )
519+ }
520+ }
521+
522+ enum NanCheck {
523+ NotNeeded ,
524+ Check ( Expression ) ,
525+ Unavailable ,
526+ }
527+
441528// Min/max do not order NaN values, so comparison rewrites are only sound when every
442529// candidate value is known to be non-NaN. Cast result dtypes are not enough: a cast
443530// from float to non-float still needs a proof about the float source values.
444531fn all_non_nan_stat (
445532 ctx : & StatsRewriteCtx < ' _ > ,
533+ nan_proof : NanProof ,
446534 expr : & Expression ,
447- ) -> VortexResult < Option < Expression > > {
535+ ) -> VortexResult < NanCheck > {
448536 if let Some ( scalar) = expr. as_opt :: < Literal > ( ) {
449537 let Some ( value) = scalar. as_primitive_opt ( ) else {
450- return Ok ( None ) ;
538+ return Ok ( NanCheck :: NotNeeded ) ;
451539 } ;
452- return Ok ( value. is_nan ( ) . then ( || lit ( false ) ) ) ;
540+ return Ok ( if value. is_nan ( ) {
541+ NanCheck :: Check ( lit ( false ) )
542+ } else {
543+ NanCheck :: NotNeeded
544+ } ) ;
453545 }
454546
455547 if expr. is :: < Cast > ( ) {
456548 if !has_nans ( & ctx. return_dtype ( expr. child ( 0 ) ) ?) {
457- return Ok ( None ) ;
549+ return Ok ( NanCheck :: NotNeeded ) ;
458550 }
459551
460- return all_non_nan_stat ( ctx, expr. child ( 0 ) ) ;
552+ return all_non_nan_stat ( ctx, nan_proof , expr. child ( 0 ) ) ;
461553 }
462554
463555 if !has_nans ( & ctx. return_dtype ( expr) ?) {
464- return Ok ( None ) ;
556+ return Ok ( NanCheck :: NotNeeded ) ;
465557 }
466558
467- let Some ( nan_count) = stat_expr ( expr, Stat :: NaNCount , ctx) else {
468- return Ok ( Some ( stat_fn (
469- expr. clone ( ) ,
470- AllNonNan . bind ( AggregateEmptyOptions ) ,
471- ) ) ) ;
472- } ;
473-
474- Ok ( Some ( or (
475- eq ( nan_count, lit ( 0u64 ) ) ,
476- stat_fn ( expr. clone ( ) , AllNonNan . bind ( AggregateEmptyOptions ) ) ,
477- ) ) )
559+ Ok ( match nan_proof {
560+ NanProof :: LegacyNanCount => match stat_expr ( expr, Stat :: NaNCount , ctx) {
561+ Some ( nan_count) => NanCheck :: Check ( eq ( nan_count, lit ( 0u64 ) ) ) ,
562+ None => NanCheck :: Unavailable ,
563+ } ,
564+ NanProof :: AllNonNan => {
565+ NanCheck :: Check ( stat_fn ( expr. clone ( ) , AllNonNan . bind ( AggregateEmptyOptions ) ) )
566+ }
567+ } )
478568}
479569
480570fn has_nans ( dtype : & DType ) -> bool {
@@ -510,31 +600,37 @@ fn stat_expr(expr: &Expression, stat: Stat, ctx: &StatsRewriteCtx<'_>) -> Option
510600
511601fn with_nan_predicate (
512602 ctx : & StatsRewriteCtx < ' _ > ,
603+ nan_proof : NanProof ,
513604 lhs : & Expression ,
514605 rhs : & Expression ,
515606 value_predicate : Expression ,
516- ) -> VortexResult < Expression > {
517- with_all_non_nan_predicate ( ctx, [ lhs, rhs] , value_predicate)
607+ ) -> VortexResult < Option < Expression > > {
608+ with_all_non_nan_predicate ( ctx, nan_proof , [ lhs, rhs] , value_predicate)
518609}
519610
520611fn with_all_non_nan_predicate < ' a > (
521612 ctx : & StatsRewriteCtx < ' _ > ,
613+ nan_proof : NanProof ,
522614 exprs : impl IntoIterator < Item = & ' a Expression > ,
523615 value_predicate : Expression ,
524- ) -> VortexResult < Expression > {
616+ ) -> VortexResult < Option < Expression > > {
525617 let mut nan_checks = Vec :: new ( ) ;
526618 for expr in exprs {
527- if let Some ( check) = all_non_nan_stat ( ctx, expr) ? {
528- nan_checks. push ( check) ;
619+ match all_non_nan_stat ( ctx, nan_proof, expr) ? {
620+ NanCheck :: NotNeeded => { }
621+ NanCheck :: Check ( check) => nan_checks. push ( check) ,
622+ NanCheck :: Unavailable => return Ok ( None ) ,
529623 }
530624 }
531625 let nan_predicate = and_collect ( nan_checks) ;
532626
533627 Ok ( match nan_predicate {
534- Some ( nan_check) => and ( nan_check, value_predicate) ,
628+ Some ( nan_check) => Some ( and ( nan_check, value_predicate) ) ,
535629 // No possible NaN-bearing expression remains, so the value predicate is
536- // already guarded.
537- None => value_predicate,
630+ // already guarded. Only one registered rule emits this unguarded
631+ // rewrite so non-float comparisons are not duplicated.
632+ None if nan_proof. emits_unguarded_rewrites ( ) => Some ( value_predicate) ,
633+ None => None ,
538634 } )
539635}
540636
@@ -666,10 +762,16 @@ mod tests {
666762 expr. satisfy ( & test_scope ( ) , & SESSION )
667763 }
668764
669- fn nan_guard ( expr : Expression ) -> Expression {
765+ fn nan_guarded ( expr : Expression , value_predicate : Expression ) -> Expression {
670766 or (
671- eq ( stat ( expr. clone ( ) , Stat :: NaNCount ) , lit ( 0u64 ) ) ,
672- stat_fn ( expr, AllNonNan . bind ( AggregateEmptyOptions ) ) ,
767+ and (
768+ eq ( stat ( expr. clone ( ) , Stat :: NaNCount ) , lit ( 0u64 ) ) ,
769+ value_predicate. clone ( ) ,
770+ ) ,
771+ and (
772+ stat_fn ( expr, AllNonNan . bind ( AggregateEmptyOptions ) ) ,
773+ value_predicate,
774+ ) ,
673775 )
674776 }
675777
@@ -897,8 +999,8 @@ mod tests {
897999
8981000 assert_eq ! (
8991001 falsify( & expr) ?,
900- Some ( and (
901- nan_guard ( col( "f" ) ) ,
1002+ Some ( nan_guarded (
1003+ col( "f" ) ,
9021004 lt_eq( cast( stat( col( "f" ) , Stat :: Max ) , dtype) , lit( 5i32 ) ) ,
9031005 ) )
9041006 ) ;
0 commit comments