-
Notifications
You must be signed in to change notification settings - Fork 15.4k
KAFKA-15022: [2/N] introduce graph to compute min cost #13996
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
bea537e
12c18a9
f73051e
f26c021
1b5556d
4ca93ec
9e67923
738b345
a2889a4
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,377 @@ | ||
| /* | ||
| * Licensed to the Apache Software Foundation (ASF) under one or more | ||
| * contributor license agreements. See the NOTICE file distributed with | ||
| * this work for additional information regarding copyright ownership. | ||
| * The ASF licenses this file to You under the Apache License, Version 2.0 | ||
| * (the "License"); you may not use this file except in compliance with | ||
| * the License. You may obtain a copy of the License at | ||
| * | ||
| * http://www.apache.org/licenses/LICENSE-2.0 | ||
| * | ||
| * Unless required by applicable law or agreed to in writing, software | ||
| * distributed under the License is distributed on an "AS IS" BASIS, | ||
| * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| * See the License for the specific language governing permissions and | ||
| * limitations under the License. | ||
| */ | ||
| package org.apache.kafka.streams.processor.internals.assignment; | ||
|
|
||
| import java.util.HashMap; | ||
| import java.util.HashSet; | ||
| import java.util.Map; | ||
| import java.util.Map.Entry; | ||
| import java.util.Objects; | ||
| import java.util.Set; | ||
| import java.util.SortedMap; | ||
| import java.util.SortedSet; | ||
| import java.util.TreeMap; | ||
| import java.util.TreeSet; | ||
| import java.util.stream.Collectors; | ||
|
|
||
| public class Graph<V extends Comparable<V>> { | ||
| public class Edge { | ||
| final V destination; | ||
| final int capacity; | ||
| final int cost; | ||
| int residualFlow; | ||
| int flow; | ||
| Edge counterEdge; | ||
| boolean forwardEdge; | ||
|
|
||
| public Edge(final V destination, final int capacity, final int cost, final int residualFlow, final int flow) { | ||
| this(destination, capacity, cost, residualFlow, flow, true); | ||
| } | ||
|
|
||
| public Edge(final V destination, final int capacity, final int cost, final int residualFlow, final int flow, | ||
| final boolean forwardEdge) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. nit: formatting (if it does not fit in one line, we should move each parameter into it's one line to simplify reading) |
||
| Objects.requireNonNull(destination); | ||
| if (capacity < 0) { | ||
| throw new IllegalArgumentException("Edge capacity cannot be negative"); | ||
| } | ||
| if (flow > capacity) { | ||
| throw new IllegalArgumentException(String.format("Edge flow %d cannot exceed capacity %d", | ||
| flow, capacity)); | ||
| } | ||
|
|
||
| this.destination = destination; | ||
| this.capacity = capacity; | ||
| this.cost = cost; | ||
| this.residualFlow = residualFlow; | ||
| this.flow = flow; | ||
| this.forwardEdge = forwardEdge; | ||
| } | ||
|
|
||
| @Override | ||
| public boolean equals(final Object other) { | ||
| if (this == other) { | ||
| return true; | ||
| } | ||
| if (other == null || other.getClass() != getClass()) { | ||
| return false; | ||
| } | ||
|
|
||
| final Graph<?>.Edge otherEdge = (Graph<?>.Edge) other; | ||
|
|
||
| return destination.equals(otherEdge.destination) && capacity == otherEdge.capacity | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. nit formatting: |
||
| && cost == otherEdge.cost && residualFlow == otherEdge.residualFlow && flow == otherEdge.flow | ||
| && forwardEdge == otherEdge.forwardEdge; | ||
| } | ||
|
|
||
| @Override | ||
| public int hashCode() { | ||
| return Objects.hash(destination, capacity, cost, residualFlow, flow, forwardEdge); | ||
| } | ||
|
|
||
| @Override | ||
| public String toString() { | ||
| return "{destination= " + destination + ", capacity=" + capacity + ", cost=" + cost | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| + ", residualFlow=" + residualFlow + ", flow=" + flow + ", forwardEdge=" + forwardEdge; | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. missing Should we also switch to one line per parameter we print to simplify reading? |
||
| } | ||
| } | ||
|
|
||
| private final SortedMap<V, SortedMap<V, Edge>> adjList = new TreeMap<>(); | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Why do we use nested
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. It's for getting the edge between two nodes efficiently. e.g. to get edge between node 0 and node 1, we can do |
||
| private final SortedSet<V> nodes = new TreeSet<>(); | ||
| private final boolean isResidualGraph; | ||
| private V sourceNode, sinkNode; | ||
|
|
||
| public Graph() { | ||
| this(false); | ||
| } | ||
|
|
||
| private Graph(final boolean isResidualGraph) { | ||
| this.isResidualGraph = isResidualGraph; | ||
| } | ||
|
|
||
| public void addEdge(final V u, final V v, final int capacity, final int cost, final int flow) { | ||
| addEdge(u, new Edge(v, capacity, cost, capacity - flow, flow)); | ||
| } | ||
|
|
||
| public Set<V> nodes() { | ||
| return nodes; | ||
| } | ||
|
|
||
| public Map<V, Edge> edges(final V node) { | ||
| return adjList.get(node); | ||
| } | ||
|
|
||
| public boolean isResidualGraph() { | ||
| return isResidualGraph; | ||
| } | ||
|
|
||
| public void setSourceNode(final V node) { | ||
| sourceNode = node; | ||
| } | ||
|
|
||
| public void setSinkNode(final V node) { | ||
| sinkNode = node; | ||
| } | ||
|
|
||
| public int totalCost() { | ||
| int totalCost = 0; | ||
| for (final Map.Entry<V, SortedMap<V, Edge>> nodeEdges : adjList.entrySet()) { | ||
| final SortedMap<V, Edge> edges = nodeEdges.getValue(); | ||
| for (final Entry<V, Edge> nodeEdge : edges.entrySet()) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. nit: we could just iterate over the |
||
| totalCost += nodeEdge.getValue().cost * nodeEdge.getValue().flow; | ||
| } | ||
| } | ||
| return totalCost; | ||
| } | ||
|
|
||
| private void addEdge(final V u, final Edge edge) { | ||
| if (!isResidualGraph) { | ||
| // Check if there's already an edge from u to v | ||
| final Map<V, Edge> edgeMap = adjList.get(edge.destination); | ||
| if (edgeMap != null && edgeMap.containsKey(u)) { | ||
| throw new IllegalArgumentException( | ||
| "There is already an edge from " + edge.destination | ||
| + " to " + u + ". Can not add an edge from " + u + " to " + edge.destination | ||
| + " since there will create a cycle between two nodes"); | ||
| } | ||
| } | ||
|
|
||
| adjList.computeIfAbsent(u, set -> new TreeMap<>()).put(edge.destination, edge); | ||
| nodes.add(u); | ||
| nodes.add(edge.destination); | ||
| } | ||
|
|
||
| /** | ||
| * Get residual graph of this graph. | ||
| * Residual graph definition: | ||
| * If there is an edge in original graph from u to v with capacity c, cost w and flow f, | ||
| * then in the new graph there are two edges e1 and e2. e1 is from u to v with capacity c - f, | ||
| * cost w and flow f. e2 is from v to u with capacity f, cost -w and flow 0. | ||
| * | ||
| * @return Residual graph | ||
| */ | ||
| public Graph<V> residualGraph() { | ||
| if (isResidualGraph) { | ||
| return this; | ||
| } | ||
|
|
||
| final Graph<V> residualGraph = new Graph<>(true); | ||
| for (final Map.Entry<V, SortedMap<V, Edge>> nodeEdges : adjList.entrySet()) { | ||
| final V node = nodeEdges.getKey(); | ||
| final SortedMap<V, Edge> edges = nodeEdges.getValue(); | ||
| for (final Entry<V, Edge> nodeEdge : edges.entrySet()) { | ||
| final Edge edge = nodeEdge.getValue(); | ||
| final Edge forwardEdge = new Edge(edge.destination, edge.capacity, edge.cost, edge.capacity - edge.flow, edge.flow); | ||
| final Edge backwardEdge = new Edge(node, edge.capacity, edge.cost * -1, edge.flow, 0, false); | ||
| forwardEdge.counterEdge = backwardEdge; | ||
| backwardEdge.counterEdge = forwardEdge; | ||
| residualGraph.addEdge(node, forwardEdge); | ||
| residualGraph.addEdge(edge.destination, backwardEdge); | ||
| } | ||
| } | ||
| return residualGraph; | ||
| } | ||
|
|
||
| /** | ||
| * Solve min cost flow with cycle canceling algorithm. | ||
| */ | ||
| public void solveMinCostFlow() { | ||
| validateMinCostGraph(); | ||
| final Graph<V> residualGraph = residualGraph(); | ||
| residualGraph.cancelNegativeCycles(); | ||
|
|
||
| for (final Entry<V, SortedMap<V, Edge>> nodeEdges : adjList.entrySet()) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Not sure what this part does.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This is to update the original graph edge to the computed flow. Basically iterating all edges in original graph and change the flow to what's in residualGraph's corresponding edge. |
||
| final V node = nodeEdges.getKey(); | ||
| for (final Entry<V, Edge> nodeEdge : nodeEdges.getValue().entrySet()) { | ||
| final V destination = nodeEdge.getKey(); | ||
| final Edge edge = nodeEdge.getValue(); | ||
| final Edge residualEdge = residualGraph.adjList.get(node).get(destination); | ||
| edge.flow = residualEdge.flow; | ||
| edge.residualFlow = residualEdge.residualFlow; | ||
| } | ||
| } | ||
| } | ||
|
|
||
| private void populateInOutFlow(final Map<V, Long> inFlow, final Map<V, Long> outFlow) { | ||
| for (final Entry<V, SortedMap<V, Edge>> nodeEdges : adjList.entrySet()) { | ||
| final V node = nodeEdges.getKey(); | ||
| if (node.equals(sinkNode)) { | ||
| throw new IllegalStateException("Sink node " + sinkNode + " shouldn't have output"); | ||
| } | ||
| for (final Entry<V, Edge> nodeEdge : nodeEdges.getValue().entrySet()) { | ||
| final V destination = nodeEdge.getKey(); | ||
| if (destination.equals(sourceNode)) { | ||
| throw new IllegalStateException("Source node " + sourceNode + " shouldn't have input " + node); | ||
| } | ||
| final Edge edge = nodeEdge.getValue(); | ||
| Long count = outFlow.get(node); | ||
| if (count == null) { | ||
| outFlow.put(node, (long) edge.flow); | ||
| } else { | ||
| outFlow.put(node, count + edge.flow); | ||
| } | ||
|
|
||
| count = inFlow.get(destination); | ||
| if (count == null) { | ||
| inFlow.put(destination, (long) edge.flow); | ||
| } else { | ||
| inFlow.put(destination, count + edge.flow); | ||
| } | ||
| } | ||
| } | ||
| } | ||
|
|
||
| private void validateMinCostGraph() { | ||
| if (isResidualGraph) { | ||
| throw new IllegalStateException("Should not be residual graph to solve min cost flow"); | ||
| } | ||
|
|
||
| /* | ||
| Check provided flow satisfying below constraints: | ||
| 1. Input flow and output flow for each node should be the same except for source and destination node | ||
| 2. Output flow of source and input flow of destination should be the same | ||
| */ | ||
|
|
||
| final Map<V, Long> inFlow = new HashMap<>(); | ||
| final Map<V, Long> outFlow = new HashMap<>(); | ||
| populateInOutFlow(inFlow, outFlow); | ||
|
|
||
| for (final Entry<V, Long> in : inFlow.entrySet()) { | ||
| if (in.getKey().equals(sourceNode) || in.getKey().equals(sinkNode)) { | ||
| continue; | ||
| } | ||
| final Long out = outFlow.get(in.getKey()); | ||
| if (!Objects.equals(in.getValue(), out)) { | ||
| throw new IllegalStateException("Input flow for node " + in.getKey() + " is " + | ||
| in.getValue() + " which doesn't match output flow " + out); | ||
| } | ||
| } | ||
|
|
||
| final Long sourceOutput = outFlow.get(sourceNode); | ||
| final Long sinkInput = inFlow.get(sinkNode); | ||
| if (!Objects.equals(sourceOutput, sinkInput)) { | ||
| throw new IllegalStateException("Output flow for source " + sourceNode + " is " + sourceOutput | ||
| + " which doesn't match input flow " + sinkInput + " for sink " + sinkNode); | ||
| } | ||
| } | ||
|
|
||
| private void cancelNegativeCycles() { | ||
| if (!isResidualGraph) { | ||
| throw new IllegalStateException("Should be residual graph to cancel negative cycles"); | ||
| } | ||
| boolean cyclePossible = true; | ||
| while (cyclePossible) { | ||
| cyclePossible = false; | ||
| for (final V node : nodes) { | ||
| final Map<V, V> parentNodes = new HashMap<>(); | ||
| final Map<V, Edge> parentEdges = new HashMap<>(); | ||
| final V possibleNodeInCycle = detectNegativeCycles(node, parentNodes, parentEdges); | ||
|
|
||
| if (possibleNodeInCycle == null) { | ||
| continue; | ||
| } | ||
|
|
||
| final Set<V> visited = new HashSet<>(); | ||
| V nodeInCycle = possibleNodeInCycle; | ||
| while (!visited.contains(nodeInCycle)) { | ||
| visited.add(nodeInCycle); | ||
| nodeInCycle = parentNodes.get(nodeInCycle); | ||
| } | ||
|
|
||
| cyclePossible = true; | ||
| cancelNegativeCycle(nodeInCycle, parentNodes, parentEdges); | ||
| } | ||
| } | ||
| } | ||
|
|
||
| private void cancelNegativeCycle(final V nodeInCycle, final Map<V, V> parentNodes, final Map<V, Edge> parentEdges) { | ||
| // Start from parentNode since nodeInCyle is used as exit condition in below loops | ||
| final V parentNode = parentNodes.get(nodeInCycle); | ||
| Edge parentEdge = parentEdges.get(nodeInCycle); | ||
|
|
||
| // Find max possible negative flow | ||
| int possibleFlow = parentEdge.residualFlow; | ||
| for (V curNode = parentNode; curNode != nodeInCycle; curNode = parentNodes.get(curNode)) { | ||
| parentEdge = parentEdges.get(curNode); | ||
| possibleFlow = Math.min(possibleFlow, parentEdge.residualFlow); | ||
| } | ||
|
|
||
| // Update graph by removing negative flow | ||
| parentEdge = parentEdges.get(nodeInCycle); | ||
| Edge counterEdge = parentEdge.counterEdge; | ||
| parentEdge.residualFlow -= possibleFlow; | ||
| if (parentEdge.forwardEdge) { | ||
| parentEdge.flow += possibleFlow; | ||
| } | ||
| counterEdge.residualFlow += possibleFlow; | ||
| if (counterEdge.forwardEdge && counterEdge.flow >= possibleFlow) { | ||
| counterEdge.flow -= possibleFlow; | ||
| } | ||
| for (V curNode = parentNode; curNode != nodeInCycle; curNode = parentNodes.get(curNode)) { | ||
| parentEdge = parentEdges.get(curNode); | ||
| counterEdge = parentEdge.counterEdge; | ||
| parentEdge.residualFlow -= possibleFlow; | ||
| if (parentEdge.forwardEdge) { | ||
| parentEdge.flow += possibleFlow; | ||
| } | ||
| counterEdge.residualFlow += possibleFlow; | ||
| if (counterEdge.forwardEdge && counterEdge.flow >= possibleFlow) { | ||
| counterEdge.flow -= possibleFlow; | ||
| } | ||
| } | ||
| } | ||
|
|
||
| /** | ||
| * Detect negative cycle using Bellman-ford shortest path algorithm. | ||
| * @param source Source node | ||
| * @param parentNodes Parent nodes to store negative cycle nodes | ||
| * @param parentEdges Parent edges to store negative cycle edges | ||
| * | ||
| * @return One node which can lead to negative cycle if exists or null if there's no negative cycle | ||
| */ | ||
| V detectNegativeCycles(final V source, final Map<V, V> parentNodes, final Map<V, Edge> parentEdges) { | ||
| // Use long to account for any overflow | ||
| final Map<V, Long> distance = nodes.stream().collect(Collectors.toMap(node -> node, node -> (long) Integer.MAX_VALUE)); | ||
| distance.put(source, 0L); | ||
| final int nodeCount = nodes.size(); | ||
|
|
||
| // Iterate nodeCount iterations since Bellaman-Ford will find shortest path in nodeCount - 1 | ||
| // iterations. If the distance can still be relaxed in nodeCount iteration, there's a negative | ||
| // cycle | ||
| for (int i = 0; i < nodeCount; i++) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. From https://en.wikipedia.org/wiki/Bellman%E2%80%93Ford_algorithm it says, so it N-1 times, but we do it N times. Why?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Good question. N-1 times will find shortest path, we need to do it one more time to see if the path can be even shorter which mean there's a negative cycle. That's why there's a check on line 335. |
||
| // Iterate through all edges | ||
| for (final Entry<V, SortedMap<V, Edge>> nodeEdges : adjList.entrySet()) { | ||
| final V u = nodeEdges.getKey(); | ||
| for (final Entry<V, Edge> nodeEdge : nodeEdges.getValue().entrySet()) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Looking into https://en.wikipedia.org/wiki/Bellman%E2%80%93Ford_algorithm that you linked from the KIP, the algorithms says "go over all edges" -- we need 2 for-loops to this but what is not totally obvious; might be worth to add a coment? |
||
| final Edge edge = nodeEdge.getValue(); | ||
| if (edge.residualFlow == 0) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Can you explain this condition?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This means their can't be a flow in this edge which essentially means these two nodes are disconnected. |
||
| continue; | ||
| } | ||
| final V v = edge.destination; | ||
| if (distance.get(v) > distance.get(u) + edge.cost) { | ||
| if (i == nodeCount - 1) { | ||
| return v; | ||
| } | ||
| distance.put(v, distance.get(u) + edge.cost); | ||
| parentNodes.put(v, u); | ||
| parentEdges.put(v, edge); | ||
| } | ||
| } | ||
| } | ||
| } | ||
| return null; | ||
| } | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Do we actually need this field?
If I read the code right, for a forward edge, it's always
capacity - flowand for a backward edge it's always0. So it seem redundant (and potentially error prone to store it expliclity)? -- Instead we could have aresidualFlow()method that compute it on-the-fly (we could also simplify the update logic when modifying flow as we only need to update theflowitself)? Or do I read the update logic insidecancelNegativeCycleincorrectly and those properties are not an invariant?There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
residualFlowfor backward edge is the flow in forward edge: https://github.com/apache/kafka/blob/trunk/streams/src/main/java/org/apache/kafka/streams/processor/internals/assignment/Graph.java#L178There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Ah. Parameter order... Thought it's
, flow, residualFlow,(not sure why you picked the "reverse" order)