Implemented sum, average and median aggregators for grouped double values

This commit is contained in:
Lennart Hensler committed 2014-02-26 14:54:11 +01:00
1 parent f39cf9a32d
commit b6726c4cc4
10 files changed
+506 -10

No files matched your search

@@ -1,4 +1,4 @@
package com.sap.sse.datamining.impl.components;
package com.sap.sse.datamining.impl.components.aggregators;
import static org.hamcrest.CoreMatchers.is;
import static org.junit.Assert.assertThat;
@@ -59,8 +59,6 @@ public class TestAbstractStoringParallelAggregationProcessor {
processElementAndVerifyThatItWasStored(processor, 7);
processor.finish();
ConcurrencyTestsUtil.sleepFor(100); //Giving the processor time to finish
assertThat("The receiver wasn't told to finish", receiverWasToldToFinish, is(true));
Integer expectedReceivedElement = 42 + 7;
assertThat(receivedElement, is(expectedReceivedElement));
@@ -71,5 +69,35 @@ public class TestAbstractStoringParallelAggregationProcessor {
ConcurrencyTestsUtil.sleepFor(100); //Giving the processor time to process the instructions
assertThat("The element store doesn't contain the previously processed element '" + element + "'", elementStore.contains(element), is(true));
}
@Test(timeout=5000)
public void testThatTheLockIsReleasedAfterStoringFailed() throws InterruptedException {
Processor<Integer> processor = new AbstractStoringParallelAggregationProcessor<Integer, Integer>(ConcurrencyTestsUtil.getExecutor(), receivers) {
@Override
protected void storeElement(Integer element) {
if (element < 0) {
throw new IllegalArgumentException("The element mustn't be negative");
}
elementStore.add(element);
}
@Override
protected Integer aggregateResult() {
Integer sum = 0;
for (Integer element : elementStore) {
sum += element;
}
return sum;
}
};
processor.onElement(-1);
processor.onElement(42);
processor.onElement(7);
processor.finish();
assertThat("The receiver wasn't told to finish", receiverWasToldToFinish, is(true));
Integer expectedReceivedElement = 42 + 7;
assertThat(receivedElement, is(expectedReceivedElement));
}
}
@@ -0,0 +1,194 @@
package com.sap.sse.datamining.impl.components.aggregators;
import static org.hamcrest.Matchers.is;
import static org.hamcrest.Matchers.notNullValue;
import static org.junit.Assert.assertThat;
import java.util.ArrayList;
import java.util.Collection;
import java.util.Collections;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Map.Entry;
import org.junit.Before;
import org.junit.Test;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.impl.components.GroupedDataEntry;
import com.sap.sse.datamining.shared.GroupKey;
import com.sap.sse.datamining.shared.impl.GenericGroupKey;
import com.sap.sse.datamining.test.util.ConcurrencyTestsUtil;
public class TestParallelDoubleAggregationProcessors {
private Collection<Processor<Map<GroupKey, Double>>> receivers;
private Map<GroupKey, Double> receivedAggregations;
@Test
public void testSumAggregationProcessor() throws InterruptedException {
Processor<GroupedDataEntry<Double>> sumAggregationProcessor = new ParallelGroupedDoubleDataSumAggregationProcessor(ConcurrencyTestsUtil.getExecutor(), receivers);
Collection<GroupedDataEntry<Double>> elements = createElements();
processElements(sumAggregationProcessor, elements);
sumAggregationProcessor.finish();
Map<GroupKey, Double> expectedReceivedAggregations = computeExpectedSumAggregations(elements);
verifyReceivedAggregations(expectedReceivedAggregations);
}
private Map<GroupKey, Double> computeExpectedSumAggregations(Collection<GroupedDataEntry<Double>> elements) {
Map<GroupKey, Double> expectedSumAggregations = new HashMap<>();
for (GroupedDataEntry<Double> element : elements) {
GroupKey key = element.getKey();
if (!expectedSumAggregations.containsKey(key)) {
expectedSumAggregations.put(key, 0.0);
}
Double currentValue = expectedSumAggregations.get(key);
expectedSumAggregations.put(key, currentValue + element.getDataEntry());
}
return expectedSumAggregations;
}
@Test
public void testAverageAggregationProcessor() throws InterruptedException {
Processor<GroupedDataEntry<Double>> averageAggregationProcessor = new ParallelGroupedDoubleDataAverageAggregationProcessor(ConcurrencyTestsUtil.getExecutor(), receivers);
Collection<GroupedDataEntry<Double>> elements = createElements();
processElements(averageAggregationProcessor, elements);
averageAggregationProcessor.finish();
Map<GroupKey, Double> expectedReceivedAggregations = computeExpectedAverageAggregations(elements);
verifyReceivedAggregations(expectedReceivedAggregations);
}
private Map<GroupKey, Double> computeExpectedAverageAggregations(Collection<GroupedDataEntry<Double>> elements) {
Map<GroupKey, Double> result = new HashMap<>();
Map<GroupKey, Double> sumAggregations = computeExpectedSumAggregations(elements);
Map<GroupKey, Double> elementAmountPerKey = countElementAmountPerKey(elements);
for (Entry<GroupKey, Double> sumAggregationEntry : sumAggregations.entrySet()) {
GroupKey key = sumAggregationEntry.getKey();
result.put(key, sumAggregationEntry.getValue() / elementAmountPerKey.get(key));
}
return result;
}
private Map<GroupKey, Double> countElementAmountPerKey(Collection<GroupedDataEntry<Double>> elements) {
Map<GroupKey, Double> elementAmountPerKey = new HashMap<>();
for (GroupedDataEntry<Double> element : elements) {
GroupKey key = element.getKey();
if (!elementAmountPerKey.containsKey(key)) {
elementAmountPerKey.put(key, 0.0);
}
Double currentAmount = elementAmountPerKey.get(key);
elementAmountPerKey.put(key, currentAmount + 1.0);
}
return elementAmountPerKey;
}
@Test
public void testMedianAggregationProcessor() throws InterruptedException {
Processor<GroupedDataEntry<Double>> medianAggregationProcessor = new ParallelGroupedDoubleDataMedianAggregationProcessor(ConcurrencyTestsUtil.getExecutor(), receivers);
Collection<GroupedDataEntry<Double>> elements = createElements();
processElements(medianAggregationProcessor, elements);
medianAggregationProcessor.finish();
Map<GroupKey, Double> expectedReceivedAggregations = computeExpectedMedianAggregations(elements);
verifyReceivedAggregations(expectedReceivedAggregations);
}
private Map<GroupKey, Double> computeExpectedMedianAggregations(Collection<GroupedDataEntry<Double>> elements) {
Map<GroupKey, Double> result = new HashMap<>();
Map<GroupKey, List<Double>> groupedValues = getGroupedValuesOf(elements);
for (Entry<GroupKey, List<Double>> groupedValuesEntry : groupedValues.entrySet()) {
result.put(groupedValuesEntry.getKey(), getMedianOf(groupedValuesEntry.getValue()));
}
return result;
}
private Map<GroupKey, List<Double>> getGroupedValuesOf(Collection<GroupedDataEntry<Double>> elements) {
Map<GroupKey, List<Double>> groupedValues = new HashMap<>();
for (GroupedDataEntry<Double> element : elements) {
GroupKey key = element.getKey();
if (!groupedValues.containsKey(key)) {
groupedValues.put(key, new ArrayList<Double>());
}
groupedValues.get(key).add(element.getDataEntry());
}
return groupedValues;
}
private Double getMedianOf(List<Double> values) {
Collections.sort(values);
if (listSizeIsEven(values)) {
int index1 = values.size() / 2;
int index2 = index1 + 1;
return (values.get(index1) + values.get(index2)) / 2;
} else {
int index = (values.size() + 1) / 2;
return values.get(index);
}
}
private boolean listSizeIsEven(List<Double> values) {
return values.size() % 2 == 0;
}
private Collection<GroupedDataEntry<Double>> createElements() {
Collection<GroupedDataEntry<Double>> elements = new ArrayList<>();
GroupKey firstGroupKey = new GenericGroupKey<Integer>(1);
elements.add(new GroupedDataEntry<Double>(firstGroupKey, 5.0));
elements.add(new GroupedDataEntry<Double>(firstGroupKey, 10.0));
elements.add(new GroupedDataEntry<Double>(firstGroupKey, 7.0));
GroupKey secondGroupKey = new GenericGroupKey<Integer>(2);
elements.add(new GroupedDataEntry<Double>(secondGroupKey, 5.0));
elements.add(new GroupedDataEntry<Double>(secondGroupKey, 3.0));
elements.add(new GroupedDataEntry<Double>(secondGroupKey, 7.0));
elements.add(new GroupedDataEntry<Double>(secondGroupKey, 7.0));
GroupKey thirdGroupKey = new GenericGroupKey<Integer>(3);
elements.add(new GroupedDataEntry<Double>(thirdGroupKey, 5.0));
elements.add(new GroupedDataEntry<Double>(thirdGroupKey, 5.0));
elements.add(new GroupedDataEntry<Double>(thirdGroupKey, 5.0));
return elements;
}
private void processElements(Processor<GroupedDataEntry<Double>> processor, Collection<GroupedDataEntry<Double>> elements) {
for (GroupedDataEntry<Double> element : elements) {
processor.onElement(element);
}
}
private void verifyReceivedAggregations(Map<GroupKey, Double> expectedReceivedAggregations) {
assertThat("No aggregation has been received.", receivedAggregations, notNullValue());
for (Entry<GroupKey, Double> expectedReceivedAggregationEntry : expectedReceivedAggregations.entrySet()) {
assertThat("The expected aggregation entry '" + expectedReceivedAggregationEntry + "' wasn't received.",
receivedAggregations.containsKey(expectedReceivedAggregationEntry.getKey()), is(true));
assertThat("The result for group '" + expectedReceivedAggregationEntry.getKey() + "' isn't correct.",
receivedAggregations.get(expectedReceivedAggregationEntry.getKey()), is(expectedReceivedAggregationEntry.getValue()));
}
}
@Before
public void initializeResultReceivers() {
Processor<Map<GroupKey, Double>> receiver = new Processor<Map<GroupKey,Double>>() {
@Override
public void onElement(Map<GroupKey, Double> element) {
receivedAggregations = element;
}
@Override
public void finish() throws InterruptedException {
}
};
receivers = new ArrayList<>();
receivers.add(receiver);
}
}
@@ -88,17 +88,21 @@ public abstract class AbstractPartitioningParallelProcessor<InputType, WorkingTy
@Override
public void finish() throws InterruptedException {
sleepUntilAllInstructionsFinished();
notifyResultReceiversToFinish();
}
protected void sleepUntilAllInstructionsFinished() throws InterruptedException {
while (areUnfinishedInstructionsLeft()) {
Thread.sleep(SLEEP_TIME_DURING_FINISHING);
}
notifyResultReceiversToFinish();
}
private boolean areUnfinishedInstructionsLeft() {
return unfinishedInstructionsCounter.getUnfinishedInstructionsAmount() > 0;
}
private void notifyResultReceiversToFinish() {
protected void notifyResultReceiversToFinish() {
for (Processor<ResultType> resultReceiver : getResultReceivers()) {
try {
resultReceiver.finish();
@@ -115,6 +119,7 @@ public abstract class AbstractPartitioningParallelProcessor<InputType, WorkingTy
private int unfinishedInstructionsAmount;
//TODO replace synchronized with ReentrantReadWriteLock
public synchronized void increment() {
unfinishedInstructionsAmount++;
}
@@ -25,4 +25,37 @@ public class GroupedDataEntry<DataType> {
return "[" + key + ", " + dataEntry + "]";
}
@Override
public int hashCode() {
final int prime = 31;
int result = 1;
result = prime * result + ((dataEntry == null) ? 0 : dataEntry.hashCode());
result = prime * result + ((key == null) ? 0 : key.hashCode());
return result;
}
@Override
public boolean equals(Object obj) {
if (this == obj)
return true;
if (obj == null)
return false;
if (getClass() != obj.getClass())
return false;
GroupedDataEntry<?> other = (GroupedDataEntry<?>) obj;
if (dataEntry == null) {
if (other.dataEntry != null)
return false;
} else if (!dataEntry.equals(other.dataEntry))
return false;
if (key == null) {
if (other.key != null)
return false;
} else if (!key.equals(other.key))
return false;
return true;
}
}
@@ -0,0 +1,47 @@
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.Collection;
import java.util.HashMap;
import java.util.Map;
import java.util.Map.Entry;
import java.util.concurrent.Executor;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.impl.components.GroupedDataEntry;
import com.sap.sse.datamining.shared.GroupKey;
public abstract class AbstractParallelGroupedDataSumAggregationProcessor<InputType, AggregatedType>
extends AbstractParallelSumAggregationProcessor<GroupedDataEntry<InputType>, Map<GroupKey, AggregatedType>> {
public AbstractParallelGroupedDataSumAggregationProcessor(Executor executor,
Collection<Processor<Map<GroupKey, AggregatedType>>> resultReceivers) {
super(executor, resultReceivers);
}
@Override
protected Map<GroupKey, AggregatedType> aggregateResult() {
Map<GroupKey, AggregatedType> result = new HashMap<>();
for (Entry<GroupedDataEntry<InputType>, Integer> elementAmountEntry : getElementAmountMap().entrySet()) {
InputType element = elementAmountEntry.getKey().getDataEntry();
Integer times = elementAmountEntry.getValue();
AggregatedType multipliedElementValue = multiply(element, times);
GroupKey groupKey = elementAmountEntry.getKey().getKey();
AggregatedType groupResult = result.get(groupKey);
result.put(groupKey, addToGroupResult(groupResult, multipliedElementValue));
}
return result;
}
private AggregatedType addToGroupResult(AggregatedType groupResult, AggregatedType multipliedElementValue) {
if (groupResult == null) {
return multipliedElementValue;
}
return add(groupResult, multipliedElementValue);
}
protected abstract AggregatedType multiply(InputType element, Integer times);
protected abstract AggregatedType add(AggregatedType firstSummand, AggregatedType secondSummand);
}
@@ -0,0 +1,33 @@
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.Collection;
import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.Executor;
import com.sap.sse.datamining.components.Processor;
public abstract class AbstractParallelSumAggregationProcessor<InputType, AggregatedType>
extends AbstractStoringParallelAggregationProcessor<InputType, AggregatedType> {
private Map<InputType, Integer> elementAmountMap;
public AbstractParallelSumAggregationProcessor(Executor executor, Collection<Processor<AggregatedType>> resultReceivers) {
super(executor, resultReceivers);
elementAmountMap = new HashMap<>();
}
@Override
protected void storeElement(InputType element) {
if (!elementAmountMap.containsKey(element)) {
elementAmountMap.put(element, 0);
}
Integer currentAmount = elementAmountMap.get(element);
elementAmountMap.put(element, currentAmount + 1);
}
protected Map<InputType, Integer> getElementAmountMap() {
return elementAmountMap;
}
}
@@ -1,16 +1,21 @@
package com.sap.sse.datamining.impl.components;
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.Collection;
import java.util.concurrent.Callable;
import java.util.concurrent.Executor;
import java.util.concurrent.locks.ReentrantReadWriteLock;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.impl.components.AbstractSimpleParallelProcessor;
public abstract class AbstractStoringParallelAggregationProcessor<InputType, AggregatedType> extends
AbstractSimpleParallelProcessor<InputType, AggregatedType> {
public abstract class AbstractStoringParallelAggregationProcessor<InputType, AggregatedType>
extends AbstractSimpleParallelProcessor<InputType, AggregatedType> {
private final ReentrantReadWriteLock storeLock;
public AbstractStoringParallelAggregationProcessor(Executor executor, Collection<Processor<AggregatedType>> resultReceivers) {
super(executor, resultReceivers);
storeLock = new ReentrantReadWriteLock();
}
@Override
@@ -18,18 +23,28 @@ public abstract class AbstractStoringParallelAggregationProcessor<InputType, Agg
return new Callable<AggregatedType>() {
@Override
public AggregatedType call() throws Exception {
storeElement(element);
storeLock.writeLock().lock();
try {
storeElement(element);
} finally {
storeLock.writeLock().unlock();
}
return AbstractStoringParallelAggregationProcessor.super.createInvalidResult();
}
};
}
/**
* Method to store the element in the concrete store. This method is only called in a way, that is thread safe, so
* that multiple threads can't corrupt the store.
*/
protected abstract void storeElement(InputType element);
@Override
public void finish() throws InterruptedException {
super.sleepUntilAllInstructionsFinished();
super.forwardResultToReceivers(aggregateResult());
super.finish();
super.notifyResultReceiversToFinish();
}
protected abstract AggregatedType aggregateResult();
@@ -0,0 +1,52 @@
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.Collection;
import java.util.HashMap;
import java.util.Map;
import java.util.Map.Entry;
import java.util.concurrent.Executor;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.impl.components.GroupedDataEntry;
import com.sap.sse.datamining.shared.GroupKey;
public class ParallelGroupedDoubleDataAverageAggregationProcessor extends
AbstractStoringParallelAggregationProcessor<GroupedDataEntry<Double>, Map<GroupKey, Double>> {
private final AbstractStoringParallelAggregationProcessor<GroupedDataEntry<Double>, Map<GroupKey, Double>> sumAggregationProcessor;
private final Map<GroupKey, Integer> elementAmountPerKey;
public ParallelGroupedDoubleDataAverageAggregationProcessor(Executor executor,
Collection<Processor<Map<GroupKey, Double>>> resultReceivers) {
super(executor, resultReceivers);
elementAmountPerKey = new HashMap<>();
sumAggregationProcessor = new ParallelGroupedDoubleDataSumAggregationProcessor(executor, resultReceivers);
}
@Override
protected void storeElement(GroupedDataEntry<Double> element) {
incrementElementAmount(element);
sumAggregationProcessor.storeElement(element);
}
private void incrementElementAmount(GroupedDataEntry<Double> element) {
GroupKey key = element.getKey();
if (!elementAmountPerKey.containsKey(key)) {
elementAmountPerKey.put(key, 0);
}
Integer currentAmount = elementAmountPerKey.get(key);
elementAmountPerKey.put(key, currentAmount + 1);
}
@Override
protected Map<GroupKey, Double> aggregateResult() {
Map<GroupKey, Double> result = new HashMap<>();
Map<GroupKey, Double> sumAggregation = sumAggregationProcessor.aggregateResult();
for (Entry<GroupKey, Double> sumAggregationEntry : sumAggregation.entrySet()) {
GroupKey key = sumAggregationEntry.getKey();
result.put(key, sumAggregationEntry.getValue() / elementAmountPerKey.get(key));
}
return result;
}
}
@@ -0,0 +1,61 @@
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.ArrayList;
import java.util.Collection;
import java.util.Collections;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Map.Entry;
import java.util.concurrent.Executor;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.impl.components.GroupedDataEntry;
import com.sap.sse.datamining.shared.GroupKey;
public class ParallelGroupedDoubleDataMedianAggregationProcessor
extends AbstractStoringParallelAggregationProcessor<GroupedDataEntry<Double>, Map<GroupKey, Double>> {
private Map<GroupKey, List<Double>> groupedValues;
public ParallelGroupedDoubleDataMedianAggregationProcessor(Executor executor,
Collection<Processor<Map<GroupKey, Double>>> resultReceivers) {
super(executor, resultReceivers);
groupedValues = new HashMap<>();
}
@Override
protected void storeElement(GroupedDataEntry<Double> element) {
GroupKey key = element.getKey();
if (!groupedValues.containsKey(key)) {
groupedValues.put(key, new ArrayList<Double>());
}
groupedValues.get(key).add(element.getDataEntry());
}
@Override
protected Map<GroupKey, Double> aggregateResult() {
Map<GroupKey, Double> result = new HashMap<>();
for (Entry<GroupKey, List<Double>> groupedValuesEntry : groupedValues.entrySet()) {
result.put(groupedValuesEntry.getKey(), getMedianOf(groupedValuesEntry.getValue()));
}
return result;
}
private Double getMedianOf(List<Double> values) {
Collections.sort(values);
if (listSizeIsEven(values)) {
int index1 = values.size() / 2;
int index2 = index1 + 1;
return (values.get(index1) + values.get(index2)) / 2;
} else {
int index = (values.size() + 1) / 2;
return values.get(index);
}
}
private boolean listSizeIsEven(List<Double> values) {
return values.size() % 2 == 0;
}
}
@@ -0,0 +1,28 @@
package com.sap.sse.datamining.impl.components.aggregators;
import java.util.Collection;
import java.util.Map;
import java.util.concurrent.Executor;
import com.sap.sse.datamining.components.Processor;
import com.sap.sse.datamining.shared.GroupKey;
public class ParallelGroupedDoubleDataSumAggregationProcessor
extends AbstractParallelGroupedDataSumAggregationProcessor<Double, Double> {
public ParallelGroupedDoubleDataSumAggregationProcessor(Executor executor,
Collection<Processor<Map<GroupKey, Double>>> resultReceivers) {
super(executor, resultReceivers);
}
@Override
protected Double multiply(Double element, Integer times) {
return element * times;
}
@Override
protected Double add(Double firstSummand, Double secondSummand) {
return firstSummand + secondSummand;
}
}