diff --git a/packages/flutter_mozilla_components/lib/src/utils/ml/cluster.dart b/packages/flutter_mozilla_components/lib/src/utils/ml/cluster.dart index a3c04109..ae138e42 100644 --- a/packages/flutter_mozilla_components/lib/src/utils/ml/cluster.dart +++ b/packages/flutter_mozilla_components/lib/src/utils/ml/cluster.dart @@ -88,11 +88,10 @@ double _getCentroidInertia( // Compute sum of squared distances to centroid for (final index in cluster) { final point = embeddings[index]; - double distanceSquared = 0.0; for (int i = 0; i < dimensions; i++) { - distanceSquared += math.pow(point[i] - centroid[i], 2); + final diff = point[i] - centroid[i]; + totalDistance += diff * diff; } - totalDistance += distanceSquared; } } diff --git a/packages/flutter_mozilla_components/lib/src/utils/ml/cluster_algo.dart b/packages/flutter_mozilla_components/lib/src/utils/ml/cluster_algo.dart index f14de14e..e0fca89b 100644 --- a/packages/flutter_mozilla_components/lib/src/utils/ml/cluster_algo.dart +++ b/packages/flutter_mozilla_components/lib/src/utils/ml/cluster_algo.dart @@ -93,7 +93,7 @@ List> initializeCentroidsSorted({ } else { centerId = (randomFunc() * nSamples).floor(); } - centers[0] = List.from(X[centerId]); + centers[0] = X[centerId]; } else { centers[0] = vectorNormalize( vectorMean(anchorIndices.map((a) => X[a]).toList()), @@ -158,7 +158,7 @@ List> initializeCentroidsSorted({ sumOfDistances = candidatesSumOfDistances[bestCandidateIdx]; // Pick best candidate - centers[c] = List.from(X[bestCandidate]); + centers[c] = X[bestCandidate]; } return centers; @@ -192,7 +192,8 @@ double euclideanDistance( }) { double sum = 0; for (int i = 0; i < point1.length; i++) { - sum += math.pow(point1[i] - point2[i], 2); + final diff = point1[i] - point2[i]; + sum += diff * diff; } return squareResult ? sum : math.sqrt(sum); } @@ -231,7 +232,8 @@ List euclideanDistancesSquared( return X.map((row) { double distSq = 0; for (int i = 0; i < row.length; i++) { - distSq += math.pow(row[i] - point[i], 2); + final diff = row[i] - point[i]; + distSq += diff * diff; } return distSq; }).toList();