Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Next Next commit
AggregateMessages: allow multiple calls to sendToSrc and sendToDst me…
…thods as well as multiple aggregation functions
  • Loading branch information
estebandonato committed Apr 24, 2017
commit a0c302b3d38cf5f848ee3805ce4fc8b9f475f580
47 changes: 30 additions & 17 deletions src/main/scala/org/graphframes/lib/AggregateMessages.scala
Original file line number Diff line number Diff line change
Expand Up @@ -18,9 +18,8 @@
package org.graphframes.lib

import org.apache.spark.sql.SQLHelpers.expr
import org.apache.spark.sql.functions.col
import org.apache.spark.sql.functions._
import org.apache.spark.sql.{Column, DataFrame}

import org.graphframes.{GraphFrame, Logging}

/**
Expand All @@ -33,14 +32,15 @@ import org.graphframes.{GraphFrame, Logging}
* triplet
* - `AggregateMessages.sendToDst()` sends a message to the destination vertex of each
* triplet
* - `AggregateMessages.agg` specifies an aggregation function for aggregating the
* messages sent to each vertex. It also runs the aggregation, computing a DataFrame
* with one row for each vertex which receives > 0 messages. The DataFrame has 2 columns:
* - `AggregateMessages.agg` specifies a series of aggregation functions for aggregating the
* messages sent to each vertex. It also runs the aggregations, computing a DataFrame
* with one row for each vertex which receives > 0 messages. The DataFrame has the following
* columns:
* - vertex column ID (named [[GraphFrame.ID]])
* - aggregate from messages sent to vertex (with the name given to the `Column` specified
* - each aggregate from messages sent to vertex (with the names given to the `Column` specified
* in `AggregateMessages.agg()`)
*
* When specifying the messages and aggregation function, the user may reference columns using:
* When specifying the messages and aggregation functions, the user may reference columns using:
* - [[AggregateMessages.src]]: column for source vertex of edge
* - [[AggregateMessages.edge]]: column for edge
* - [[AggregateMessages.dst]]: column for destination vertex of edge
Expand All @@ -62,22 +62,22 @@ class AggregateMessages private[graphframes] (private val g: GraphFrame)

import org.graphframes.GraphFrame.{DST, ID, SRC}

private var msgToSrc: Option[Column] = None
private var msgToSrc: Seq[Column] = Vector()

/** Send message to source vertex */
def sendToSrc(value: Column): this.type = {
msgToSrc = Some(value)
msgToSrc :+= value
this
}

/** Send message to source vertex, specifying SQL expression as a String */
def sendToSrc(value: String): this.type = sendToSrc(expr(value))

private var msgToDst: Option[Column] = None
private var msgToDst: Seq[Column] = Vector()

/** Send message to destination vertex */
def sendToDst(value: Column): this.type = {
msgToDst = Some(value)
msgToDst :+= value
this
}

