Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,6 @@
import java.util.Objects;
import java.util.Optional;
import java.util.Set;
import java.util.function.BiFunction;
import java.util.function.Function;
import java.util.stream.Collectors;
import java.util.stream.IntStream;
Expand Down Expand Up @@ -230,11 +229,6 @@ public TaskStateAssignment[] getDownstreamAssignments() {
return downstreamAssignments;
}

private static int getAssignmentIndex(
TaskStateAssignment[] assignments, TaskStateAssignment assignment) {
return Arrays.asList(assignments).indexOf(assignment);
}

public TaskStateAssignment[] getUpstreamAssignments() {
if (upstreamAssignments == null) {
upstreamAssignments =
Expand Down Expand Up @@ -272,12 +266,6 @@ public OperatorSubtaskState getSubtaskState(OperatorInstanceID instanceID) {
instanceID,
inputOperatorID,
getUpstreamAssignments(),
(assignment, recompute) -> {
int assignmentIndex =
getAssignmentIndex(
assignment.getDownstreamAssignments(), this);
return assignment.getOutputMapping(assignmentIndex, recompute);
},
inputSubtaskMappings,
this::getInputMapping,
true))
Expand Down Expand Up @@ -320,11 +308,6 @@ private InflightDataRescalingDescriptor computeOutputRescalingDescriptor(
instanceID,
outputOperatorID,
getDownstreamAssignments(),
(downstreamAssignment, recompute) -> {
int assignmentIndex =
getAssignmentIndex(downstreamAssignment.getUpstreamAssignments(), this);
return downstreamAssignment.getInputMapping(assignmentIndex, recompute);
},
outputSubtaskMappings,
this::getOutputMapping,
false);
Expand Down Expand Up @@ -355,7 +338,6 @@ private InflightDataRescalingDescriptor createRescalingDescriptor(
OperatorInstanceID instanceID,
OperatorID expectedOperatorID,
TaskStateAssignment[] connectedAssignments,
BiFunction<TaskStateAssignment, Boolean, SubtasksRescaleMapping> mappingRetriever,
Map<Integer, SubtasksRescaleMapping> subtaskGateOrPartitionMappings,
Function<Integer, SubtasksRescaleMapping> subtaskMappingCalculator,
boolean isInput) {
Expand All @@ -364,8 +346,11 @@ private InflightDataRescalingDescriptor createRescalingDescriptor(
}

SubtasksRescaleMapping[] rescaledChannelsMappings =
Arrays.stream(connectedAssignments)
.map(assignment -> mappingRetriever.apply(assignment, false))
IntStream.range(0, connectedAssignments.length)
.mapToObj(
index ->
getConnectedMapping(
isInput, index, connectedAssignments[index], false))
.toArray(SubtasksRescaleMapping[]::new);

// no state on input and output, especially for any aligned checkpoint
Expand All @@ -378,7 +363,6 @@ private InflightDataRescalingDescriptor createRescalingDescriptor(
createGateOrPartitionRescalingDescriptors(
instanceID,
connectedAssignments,
assignment -> mappingRetriever.apply(assignment, true),
subtaskGateOrPartitionMappings,
subtaskMappingCalculator,
rescaledChannelsMappings,
Expand All @@ -398,7 +382,6 @@ private InflightDataRescalingDescriptor createRescalingDescriptor(
createGateOrPartitionRescalingDescriptors(
OperatorInstanceID instanceID,
TaskStateAssignment[] connectedAssignments,
Function<TaskStateAssignment, SubtasksRescaleMapping> mappingCalculator,
Map<Integer, SubtasksRescaleMapping> subtaskGateOrPartitionMappings,
Function<Integer, SubtasksRescaleMapping> subtaskMappingCalculator,
SubtasksRescaleMapping[] rescaledChannelsMappings,
Expand All @@ -415,8 +398,11 @@ private InflightDataRescalingDescriptor createRescalingDescriptor(
Optional.ofNullable(rescaledChannelsMappings[partition])
.orElseGet(
() ->
mappingCalculator.apply(
connectedAssignment));
getConnectedMapping(
isInput,
partition,
connectedAssignment,
true));
SubtasksRescaleMapping subtaskMapping =
Optional.ofNullable(
subtaskGateOrPartitionMappings.get(partition))
Expand Down Expand Up @@ -485,6 +471,11 @@ private SubtasksRescaleMapping getOutputMapping(int assignmentIndex, boolean rec
}
}

private SubtasksRescaleMapping getOutputMapping(
IntermediateDataSetID resultId, boolean recompute) {
return getOutputMapping(findResultPartitionIndex(resultId), recompute);
}

private SubtasksRescaleMapping getInputMapping(int assignmentIndex, boolean recompute) {
SubtasksRescaleMapping mapping = inputSubtaskMappings.get(assignmentIndex);
if (recompute && mapping == null) {
Expand All @@ -494,6 +485,31 @@ private SubtasksRescaleMapping getInputMapping(int assignmentIndex, boolean reco
}
}

private SubtasksRescaleMapping getInputMapping(
IntermediateDataSetID resultId, boolean recompute) {
return getInputMapping(findInputGateIndex(resultId), recompute);
}

/**
* Resolves the mapping on {@code connectedAssignment} that corresponds to {@code index} on
* {@code this} assignment, disambiguating by {@link IntermediateDataSetID} rather than by array
* position (multiple edges can connect the same pair of job vertices).
*/
private SubtasksRescaleMapping getConnectedMapping(
boolean isInput,
int index,
TaskStateAssignment connectedAssignment,
boolean recompute) {
if (isInput) {
IntermediateDataSetID resultId = executionJobVertex.getInputs().get(index).getId();
return connectedAssignment.getOutputMapping(resultId, recompute);
} else {
IntermediateDataSetID resultId =
executionJobVertex.getProducedDataSets()[index].getId();
return connectedAssignment.getInputMapping(resultId, recompute);
}
}

public SubtasksRescaleMapping getOutputMapping(int partitionIndex) {
final TaskStateAssignment downstreamAssignment = getDownstreamAssignments()[partitionIndex];
final IntermediateResult output = executionJobVertex.getProducedDataSets()[partitionIndex];
Expand Down Expand Up @@ -547,12 +563,8 @@ public boolean hasInFlightDataForInputGate(int gateIndex) {
if (upstreamAssignment != null && upstreamAssignment.hasOutputState()) {
IntermediateResult inputResult = executionJobVertex.getInputs().get(gateIndex);
IntermediateDataSetID resultId = inputResult.getId();
IntermediateResult[] producedDataSets = inputResult.getProducer().getProducedDataSets();
for (int i = 0; i < producedDataSets.length; i++) {
if (producedDataSets[i].getId().equals(resultId)) {
return upstreamAssignment.outputStatePartitions.contains(i);
}
}
return upstreamAssignment.outputStatePartitions.contains(
upstreamAssignment.findResultPartitionIndex(resultId));
}

return false;
Expand All @@ -571,12 +583,8 @@ public boolean hasInFlightDataForResultPartition(int partitionIndex) {
IntermediateResult producedResult =
executionJobVertex.getProducedDataSets()[partitionIndex];
IntermediateDataSetID resultId = producedResult.getId();
List<IntermediateResult> inputs = downstreamAssignment.executionJobVertex.getInputs();
for (int i = 0; i < inputs.size(); i++) {
if (inputs.get(i).getId().equals(resultId)) {
return downstreamAssignment.inputStateGates.contains(i);
}
}
return downstreamAssignment.inputStateGates.contains(
downstreamAssignment.findInputGateIndex(resultId));
}
return false;
}
Expand Down Expand Up @@ -642,15 +650,35 @@ private int findInputGateIdxForResultPartition(int partitionIndex) {

IntermediateResult producedResult =
executionJobVertex.getProducedDataSets()[partitionIndex];
IntermediateDataSetID resultId = producedResult.getId();
List<IntermediateResult> inputs = downstreamAssignment.executionJobVertex.getInputs();
return downstreamAssignment.findInputGateIndex(producedResult.getId());
}

private int findInputGateIndex(IntermediateDataSetID resultId) {
List<IntermediateResult> inputs = executionJobVertex.getInputs();
for (int i = 0; i < inputs.size(); i++) {
if (inputs.get(i).getId().equals(resultId)) {
return i;
}
}
throw new IllegalArgumentException(
"No channel rescaler found during rescaling of channel state");
"No input gate found for intermediate data set "
+ resultId
+ " in "
+ executionJobVertex.getName());
}

private int findResultPartitionIndex(IntermediateDataSetID resultId) {
IntermediateResult[] producedDataSets = executionJobVertex.getProducedDataSets();
for (int i = 0; i < producedDataSets.length; i++) {
if (producedDataSets[i].getId().equals(resultId)) {
return i;
}
}
throw new IllegalArgumentException(
"No result partition found for intermediate data set "
+ resultId
+ " in "
+ executionJobVertex.getName());
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -579,6 +579,72 @@ void testChannelStateAssignmentDownscalingTwoDifferentGates()
RESCALING))));
}

@Test
void testChannelStateAssignmentUsesResultIdForDuplicateJobVertexConnections()
throws JobException, JobExecutionException {
int oldParallelism = 3;
int newParallelism = 2;
JobVertex upstream = createJobVertex(new OperatorID(), newParallelism);
JobVertex downstream = createJobVertex(new OperatorID(), newParallelism);
OperatorID upstreamOperator = upstream.getOperatorIDs().get(0).getGeneratedOperatorID();
OperatorID downstreamOperator = downstream.getOperatorIDs().get(0).getGeneratedOperatorID();
Random random = new Random();

OperatorState upstreamState =
new OperatorState("", "", upstreamOperator, oldParallelism, MAX_P);
OperatorState downstreamState =
new OperatorState("", "", downstreamOperator, oldParallelism, MAX_P);
for (int i = 0; i < oldParallelism; i++) {
upstreamState.putState(
i,
OperatorSubtaskState.builder()
.setResultSubpartitionState(
new StateObjectCollection<>(
asList(
createNewResultSubpartitionStateHandle(
10, 0, random),
createNewResultSubpartitionStateHandle(
10, 1, random))))
.build());
downstreamState.putState(
i,
OperatorSubtaskState.builder()
.setInputChannelState(
new StateObjectCollection<>(
asList(
createNewInputChannelStateHandle(10, 0, random),
createNewInputChannelStateHandle(
10, 1, random))))
.build());
}
Map<OperatorID, OperatorState> states = new HashMap<>();
states.put(upstreamOperator, upstreamState);
states.put(downstreamOperator, downstreamState);

connectVertices(upstream, downstream, RANGE, RANGE);
connectVertices(upstream, downstream, ROUND_ROBIN, ROUND_ROBIN);

Map<OperatorID, ExecutionJobVertex> vertices = toExecutionVertices(upstream, downstream);

new StateAssignmentOperation(0, new HashSet<>(vertices.values()), states, false, false)
.assignStates();

InflightDataRescalingDescriptor outputDescriptor =
getAssignedState(vertices.get(upstreamOperator), upstreamOperator, 0)
.getOutputRescalingDescriptor();
InflightDataRescalingDescriptor inputDescriptor =
getAssignedState(vertices.get(downstreamOperator), downstreamOperator, 0)
.getInputRescalingDescriptor();
assertThat(outputDescriptor.getChannelMapping(0))
.isEqualTo(RANGE.getNewToOldSubtasksMapping(oldParallelism, newParallelism));
assertThat(outputDescriptor.getChannelMapping(1))
.isEqualTo(ROUND_ROBIN.getNewToOldSubtasksMapping(oldParallelism, newParallelism));
assertThat(inputDescriptor.getChannelMapping(0))
.isEqualTo(RANGE.getNewToOldSubtasksMapping(oldParallelism, newParallelism));
assertThat(inputDescriptor.getChannelMapping(1))
.isEqualTo(ROUND_ROBIN.getNewToOldSubtasksMapping(oldParallelism, newParallelism));
}

private InflightDataGateOrPartitionRescalingDescriptor gate(
int[] oldIndices,
RescaleMappings rescaleMapping,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -345,6 +345,15 @@ public static void waitForSubtasksToFinish(
/** Wait for one checkpoint with in-flight buffers. */
public static String waitForCheckpointWithInflightBuffers(JobID jobID, MiniCluster miniCluster)
throws Exception {
return waitForCheckpointWithInflightBuffers(jobID, miniCluster, 1);
}

/**
* Wait for at least {@code minCompletedCheckpoints} completed checkpoints and return the latest
* checkpoint with in-flight buffers.
*/
public static String waitForCheckpointWithInflightBuffers(
JobID jobID, MiniCluster miniCluster, long minCompletedCheckpoints) throws Exception {
CompletableFuture<String> checkpointPath = new CompletableFuture<>();
waitForCheckpoints(
jobID,
Expand All @@ -353,6 +362,10 @@ public static String waitForCheckpointWithInflightBuffers(JobID jobID, MiniClust
if (checkpointStatsSnapshot == null) {
return false;
}
if (checkpointStatsSnapshot.getCounts().getNumberOfCompletedCheckpoints()
< minCompletedCheckpoints) {
return false;
}
CompletedCheckpointStats latestCompletedCheckpoint =
checkpointStatsSnapshot.getHistory().getLatestCompletedCheckpoint();

Expand Down
Loading