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
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
34
35
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
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
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
161
162
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
201
202 }