Expand All @@ -92,6 +92,7 @@ class AggregateMessages private[graphframes] (private val g: GraphFrame)
* This returns a DataFrame with schema:
* - column "id": vertex ID
* - aggCol: aggregate result
* - aggCols: one column with the result of each additional defined aggregation
* If you need to join this with the original [[GraphFrame.vertices]], you can run an inner
* join of the form:
* {{{
Expand All @@ -100,18 +101,30 @@ class AggregateMessages private[graphframes] (private val g: GraphFrame)
* aggResult.join(g.vertices, ID)
* }}}
*/
def agg(aggCol: Column): DataFrame = {
def agg(aggCol: Column, aggCols: Column*): DataFrame = {
def removeColumnNamePrefix(columnName: String) = columnName match {
case cn if cn.startsWith(s"${GraphFrame.SRC}.") => cn.substring(s"${GraphFrame.SRC}.".length)
case cn if cn.startsWith(s"${GraphFrame.DST}.") => cn.substring(s"${GraphFrame.DST}.".length)
case cn if cn.startsWith(s"${GraphFrame.EDGE}.") => cn.substring(s"${GraphFrame.EDGE}.".length)
case cn => cn
}
def msgColumn(df: DataFrame) = df.columns.filter(_ != ID) match {
case Array(c) => df.withColumnRenamed(c, "MSG")
case columns => df.select(
df(ID),
struct(columns.map(c => df(s"`${c}`").as(removeColumnNamePrefix(c))) :_*).as("MSG"))
}
require(msgToSrc.nonEmpty || msgToDst.nonEmpty, s"To run GraphFrame.aggregateMessages," +
s" messages must be sent to src, dst, or both. Set using sendToSrc(), sendToDst().")
val triplets = g.triplets
val sentMsgsToSrc = msgToSrc.map { msg =>
val msgsToSrc = triplets.select(msg.as("MSG"), triplets(SRC)(ID).as(ID))
val sentMsgsToSrc = msgToSrc.headOption.map { _ =>
val msgsToSrc = msgColumn(triplets.select((msgToSrc :+ triplets(SRC)(ID).as(ID)): _*))
// Inner join: only send messages to vertices with edges
msgsToSrc.join(g.vertices, ID)
.select(msgsToSrc("MSG"), col(ID))
}
val sentMsgsToDst = msgToDst.map { msg =>
val msgsToDst = triplets.select(msg.as("MSG"), triplets(DST)(ID).as(ID))
val sentMsgsToDst = msgToDst.headOption.map { _ =>
val msgsToDst = msgColumn(triplets.select((msgToDst :+ triplets(DST)(ID).as(ID)): _*))
msgsToDst.join(g.vertices, ID)
.select(msgsToDst("MSG"), col(ID))
}
Expand All @@ -124,7 +137,7 @@ class AggregateMessages private[graphframes] (private val g: GraphFrame)
// Should never happen. Specify this case to avoid compilation warnings.
throw new RuntimeException("AggregateMessages: No messages were specified to be sent.")
}
unionedMsgs.groupBy(ID).agg(aggCol)
unionedMsgs.groupBy(ID).agg(aggCol, aggCols:_*)
}
}

Expand Down
36 changes: 33 additions & 3 deletions src/test/scala/org/graphframes/lib/AggregateMessagesSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,13 @@

package org.graphframes.lib

import scala.collection.mutable
import org.apache.spark.sql.Row

import scala.collection.mutable
import org.apache.spark.sql.functions._

import org.apache.spark.sql.types._
import org.graphframes.examples.Graphs
import org.graphframes.{GraphFrameTestSparkContext, SparkFunSuite}
import org.graphframes.{GraphFrame, GraphFrameTestSparkContext, SparkFunSuite, TestUtils}


class AggregateMessagesSuite extends SparkFunSuite with GraphFrameTestSparkContext {
Expand Down Expand Up @@ -60,4 +61,33 @@ class AggregateMessagesSuite extends SparkFunSuite with GraphFrameTestSparkConte
assert(aggMap(user) === trueAgg(user), s"Failure on user $user")
}
}

test("aggregateMessages with multiple message and aggregation columns") {
val AM = AggregateMessages
val vertices = sqlContext.createDataFrame(
List((1, 30, 3), (2, 40, 4), (3, 50, 5), (4, 60, 6))).toDF("id", "att1", "att2")
val edges = sqlContext.createDataFrame(List(1 -> 2, 2 -> 3, 1 -> 4)).toDF("src", "dst")
val expectedValues = Map(1 -> (100l, 5.0), 2 -> (80l, 4.0), 3 -> (40l, 4.0), 4 -> (30l, 3.0))

val g = GraphFrame(vertices, edges)
val agg = g.aggregateMessages
.sendToDst(AM.src("att1"))
.sendToSrc(AM.dst("att1"))
.sendToDst(AM.src("att2"))
.sendToSrc(AM.dst("att2"))
.agg(
sum(AM.msg("att1")).as("sum_att1"),
avg(AM.msg("att2")).as("avg_att2"))

//validate schema
assert(agg.schema.size === 3)
TestUtils.checkColumnType(agg.schema, "id", IntegerType)
TestUtils.checkColumnType(agg.schema, "sum_att1", LongType)
TestUtils.checkColumnType(agg.schema, "avg_att2", DoubleType)

//validate content
assert(agg.collect().map { case Row(id: Int, sumAtt1: Long, avgAtt2: Double) =>
id -> (sumAtt1, avgAtt2)
}.toMap === expectedValues)
}
}