Skip to content
Closed
Show file tree
Hide file tree
Changes from 4 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
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,8 @@ abstract class Optimizer(sessionCatalog: SessionCatalog, conf: CatalystConf)
SimplifyCaseConversionExpressions,
RewriteCorrelatedScalarSubquery,
EliminateSerialization,
RemoveAliasOnlyProject) ::
RemoveRedundantAliases,
RemoveRedundantProject) ::
Batch("Check Cartesian Products", Once,
CheckCartesianProducts(conf)) ::
Batch("Decimal Optimizations", fixedPoint,
Expand Down Expand Up @@ -154,56 +155,108 @@ class SimpleTestOptimizer extends Optimizer(
new SimpleCatalystConf(caseSensitiveAnalysis = true))

/**
* Removes the Project only conducting Alias of its child node.
* It is created mainly for removing extra Project added in EliminateSerialization rule,
* but can also benefit other operators.
* Remove redundant aliases from a query plan. A redundant alias is an alias that does not change
* the name or metadata of a column, and does not deduplicate it.
*/
object RemoveAliasOnlyProject extends Rule[LogicalPlan] {
object RemoveRedundantAliases extends Rule[LogicalPlan] {

/**
* Returns true if the project list is semantically same as child output, after strip alias on
* attribute.
* Replace the attributes in an expression using the given mapping.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

looks like this doc is wrong?

*/
private def isAliasOnly(
projectList: Seq[NamedExpression],
childOutput: Seq[Attribute]): Boolean = {
if (projectList.length != childOutput.length) {
false
} else {
stripAliasOnAttribute(projectList).zip(childOutput).forall {
case (a: Attribute, o) if a semanticEquals o => true
case _ => false
}
private def createAttributeMapping(current: LogicalPlan, next: LogicalPlan)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can you explain what current and next means here?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Current is plan before we remove redundant aliases, and next is the plan after we have remove the redundant aliases. I'll update the doc.

: Seq[(Attribute, Attribute)] = {
current.output.zip(next.output).filterNot {
case (a1, a2) => a1.semanticEquals(a2)
}
}

private def stripAliasOnAttribute(projectList: Seq[NamedExpression]) = {
projectList.map {
// Alias with metadata can not be stripped, or the metadata will be lost.
// If the alias name is different from attribute name, we can't strip it either, or we may
// accidentally change the output schema name of the root plan.
case a @ Alias(attr: Attribute, name) if a.metadata == Metadata.empty && name == attr.name =>
attr
case other => other
}
/**
* Remove the top-level alias from an expression when it is redundant.
*/
private def removeRedundantAlias(e: Expression, blacklist: AttributeSet): Expression = e match {
// Alias with metadata can not be stripped, or the metadata will be lost.
// If the alias name is different from attribute name, we can't strip it either, or we
// may accidentally change the output schema name of the root plan.
case a @ Alias(attr: Attribute, name)
if a.metadata == Metadata.empty && name == attr.name && !blacklist.contains(attr) =>
attr
case a => a
}

def apply(plan: LogicalPlan): LogicalPlan = {
val aliasOnlyProject = plan.collectFirst {
case p @ Project(pList, child) if isAliasOnly(pList, child.output) => p
/**
* Get an appropriate alias cleaning method for the given node.
*
* We currently clean Project, Aggregate & Window nodes.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

so this is an improvement right? previously we only clean Project. However I think this method is over engineered, we can just create a def needClean(plan: LogicalPlan): Boolean

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah that is an improvement. I added all LogicalPlan nodes that are producing new attributes using named expressions.

I will inline this method.

*/
private def getAliasCleaner(plan: LogicalPlan): (Expression, AttributeSet) => Expression = {
plan match {
case _: Project => removeRedundantAlias
case _: Aggregate => removeRedundantAlias
case _: Window => removeRedundantAlias
case _ => (e, _) => e
}
}

aliasOnlyProject.map { case proj =>
val attributesToReplace = proj.output.zip(proj.child.output).filterNot {
case (a1, a2) => a1 semanticEquals a2
}
val attrMap = AttributeMap(attributesToReplace)
plan transform {
case plan: Project if plan eq proj => plan.child
case plan => plan transformExpressions {
case a: Attribute if attrMap.contains(a) => attrMap(a)
/**
* Remove redundant alias expression from a LogicalPlan and its subtree. A blacklist is used to
* prevent the removal of seemingly redundant aliases used to deduplicate the input for a (self)
* join.
*/
private def removeRedundantAliases(plan: LogicalPlan, blacklist: AttributeSet): LogicalPlan = {
plan match {
// A join has to be treated differently, because the left and the right side of the join are
// not allowed to use the same attributes. We use a blacklist to prevent us from creating a
// situation in which this happens; the rule will only remove an alias if its child
// attribute is not on the black list.
case Join(left, right, joinType, condition) =>
val newLeft = removeRedundantAliases(left, blacklist ++ right.outputSet)
val newRight = removeRedundantAliases(right, blacklist ++ newLeft.outputSet)
val mapping = AttributeMap(
createAttributeMapping(left, newLeft) ++
createAttributeMapping(right, newRight))
val newCondition = condition.map(_.transform {
case a: Attribute => mapping.getOrElse(a, a)
})
Join(newLeft, newRight, joinType, newCondition)

case _ =>
// Drop blacklisted attributes that are masked in the current project. This allows us to
// remove redundant aliases in the subtree.
val childBlacklist = blacklist -- (plan.inputSet -- plan.outputSet)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is this branch needed because Union reuse the output of left side? can we remove it if we fix Union?

@hvanhovell hvanhovell Feb 3, 2017

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You mean the case _ => right? That is needed for everything which is not a Join. We are doing manual tree traversal here.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

sorry I mean childBlacklist. We can just use blacklist if Union is fixed right?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The child blacklist is an optimization. I can remove an attribute from the child's blacklist if I know that it is being created in the current node. This way I give the rule more freedom in removing attributes. The thing is that situation should only happen when there are multiple self joins, and this might be an over optimization.


// Remove redundant aliases in the subtree(s).
val currentNextAttrPairs = mutable.Buffer.empty[(Attribute, Attribute)]
val newNode = plan.mapChildren { child =>
val newChild = removeRedundantAliases(child, childBlacklist)
currentNextAttrPairs ++= createAttributeMapping(child, newChild)
newChild
}
}
}.getOrElse(plan)

// Create the attribute mapping. Note that the currentNextAttrPairs can contain duplicate
// keys in case of Union (this is caused by the PushProjectionThroughUnion rule); in this
// case we use the the first mapping (which should be provided by the first child).
val mapping = AttributeMap(currentNextAttrPairs)

// Transform the expressions.
val cleanExpression = getAliasCleaner(plan)
newNode.mapExpressions { expr =>
val newExpr = expr.transform {
case a: Attribute => mapping.getOrElse(a, a)
}
cleanExpression(newExpr, blacklist)
}
}
}

def apply(plan: LogicalPlan): LogicalPlan = removeRedundantAliases(plan, AttributeSet.empty)
}

/**
* Remove projections from the query plan that do not make any modifications.
*/
object RemoveRedundantProject extends Rule[LogicalPlan] {
def apply(plan: LogicalPlan): LogicalPlan = plan transform {
case p @ Project(_, child) if p.output == child.output => child
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -242,31 +242,7 @@ abstract class QueryPlan[PlanType <: QueryPlan[PlanType]] extends TreeNode[PlanT
* @param rule the rule to be applied to every expression in this operator.
*/
def transformExpressionsDown(rule: PartialFunction[Expression, Expression]): this.type = {
var changed = false

@inline def transformExpressionDown(e: Expression): Expression = {
val newE = e.transformDown(rule)
if (newE.fastEquals(e)) {
e
} else {
changed = true
newE
}
}

def recursiveTransform(arg: Any): AnyRef = arg match {
case e: Expression => transformExpressionDown(e)
case Some(e: Expression) => Some(transformExpressionDown(e))
case m: Map[_, _] => m
case d: DataType => d // Avoid unpacking Structs
case seq: Traversable[_] => seq.map(recursiveTransform)
case other: AnyRef => other
case null => null
}

val newArgs = mapProductIterator(recursiveTransform)

if (changed) makeCopy(newArgs).asInstanceOf[this.type] else this
mapExpressions(_.transformDown(rule))
}

/**
Expand All @@ -276,10 +252,18 @@ abstract class QueryPlan[PlanType <: QueryPlan[PlanType]] extends TreeNode[PlanT
* @return
*/
def transformExpressionsUp(rule: PartialFunction[Expression, Expression]): this.type = {
mapExpressions(_.transformUp(rule))
}

/**
* Apply a map function to each expression present in this query operator, and return a new
* query operator based on the mapped expressions.
*/
def mapExpressions(f: Expression => Expression): this.type = {
var changed = false

@inline def transformExpressionUp(e: Expression): Expression = {
val newE = e.transformUp(rule)
@inline def transformExpression(e: Expression): Expression = {
val newE = f(e)
if (newE.fastEquals(e)) {
e
} else {
Expand All @@ -289,8 +273,8 @@ abstract class QueryPlan[PlanType <: QueryPlan[PlanType]] extends TreeNode[PlanT
}

def recursiveTransform(arg: Any): AnyRef = arg match {
case e: Expression => transformExpressionUp(e)
case Some(e: Expression) => Some(transformExpressionUp(e))
case e: Expression => transformExpression(e)
case Some(e: Expression) => Some(transformExpression(e))
case m: Map[_, _] => m
case d: DataType => d // Avoid unpacking Structs
case seq: Traversable[_] => seq.map(recursiveTransform)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ abstract class LogicalPlan extends QueryPlan[LogicalPlan] with Logging {
*/
def resolveOperators(rule: PartialFunction[LogicalPlan, LogicalPlan]): LogicalPlan = {
if (!analyzed) {
val afterRuleOnChildren = transformChildren(rule, (t, r) => t.resolveOperators(r))
val afterRuleOnChildren = mapChildren(_.resolveOperators(rule))
if (this fastEquals afterRuleOnChildren) {
CurrentOrigin.withOrigin(origin) {
rule.applyOrElse(this, identity[LogicalPlan])
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -190,26 +190,6 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {
arr
}

/**
* Returns a copy of this node where `f` has been applied to all the nodes children.
*/
def mapChildren(f: BaseType => BaseType): BaseType = {
var changed = false
val newArgs = mapProductIterator {
case arg: TreeNode[_] if containsChild(arg) =>
val newChild = f(arg.asInstanceOf[BaseType])
if (newChild fastEquals arg) {
arg
} else {
changed = true
newChild
}
case nonChild: AnyRef => nonChild
case null => null
}
if (changed) makeCopy(newArgs) else this
}

/**
* Returns a copy of this node with the children replaced.
* TODO: Validate somewhere (in debug mode?) that children are ordered correctly.
Expand Down Expand Up @@ -289,9 +269,9 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {

// Check if unchanged and then possibly return old copy to avoid gc churn.
if (this fastEquals afterRule) {
transformChildren(rule, (t, r) => t.transformDown(r))
mapChildren(_.transformDown(rule))
} else {
afterRule.transformChildren(rule, (t, r) => t.transformDown(r))
afterRule.mapChildren(_.transformDown(rule))
}
}

Expand All @@ -303,7 +283,7 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {
* @param rule the function use to transform this nodes children
*/
def transformUp(rule: PartialFunction[BaseType, BaseType]): BaseType = {
val afterRuleOnChildren = transformChildren(rule, (t, r) => t.transformUp(r))
val afterRuleOnChildren = mapChildren(_.transformUp(rule))
if (this fastEquals afterRuleOnChildren) {
CurrentOrigin.withOrigin(origin) {
rule.applyOrElse(this, identity[BaseType])
Expand All @@ -316,26 +296,22 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {
}

/**
* Returns a copy of this node where `rule` has been recursively applied to all the children of
* this node. When `rule` does not apply to a given node it is left unchanged.
* @param rule the function used to transform this nodes children
* Returns a copy of this node where `f` has been applied to all the nodes children.
*/
protected def transformChildren(
rule: PartialFunction[BaseType, BaseType],
nextOperation: (BaseType, PartialFunction[BaseType, BaseType]) => BaseType): BaseType = {
def mapChildren(f: BaseType => BaseType): BaseType = {
if (children.nonEmpty) {
var changed = false
val newArgs = mapProductIterator {
case arg: TreeNode[_] if containsChild(arg) =>
val newChild = nextOperation(arg.asInstanceOf[BaseType], rule)
val newChild = f(arg.asInstanceOf[BaseType])
if (!(newChild fastEquals arg)) {
changed = true
newChild
} else {
arg
}
case Some(arg: TreeNode[_]) if containsChild(arg) =>
val newChild = nextOperation(arg.asInstanceOf[BaseType], rule)
val newChild = f(arg.asInstanceOf[BaseType])
if (!(newChild fastEquals arg)) {
changed = true
Some(newChild)
Expand All @@ -344,7 +320,7 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {
}
case m: Map[_, _] => m.mapValues {
case arg: TreeNode[_] if containsChild(arg) =>
val newChild = nextOperation(arg.asInstanceOf[BaseType], rule)
val newChild = f(arg.asInstanceOf[BaseType])
if (!(newChild fastEquals arg)) {
changed = true
newChild
Expand All @@ -356,16 +332,16 @@ abstract class TreeNode[BaseType <: TreeNode[BaseType]] extends Product {
case d: DataType => d // Avoid unpacking Structs
case args: Traversable[_] => args.map {
case arg: TreeNode[_] if containsChild(arg) =>
val newChild = nextOperation(arg.asInstanceOf[BaseType], rule)
val newChild = f(arg.asInstanceOf[BaseType])
if (!(newChild fastEquals arg)) {
changed = true
newChild
} else {
arg
}
case tuple@(arg1: TreeNode[_], arg2: TreeNode[_]) =>
val newChild1 = nextOperation(arg1.asInstanceOf[BaseType], rule)
val newChild2 = nextOperation(arg2.asInstanceOf[BaseType], rule)
val newChild1 = f(arg1.asInstanceOf[BaseType])
val newChild2 = f(arg2.asInstanceOf[BaseType])
if (!(newChild1 fastEquals arg1) || !(newChild2 fastEquals arg2)) {
changed = true
(newChild1, newChild2)
Expand Down
Loading