View Javadoc
1   package net.bmahe.genetics4j.samples.clustering;
2   
3   import java.util.HashMap;
4   import java.util.HashSet;
5   import java.util.Map;
6   import java.util.Objects;
7   import java.util.Set;
8   
9   import org.apache.commons.lang3.Validate;
10  
11  import net.bmahe.genetics4j.core.Fitness;
12  
13  public class FitnessUtils {
14  
15  	// tag::a_i[]
16  	private static double aI(final double[][] data, final double[][] distances,
17  			final Map<Integer, Set<Integer>> clusterToMembers, final int clusterIndex, final int i) {
18  		Objects.requireNonNull(data);
19  		Objects.requireNonNull(distances);
20  		Objects.requireNonNull(clusterToMembers);
21  
22  		final var members = clusterToMembers.get(clusterIndex);
23  
24  		double sumDistances = 0.0;
25  		for (final int memberIndex : members) {
26  			if (memberIndex != i) {
27  				sumDistances += distances[i][memberIndex];
28  			}
29  		}
30  
31  		return sumDistances / ((double) members.size() - 1.0d);
32  	}
33  	// end::a_i[]
34  
35  	// tag::b_i[]
36  	private static double bI(final double[][] data, final double[][] distances,
37  			final Map<Integer, Set<Integer>> clusterToMembers, final int numClusters, final int clusterIndex,
38  			final int i) {
39  		Objects.requireNonNull(data);
40  		Objects.requireNonNull(distances);
41  		Objects.requireNonNull(clusterToMembers);
42  		Validate.isTrue(numClusters > 0);
43  		Validate.inclusiveBetween(0, numClusters - 1, clusterIndex);
44  
45  		double minMean = -1;
46  		for (int otherClusterIndex = 0; otherClusterIndex < numClusters; otherClusterIndex++) {
47  
48  			if (otherClusterIndex != clusterIndex) {
49  
50  				final var members = clusterToMembers.get(otherClusterIndex);
51  
52  				if (members != null && members.isEmpty() == false) {
53  					double sumDistances = 0.0;
54  					for (final int memberIndex : members) {
55  						sumDistances += distances[i][memberIndex];
56  					}
57  
58  					final double meanDistance = sumDistances / members.size();
59  
60  					if (minMean < 0 || meanDistance < minMean) {
61  						minMean = meanDistance;
62  					}
63  				}
64  			}
65  		}
66  
67  		return minMean;
68  	}
69  	// end::b_i[]
70  
71  	public static int[] assignDataToClusters(final double[][] data, double[][] distances, final double[][] clusters) {
72  
73  		final double[] closestClusterDistance = new double[data.length];
74  		final int[] closestClusterIndex = new int[data.length];
75  
76  		for (int i = 0; i < data.length; i++) {
77  			closestClusterIndex[i] = -1;
78  
79  			final double dataX = data[i][0];
80  			final double dataY = data[i][1];
81  
82  			for (int clusterIndex = 0; clusterIndex < clusters.length; clusterIndex++) {
83  				final double clusterX = clusters[clusterIndex][0];
84  				final double clusterY = clusters[clusterIndex][1];
85  
86  				final double distance = Math
87  						.sqrt(((clusterX - dataX) * (clusterX - dataX)) + ((clusterY - dataY) * (clusterY - dataY)));
88  
89  				if (closestClusterIndex[i] == -1 || distance < closestClusterDistance[i]) {
90  					closestClusterIndex[i] = clusterIndex;
91  					closestClusterDistance[i] = distance;
92  				}
93  			}
94  		}
95  
96  		return closestClusterIndex;
97  	}
98  
99  	public static double computeSilhouetteScore(final double[][] data, double[][] distances, final int numClusters,
100 			final Map<Integer, Set<Integer>> clusterToMembers, final int[] closestClusterIndex, final int i) {
101 
102 		final int clusterI = closestClusterIndex[i];
103 
104 		double silhouetteScore = 0.0D;
105 		if (clusterToMembers.getOrDefault(clusterI, Set.of()).size() > 1) {
106 			final double ai = aI(data, distances, clusterToMembers, clusterI, i);
107 			final double bi = bI(data, distances, clusterToMembers, numClusters, clusterI, i);
108 
109 			silhouetteScore = (bi - ai) / Math.max(ai, bi);
110 		}
111 
112 		return silhouetteScore;
113 	}
114 
115 	public static double computeSumSquaredErrors(final double[][] data, double[][] distances, final double[][] clusters,
116 			final Map<Integer, Set<Integer>> clusterToMembers, final int[] closestClusterIndex) {
117 
118 		double sumSquareErrors = 0.0D;
119 		for (int i = 0; i < data.length; i++) {
120 			final double[] cluster = clusters[closestClusterIndex[i]];
121 
122 			sumSquareErrors += (cluster[0] - data[i][0]) * (cluster[0] - data[i][0]);
123 			sumSquareErrors += (cluster[1] - data[i][1]) * (cluster[1] - data[i][1]);
124 		}
125 
126 		return sumSquareErrors;
127 	}
128 
129 	// tag::fitness[]
130 	public static Fitness<Double> computeFitness(final int numDataPoints, final double[][] data, double[][] distances,
131 			final int numClusters) {
132 		Objects.requireNonNull(data);
133 		Objects.requireNonNull(distances);
134 		Validate.isTrue(numDataPoints > 0);
135 		Validate.isTrue(numDataPoints == data.length);
136 		Validate.isTrue(numDataPoints == distances.length);
137 		Validate.isTrue(numClusters > 0);
138 
139 		return genoType -> {
140 
141 			final double[][] clusters = PhenotypeUtils.toPhenotype(genoType);
142 
143 			final int[] closestClusterIndex = assignDataToClusters(data, distances, clusters);
144 
145 			final Map<Integer, Set<Integer>> clusterToMembers = new HashMap<>();
146 
147 			for (int i = 0; i < numDataPoints; i++) {
148 				final var members = clusterToMembers.computeIfAbsent(closestClusterIndex[i], k -> new HashSet<>());
149 				members.add(i);
150 			}
151 
152 			double sum_si = 0.0;
153 			for (int i = 0; i < numDataPoints; i++) {
154 				sum_si += computeSilhouetteScore(data, distances, numClusters, clusterToMembers, closestClusterIndex, i);
155 			}
156 
157 			return sum_si;
158 		};
159 	}
160 	// end::fitness[]
161 
162 	// Copy/pasted for the Clustering doc tag::fitness_with_sse[]
163 	public static Fitness<Double> computeFitnessWithSSE(final int numDataPoints, final double[][] data,
164 			double[][] distances, final int numClusters) {
165 		Objects.requireNonNull(data);
166 		Objects.requireNonNull(distances);
167 		Validate.isTrue(numDataPoints > 0);
168 		Validate.isTrue(numDataPoints == data.length);
169 		Validate.isTrue(numDataPoints == distances.length);
170 		Validate.isTrue(numClusters > 0);
171 
172 		return genoType -> {
173 
174 			final double[][] clusters = PhenotypeUtils.toPhenotype(genoType);
175 
176 			final int[] closestClusterIndex = assignDataToClusters(data, distances, clusters);
177 
178 			final Map<Integer, Set<Integer>> clusterToMembers = new HashMap<>();
179 
180 			for (int i = 0; i < numDataPoints; i++) {
181 				final var members = clusterToMembers.computeIfAbsent(closestClusterIndex[i], k -> new HashSet<>());
182 				members.add(i);
183 			}
184 
185 			double sum_si = 0.0;
186 			for (int i = 0; i < numDataPoints; i++) {
187 				sum_si += computeSilhouetteScore(data, distances, numClusters, clusterToMembers, closestClusterIndex, i);
188 			}
189 
190 			final double sumSquaredError = computeSumSquaredErrors(
191 					data,
192 						distances,
193 						clusters,
194 						clusterToMembers,
195 						closestClusterIndex);
196 
197 			return sum_si + (1.0 / sumSquaredError);
198 		};
199 	}
200 	// end::fitness_with_sse[]
201 
202 }