Skip to content

Commit b222454

Browse files
committed
Split NaN stat rewrite proofs
Signed-off-by: "Nicholas Gates" <nick@nickgates.com>
1 parent 25e9400 commit b222454

1 file changed

Lines changed: 151 additions & 49 deletions

File tree

vortex-array/src/stats/rewrite/builtins.rs

Lines changed: 151 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,8 @@ use crate::stats::session::StatsSession;
5252

5353
/// Register built-in stats rewrite rules.
5454
pub(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

7190
impl 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

337387
impl 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

386455
impl 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.
444531
fn 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

480570
fn has_nans(dtype: &DType) -> bool {
@@ -510,31 +600,37 @@ fn stat_expr(expr: &Expression, stat: Stat, ctx: &StatsRewriteCtx<'_>) -> Option
510600

511601
fn 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

520611
fn 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

Comments
 (0)