Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
251 changes: 218 additions & 33 deletions crates/asap-aware-mapping/src/frequency_rewrite.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,11 @@ fn expand(
fn substitute(expr: &ScalarExpr, cols: &[ProjectItem]) -> Option<ScalarExpr> {
Some(match expr {
ScalarExpr::Column(id) => cols.get(*id)?.expr.clone(),
ScalarExpr::Literal(_) => expr.clone(),
ScalarExpr::Negative { expr, semantics } => ScalarExpr::Negative {
expr: Box::new(substitute(expr, cols)?),
semantics: *semantics,
},
ScalarExpr::Cast {
expr,
to,
Expand Down Expand Up @@ -64,12 +69,7 @@ fn uncast(expr: &ScalarExpr) -> &ScalarExpr {
}

pub(super) fn frequency_l2_rewrite(root: &Rc<OperatorNode>) -> Option<Rc<OperatorNode>> {
let NonASAPOp::Project {
cols,
child,
qualifier,
} = root.non_asap()?
else {
let NonASAPOp::Project { cols, child, .. } = root.non_asap()? else {
return None;
};
let [item] = cols.as_slice() else {
Expand Down Expand Up @@ -119,66 +119,251 @@ pub(super) fn frequency_l2_rewrite(root: &Rc<OperatorNode>) -> Option<Rc<Operato
if product.scalar_type(&inner.schema).ok()?.0 != DataType::Float64 {
return None;
}
if !matches!(
(uncast(left), uncast(right)),
(ScalarExpr::Column(1), ScalarExpr::Column(1))
) {
return None;
}
let (key, accuracy, input) = grouped_unit_count(&inner)?;
sql_frequency_result(
root,
input,
AggIntent::FrequencyL2 {
col: Some(key),
accuracy,
},
"frequency_l2",
ScalarExpr::Column(0),
)
}

// Both frequency rules require a complete, unfiltered count per non-NULL identity.
fn grouped_unit_count(
node: &Rc<OperatorNode>,
) -> Option<(usize, asap_types::types::AccuracyTarget, Rc<OperatorNode>)> {
let NonASAPOp::Aggregate {
reduction: Reduction::Reduce(keys),
measures,
filters,
having: None,
child,
..
} = inner.non_asap()?
} = node.non_asap()?
else {
return None;
};
let [key] = keys.keys() else {
return None;
};
if keys.is_without() || any_measure_filtered(filters) {
return None;
}
let [AggIntent::Count { accuracy }] = measures.as_slice() else {
return None;
};
if !matches!(
(uncast(left), uncast(right)),
(ScalarExpr::Column(1), ScalarExpr::Column(1))
) {
return None;
}
let field = child.schema.fields.get(*key)?;
// COUNT(*) GROUP BY NULL creates a real group; the frequency intent skips it.
if field.nullable
if keys.is_without()
|| any_measure_filtered(filters)
|| field.nullable
|| !matches!(
field.plain_dtype()?,
DataType::Bool | DataType::Int64 | DataType::Utf8
)
{
return None;
}
let aggregate = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Aggregate {
reduction: Reduction::Reduce(GroupKeys::none()),
measures: vec![AggIntent::FrequencyL2 {
col: Some(*key),
accuracy: accuracy.clone(),
}],
output_names: vec!["frequency_l2".into()],
filters: vec![],
Some((*key, accuracy.clone(), Rc::clone(child)))
}

fn probability_term(product: &ScalarExpr) -> Option<&ScalarExpr> {
let ScalarExpr::Arithmetic {
op: ArithmeticOpKind::Mul,
left,
right,
semantics: ExprSemantics::Sql,
} = product
else {
return None;
};
for (probability, logarithm) in [(left, right), (right, left)] {
let ScalarExpr::FunctionCall { name, args } = logarithm.as_ref() else {
continue;
};
if name.eq_ignore_ascii_case("ln") && args.as_slice() == [probability.as_ref().clone()] {
return Some(probability);
}
}
None
}

fn unit_count_term(expr: &ScalarExpr) -> bool {
match uncast(expr) {
ScalarExpr::Column(1) => true,
ScalarExpr::Arithmetic { op: ArithmeticOpKind::Mul, left, right, semantics: ExprSemantics::Sql } => {
[(left, right), (right, left)].into_iter().any(|(count, scale)| matches!(uncast(count), ScalarExpr::Column(1)) && matches!(scale.as_ref(), ScalarExpr::Literal(ScalarValue::Float64(value)) if *value == 1.0))
},
_ => false,
}
}

pub(super) fn frequency_entropy_rewrite(root: &Rc<OperatorNode>) -> Option<Rc<OperatorNode>> {
use asap_types::pre_asap::{WindowFrameBound, WindowFrameOffset, WindowFuncKind};
let NonASAPOp::Project { cols, child, .. } = root.non_asap()? else {
return None;
};
let [item] = cols.as_slice() else {
return None;
};
let (expr, outer) = expand(item.expr.clone(), Rc::clone(child))?;
let ScalarExpr::Negative {
expr,
semantics: ExprSemantics::Sql,
} = expr
else {
return None;
};
if !matches!(uncast(&expr), ScalarExpr::Column(0)) {
return None;
}
let NonASAPOp::Aggregate {
reduction: Reduction::Reduce(keys),
measures,
filters,
having: None,
child: Rc::clone(child),
child,
..
} = outer.non_asap()?
else {
return None;
};
if keys.is_without() || !keys.keys().is_empty() || any_measure_filtered(filters) {
return None;
}
let [AggIntent::Sum { col: Some(col) }] = measures.as_slice() else {
return None;
};
let (product, window) = expand(ScalarExpr::Column(*col), Rc::clone(child))?;
let probability = probability_term(&product)?;
let ScalarExpr::Arithmetic {
op: ArithmeticOpKind::Div,
left,
right,
semantics: ExprSemantics::Sql,
} = probability
else {
return None;
};
if !unit_count_term(left)
|| !matches!(uncast(right), ScalarExpr::Column(2))
|| probability.scalar_type(&window.schema).ok()?.0 != DataType::Float64
{
return None;
}
let NonASAPOp::SQLWindowFunc {
func: WindowFuncKind::Sum,
args,
partition_by,
order_by,
frame: Some(frame),
child,
..
} = window.non_asap()?
else {
return None;
};
if partition_by.is_without()
|| !partition_by.keys().is_empty()
|| !order_by.is_empty()
|| args.as_slice() != [ScalarExpr::Column(1)]
{
return None;
}
if !matches!(
frame.start_bound,
WindowFrameBound::Preceding(WindowFrameOffset::Scalar(ScalarValue::Null))
) || !matches!(
frame.end_bound,
WindowFrameBound::Following(WindowFrameOffset::Scalar(ScalarValue::Null))
) {
return None;
}
let (key, accuracy, input) = grouped_unit_count(child)?;
let nats = ScalarExpr::Arithmetic {
op: ArithmeticOpKind::Mul,
left: Box::new(ScalarExpr::Column(0)),
right: Box::new(ScalarExpr::Literal(ScalarValue::Float64(
std::f64::consts::LN_2,
))),
semantics: ExprSemantics::Sql,
};
// -SUM(p*LN(p)) is negative zero for a single-identity population.
let nats = ScalarExpr::Negative {
expr: Box::new(ScalarExpr::Arithmetic {
op: ArithmeticOpKind::Sub,
left: Box::new(ScalarExpr::Literal(ScalarValue::Float64(0.0))),
right: Box::new(nats),
semantics: ExprSemantics::Sql,
}),
semantics: ExprSemantics::Sql,
};
sql_frequency_result(
root,
input,
AggIntent::FrequencyEntropy {
col: Some(key),
accuracy,
},
"frequency_entropy",
nats,
)
}

// Both rules use an exact population guard: an approximate statistic may be
// zero even for nonempty input, and must not control SQL's NULL result.
fn sql_frequency_result(
root: &Rc<OperatorNode>,
input: Rc<OperatorNode>,
measure: AggIntent,
name: &str,
value: ScalarExpr,
) -> Option<Rc<OperatorNode>> {
use asap_types::{ir::Predicate, pre_asap::JoinKind, types::AccuracyTarget};
let NonASAPOp::Project { qualifier, .. } = root.non_asap()? else {
return None;
};
let aggregate = |measure, name: &str| {
OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Aggregate {
reduction: Reduction::Reduce(GroupKeys::none()),
measures: vec![measure],
output_names: vec![name.into()],
filters: vec![],
having: None,
child: Rc::clone(&input),
}))
.ok()
};
let statistic = aggregate(measure, name)?;
let count = aggregate(
AggIntent::Count {
accuracy: AccuracyTarget::Exact,
},
"population_count",
)?;
let child = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Join {
kind: JoinKind::Cross,
pred: Predicate(ScalarExpr::Literal(ScalarValue::Boolean(true))),
left: statistic,
right: count,
}))
.ok()?;
// L2 is positive for any nonempty unit-update population. Restore SQL SUM's
// NULL on an empty relation without introducing another count computation.
let rewritten = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Project {
cols: vec![ProjectItem {
alias: Some(root.schema.fields.first()?.name.clone()),
expr: ScalarExpr::Case {
operand: None,
branches: vec![(
ScalarExpr::Compare {
left: Box::new(ScalarExpr::Column(0)),
left: Box::new(ScalarExpr::Column(1)),
op: CompareOpKind::Eq,
right: Box::new(ScalarExpr::Literal(ScalarValue::Float64(0.0))),
right: Box::new(ScalarExpr::Literal(ScalarValue::Int64(0))),
semantics: ExprSemantics::Sql,
},
ScalarExpr::Cast {
Expand All @@ -187,11 +372,11 @@ pub(super) fn frequency_l2_rewrite(root: &Rc<OperatorNode>) -> Option<Rc<Operato
try_cast: false,
},
)],
else_expr: Some(Box::new(ScalarExpr::Column(0))),
else_expr: Some(Box::new(value)),
},
}],
qualifier: qualifier.clone(),
child: aggregate,
child,
}))
.ok()?;
(root.schema == rewritten.schema).then_some(rewritten)
Expand Down
14 changes: 13 additions & 1 deletion crates/asap-aware-mapping/src/rewrite.rs
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,10 @@
//! frequency L2 alternative through the same semantic strategy. Projection
//! lineage, predicates, accuracy and empty-input NULL are preserved. Integer
//! products and nullable grouping keys are excluded because overflow and NULL
//! groups have observable SQL behavior. See `frequency_rewrite` for the rule.
//! groups have observable SQL behavior. Normalized natural-log entropy adds an
//! entropy alternative with explicit bits-to-nats conversion. Both frequency
//! rules retain SQL empty-input NULL using an exact population guard. See
//! `frequency_rewrite` for the rules.
//!
//! ## Scope
//!
Expand Down Expand Up @@ -408,9 +411,18 @@ impl ReplacementStrategy for SemanticEquivalentRewriteStrategy {
avg_rewrite_target(target.root).is_some()
|| composed_aggregate_rewrite(target.root).is_some()
|| crate::frequency_rewrite::frequency_l2_rewrite(target.root).is_some()
|| crate::frequency_rewrite::frequency_entropy_rewrite(target.root).is_some()
}

fn replacements(&self, target: &TargetSubDAG<'_>) -> Vec<ReplacementSubDAG> {
if let Some(rewritten) = crate::frequency_rewrite::frequency_entropy_rewrite(target.root) {
return vec![ReplacementSubDAG {
strategy: "SemanticEquivalentRewriteStrategy",
replacement: Replacement::SubDAG(rewritten),
provenance: crate::replacement::ReplacementProvenance::LogicalRewrite,
rationale: "recognize SQL natural-log entropy with explicit bits-to-nats conversion and exact empty-population guard".into(),
}];
}
if let Some(rewritten) = crate::frequency_rewrite::frequency_l2_rewrite(target.root) {
return vec![ReplacementSubDAG {
strategy: "SemanticEquivalentRewriteStrategy",
Expand Down
Loading
Loading