first implementation
This commit is contained in:
@@ -134,6 +134,45 @@ class GeckoInferenceRepository extends _$GeckoInferenceRepository {
|
|||||||
return neighbors;
|
return neighbors;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
Future<List<List<String>>?> suggestClusters({
|
||||||
|
required List<String> unassignedDocumentsInput,
|
||||||
|
}) async {
|
||||||
|
if (!ref.read(
|
||||||
|
generalSettingsWithDefaultsProvider.select(
|
||||||
|
(settings) => settings.enableLocalAiFeatures,
|
||||||
|
),
|
||||||
|
)) {
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
|
final processedDocuments = <String, String>{};
|
||||||
|
final unassignedDocumentsProcessed = unassignedDocumentsInput.map((doc) {
|
||||||
|
final processed = preprocessText(doc);
|
||||||
|
if (processed != doc) {
|
||||||
|
processedDocuments[processed] = doc;
|
||||||
|
}
|
||||||
|
|
||||||
|
return processed;
|
||||||
|
}).toList();
|
||||||
|
|
||||||
|
final embeddings = await generateDocumentEmbeddings(
|
||||||
|
unassignedDocumentsProcessed,
|
||||||
|
);
|
||||||
|
|
||||||
|
final clusters = embeddings.mapNotNull(
|
||||||
|
(embeddings) => clusterEmbeddings(embeddings: embeddings.values.toList())
|
||||||
|
.map(
|
||||||
|
(cluster) =>
|
||||||
|
cluster.map((i) => embeddings.keys.elementAt(i)).toList(),
|
||||||
|
)
|
||||||
|
.toList(),
|
||||||
|
);
|
||||||
|
|
||||||
|
print(clusters);
|
||||||
|
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
Future<Map<String, List<double>>?> generateDocumentEmbeddings(
|
Future<Map<String, List<double>>?> generateDocumentEmbeddings(
|
||||||
List<String> documents,
|
List<String> documents,
|
||||||
) async {
|
) async {
|
||||||
@@ -196,6 +235,32 @@ Future<String?> containerTopic(Ref ref, String containerId) async {
|
|||||||
return topic;
|
return topic;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@Riverpod()
|
||||||
|
Future<List<List<String>>?> suggestClusters(Ref ref) async {
|
||||||
|
final unassignedTitles = await ref.watch(
|
||||||
|
containerTabsDataProvider(null).selectAsync(
|
||||||
|
(tabData) => EquatableValue(
|
||||||
|
tabData
|
||||||
|
.where((tab) => tab.title.isNotEmpty)
|
||||||
|
.map((tab) => (tab.id, tab.title!))
|
||||||
|
.toSet(),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
);
|
||||||
|
|
||||||
|
if (unassignedTitles.value.isNotEmpty) {
|
||||||
|
await ref
|
||||||
|
.read(geckoInferenceRepositoryProvider.notifier)
|
||||||
|
.suggestClusters(
|
||||||
|
unassignedDocumentsInput: unassignedTitles.value
|
||||||
|
.map((tab) => tab.$2)
|
||||||
|
.toList(),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
return null;
|
||||||
|
}
|
||||||
|
|
||||||
@Riverpod()
|
@Riverpod()
|
||||||
Future<List<String>?> containerTabSuggestions(
|
Future<List<String>?> containerTabSuggestions(
|
||||||
Ref ref,
|
Ref ref,
|
||||||
|
|||||||
+42
-1
@@ -42,7 +42,7 @@ final class GeckoInferenceRepositoryProvider
|
|||||||
}
|
}
|
||||||
|
|
||||||
String _$geckoInferenceRepositoryHash() =>
|
String _$geckoInferenceRepositoryHash() =>
|
||||||
r'6099998a3f49fc4faf650878ab4afd6268c593e2';
|
r'1f08047bee912a475bfd0851f8a9868a41bca836';
|
||||||
|
|
||||||
abstract class _$GeckoInferenceRepository extends $Notifier<void> {
|
abstract class _$GeckoInferenceRepository extends $Notifier<void> {
|
||||||
void build();
|
void build();
|
||||||
@@ -132,6 +132,47 @@ final class ContainerTopicFamily extends $Family
|
|||||||
String toString() => r'containerTopicProvider';
|
String toString() => r'containerTopicProvider';
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@ProviderFor(suggestClusters)
|
||||||
|
const suggestClustersProvider = SuggestClustersProvider._();
|
||||||
|
|
||||||
|
final class SuggestClustersProvider
|
||||||
|
extends
|
||||||
|
$FunctionalProvider<
|
||||||
|
AsyncValue<List<List<String>>?>,
|
||||||
|
List<List<String>>?,
|
||||||
|
FutureOr<List<List<String>>?>
|
||||||
|
>
|
||||||
|
with
|
||||||
|
$FutureModifier<List<List<String>>?>,
|
||||||
|
$FutureProvider<List<List<String>>?> {
|
||||||
|
const SuggestClustersProvider._()
|
||||||
|
: super(
|
||||||
|
from: null,
|
||||||
|
argument: null,
|
||||||
|
retry: null,
|
||||||
|
name: r'suggestClustersProvider',
|
||||||
|
isAutoDispose: true,
|
||||||
|
dependencies: null,
|
||||||
|
$allTransitiveDependencies: null,
|
||||||
|
);
|
||||||
|
|
||||||
|
@override
|
||||||
|
String debugGetCreateSourceHash() => _$suggestClustersHash();
|
||||||
|
|
||||||
|
@$internal
|
||||||
|
@override
|
||||||
|
$FutureProviderElement<List<List<String>>?> $createElement(
|
||||||
|
$ProviderPointer pointer,
|
||||||
|
) => $FutureProviderElement(pointer);
|
||||||
|
|
||||||
|
@override
|
||||||
|
FutureOr<List<List<String>>?> create(Ref ref) {
|
||||||
|
return suggestClusters(ref);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
String _$suggestClustersHash() => r'9d212158b80ba53afe305db938ab399b1e931d88';
|
||||||
|
|
||||||
@ProviderFor(containerTabSuggestions)
|
@ProviderFor(containerTabSuggestions)
|
||||||
const containerTabSuggestionsProvider = ContainerTabSuggestionsFamily._();
|
const containerTabSuggestionsProvider = ContainerTabSuggestionsFamily._();
|
||||||
|
|
||||||
|
|||||||
@@ -1,2 +1,3 @@
|
|||||||
|
export 'src/utils/ml/cluster.dart';
|
||||||
export 'src/utils/ml/embedding_text_processing.dart';
|
export 'src/utils/ml/embedding_text_processing.dart';
|
||||||
export 'src/utils/ml/nearest_neighbor.dart';
|
export 'src/utils/ml/nearest_neighbor.dart';
|
||||||
|
|||||||
@@ -0,0 +1,163 @@
|
|||||||
|
import 'dart:math' as math;
|
||||||
|
|
||||||
|
import 'package:flutter_mozilla_components/src/utils/ml/cluster_algo.dart';
|
||||||
|
|
||||||
|
/// Clusters embeddings using K-means algorithm with configurable parameters
|
||||||
|
List<List<int>> clusterEmbeddings({
|
||||||
|
required List<List<double>> embeddings,
|
||||||
|
int? k,
|
||||||
|
double Function()? randomFunc,
|
||||||
|
int clusteringTriesPerK = 3,
|
||||||
|
}) {
|
||||||
|
k ??= 0;
|
||||||
|
int startK = k;
|
||||||
|
int endK = k + 1;
|
||||||
|
|
||||||
|
if (k == 0) {
|
||||||
|
startK = 2;
|
||||||
|
// Find a reasonable max # of clusters
|
||||||
|
endK =
|
||||||
|
math.min(
|
||||||
|
(math.log(embeddings.length) * 2.0).floor(),
|
||||||
|
embeddings.length,
|
||||||
|
) +
|
||||||
|
1;
|
||||||
|
}
|
||||||
|
|
||||||
|
List<List<int>>? bestResult;
|
||||||
|
double bestResultSilScore = -100.0;
|
||||||
|
|
||||||
|
for (int curK = startK; curK < endK; curK++) {
|
||||||
|
List<List<int>>? bestItemsForK;
|
||||||
|
double bestInertiaForK = 500000000000;
|
||||||
|
|
||||||
|
for (int j = 0; j < clusteringTriesPerK; j++) {
|
||||||
|
final allItems = kmeansPlusPlus(
|
||||||
|
data: embeddings,
|
||||||
|
k: curK,
|
||||||
|
randomFunc: randomFunc,
|
||||||
|
freezeAnchorsInZeroCluster: false, // Not needed since no anchors
|
||||||
|
);
|
||||||
|
|
||||||
|
final inertia = _getCentroidInertia(allItems, embeddings);
|
||||||
|
if (inertia < bestInertiaForK) {
|
||||||
|
bestInertiaForK = inertia;
|
||||||
|
bestItemsForK = allItems;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if (bestItemsForK != null) {
|
||||||
|
final silScores = silhouetteCoefficients(embeddings, bestItemsForK);
|
||||||
|
final avgSil = silScores.reduce((a, b) => a + b) / silScores.length;
|
||||||
|
|
||||||
|
if (avgSil > bestResultSilScore) {
|
||||||
|
bestResultSilScore = avgSil;
|
||||||
|
bestResult = bestItemsForK;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return bestResult ?? [];
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Computes the inertia (sum of squared distances to centroids) for clusters
|
||||||
|
double _getCentroidInertia(
|
||||||
|
List<List<int>> clusters,
|
||||||
|
List<List<double>> embeddings,
|
||||||
|
) {
|
||||||
|
double totalDistance = 0.0;
|
||||||
|
|
||||||
|
for (final cluster in clusters) {
|
||||||
|
if (cluster.isEmpty) continue;
|
||||||
|
|
||||||
|
// Compute centroid
|
||||||
|
final dimensions = embeddings[0].length;
|
||||||
|
final centroid = List<double>.filled(dimensions, 0.0);
|
||||||
|
|
||||||
|
for (final index in cluster) {
|
||||||
|
final point = embeddings[index];
|
||||||
|
for (int i = 0; i < dimensions; i++) {
|
||||||
|
centroid[i] += point[i];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int i = 0; i < dimensions; i++) {
|
||||||
|
centroid[i] /= cluster.length;
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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);
|
||||||
|
}
|
||||||
|
totalDistance += distanceSquared;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return totalDistance;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Computes silhouette coefficients for clusters
|
||||||
|
List<double> silhouetteCoefficients(
|
||||||
|
List<List<double>> embeddings,
|
||||||
|
List<List<int>> clusters,
|
||||||
|
) {
|
||||||
|
final silhouettes = <double>[];
|
||||||
|
|
||||||
|
for (int clusterIdx = 0; clusterIdx < clusters.length; clusterIdx++) {
|
||||||
|
final cluster = clusters[clusterIdx];
|
||||||
|
if (cluster.length <= 1) {
|
||||||
|
silhouettes.add(0.0);
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
double clusterSilhouette = 0.0;
|
||||||
|
|
||||||
|
for (final pointIdx in cluster) {
|
||||||
|
final point = embeddings[pointIdx];
|
||||||
|
|
||||||
|
// Compute average intra-cluster distance (a)
|
||||||
|
double intraDistance = 0.0;
|
||||||
|
for (final otherIdx in cluster) {
|
||||||
|
if (pointIdx != otherIdx) {
|
||||||
|
intraDistance += euclideanDistance(point, embeddings[otherIdx]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
final a = cluster.length > 1 ? intraDistance / (cluster.length - 1) : 0.0;
|
||||||
|
|
||||||
|
// Compute minimum average inter-cluster distance (b)
|
||||||
|
double minInterDistance = double.infinity;
|
||||||
|
for (
|
||||||
|
int otherClusterIdx = 0;
|
||||||
|
otherClusterIdx < clusters.length;
|
||||||
|
otherClusterIdx++
|
||||||
|
) {
|
||||||
|
if (otherClusterIdx == clusterIdx) continue;
|
||||||
|
|
||||||
|
final otherCluster = clusters[otherClusterIdx];
|
||||||
|
if (otherCluster.isEmpty) continue;
|
||||||
|
|
||||||
|
double interDistance = 0.0;
|
||||||
|
for (final otherIdx in otherCluster) {
|
||||||
|
interDistance += euclideanDistance(point, embeddings[otherIdx]);
|
||||||
|
}
|
||||||
|
final avgInterDistance = interDistance / otherCluster.length;
|
||||||
|
minInterDistance = math.min(minInterDistance, avgInterDistance);
|
||||||
|
}
|
||||||
|
|
||||||
|
final b = minInterDistance == double.infinity ? 0.0 : minInterDistance;
|
||||||
|
|
||||||
|
// Compute silhouette for this point
|
||||||
|
final silhouette = (a == 0.0 && b == 0.0)
|
||||||
|
? 0.0
|
||||||
|
: (b - a) / math.max(a, b);
|
||||||
|
clusterSilhouette += silhouette;
|
||||||
|
}
|
||||||
|
|
||||||
|
silhouettes.add(clusterSilhouette / cluster.length);
|
||||||
|
}
|
||||||
|
|
||||||
|
return silhouettes;
|
||||||
|
}
|
||||||
@@ -0,0 +1,296 @@
|
|||||||
|
import 'dart:math' as math;
|
||||||
|
|
||||||
|
/// Performs K-Means clustering with K-Means++ initialization of centroids.
|
||||||
|
/// If an existing cluster is specified with [anchorIndices], then one of the centroids
|
||||||
|
/// is the average of the embeddings of the items in the cluster.
|
||||||
|
List<List<int>> kmeansPlusPlus({
|
||||||
|
required List<List<double>> data,
|
||||||
|
required int k,
|
||||||
|
int? maxIterations,
|
||||||
|
double Function()? randomFunc,
|
||||||
|
List<int> anchorIndices = const [],
|
||||||
|
List<int> preassignedIndices = const [],
|
||||||
|
bool freezeAnchorsInZeroCluster = true,
|
||||||
|
}) {
|
||||||
|
randomFunc ??= math.Random().nextDouble;
|
||||||
|
maxIterations ??= 300;
|
||||||
|
|
||||||
|
final dimensions = data[0].length;
|
||||||
|
final centroids = initializeCentroidsSorted(
|
||||||
|
X: data,
|
||||||
|
k: k,
|
||||||
|
randomFunc: randomFunc,
|
||||||
|
anchorIndices: anchorIndices,
|
||||||
|
);
|
||||||
|
|
||||||
|
List<List<int>> resultClusters = [];
|
||||||
|
final anchorSet = Set<int>.from(anchorIndices);
|
||||||
|
final preassignedSet = Set<int>.from(preassignedIndices);
|
||||||
|
|
||||||
|
for (int iter = 0; iter < maxIterations; iter++) {
|
||||||
|
resultClusters = List.generate(k, (_) => <int>[]);
|
||||||
|
bool hasChanged = false;
|
||||||
|
|
||||||
|
// Assign each data point to the nearest centroid
|
||||||
|
for (int i = 0; i < data.length; i++) {
|
||||||
|
if (freezeAnchorsInZeroCluster && anchorSet.contains(i)) {
|
||||||
|
resultClusters[0].add(i);
|
||||||
|
} else {
|
||||||
|
final point = data[i];
|
||||||
|
final centroidIndex = getClosestCentroid(
|
||||||
|
point,
|
||||||
|
centroids,
|
||||||
|
excludeIndex: preassignedSet.contains(i) && !anchorSet.contains(i)
|
||||||
|
? 0
|
||||||
|
: -1,
|
||||||
|
);
|
||||||
|
resultClusters[centroidIndex].add(i);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Recompute centroids
|
||||||
|
for (int j = 0; j < k; j++) {
|
||||||
|
final newCentroid = _computeCentroid(resultClusters[j], data, dimensions);
|
||||||
|
if (!_arePointsEqual(centroids[j], newCentroid)) {
|
||||||
|
centroids[j] = newCentroid;
|
||||||
|
hasChanged = true;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop if centroids don't change
|
||||||
|
if (!hasChanged) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return resultClusters;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Kmeans++ initialization of centroids by finding ones farther than one another
|
||||||
|
List<List<double>> initializeCentroidsSorted({
|
||||||
|
required List<List<double>> X,
|
||||||
|
required int k,
|
||||||
|
required double Function() randomFunc,
|
||||||
|
int? numTrials,
|
||||||
|
List<int> anchorIndices = const [],
|
||||||
|
}) {
|
||||||
|
final nSamples = X.length;
|
||||||
|
final nFeatures = X[0].length;
|
||||||
|
final centers = List.generate(k, (_) => List<double>.filled(nFeatures, 0.0));
|
||||||
|
numTrials ??= 2 + (math.log(k) / math.ln10).floor();
|
||||||
|
|
||||||
|
void zeroOutAnchorItems(List<double> arr) {
|
||||||
|
for (final a in anchorIndices) {
|
||||||
|
arr[a] = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// First center is random unless anchor is specified
|
||||||
|
int centerId;
|
||||||
|
if (anchorIndices.length <= 1) {
|
||||||
|
if (anchorIndices.length == 1) {
|
||||||
|
centerId = anchorIndices[0];
|
||||||
|
} else {
|
||||||
|
centerId = (randomFunc() * nSamples).floor();
|
||||||
|
}
|
||||||
|
centers[0] = List<double>.from(X[centerId]);
|
||||||
|
} else {
|
||||||
|
centers[0] = vectorNormalize(
|
||||||
|
vectorMean(anchorIndices.map((a) => X[a]).toList()),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get closest distances
|
||||||
|
final closestDistSq = euclideanDistancesSquared(centers[0], X);
|
||||||
|
double sumOfDistances = closestDistSq.reduce((sum, dist) => sum + dist);
|
||||||
|
|
||||||
|
// Pick the remaining nClusters-1 points
|
||||||
|
for (int c = 1; c < k; c++) {
|
||||||
|
// Choose center candidates by sampling
|
||||||
|
final randVals = List.generate(
|
||||||
|
numTrials,
|
||||||
|
(_) => randomFunc() * sumOfDistances,
|
||||||
|
);
|
||||||
|
final closestDistSqForSamples = List<double>.from(closestDistSq);
|
||||||
|
|
||||||
|
if (anchorIndices.length > 1) {
|
||||||
|
zeroOutAnchorItems(closestDistSqForSamples);
|
||||||
|
}
|
||||||
|
|
||||||
|
final cumulativeProbs = stableCumsum(closestDistSqForSamples);
|
||||||
|
final candidateIds = randVals
|
||||||
|
.map((randVal) => searchSorted(cumulativeProbs, randVal))
|
||||||
|
.where((candId) => candId < nSamples)
|
||||||
|
.toList();
|
||||||
|
|
||||||
|
// Compute distances to center candidates
|
||||||
|
final distancesToCandidates = candidateIds
|
||||||
|
.map((candidateId) => euclideanDistancesSquared(X[candidateId], X))
|
||||||
|
.toList();
|
||||||
|
|
||||||
|
// Update closest distances squared and potential for each candidate
|
||||||
|
final candidatesSumOfDistances = distancesToCandidates.map((distances) {
|
||||||
|
double sum = 0;
|
||||||
|
for (int j = 0; j < closestDistSq.length; j++) {
|
||||||
|
sum += math.min(closestDistSq[j], distances[j]);
|
||||||
|
}
|
||||||
|
return sum;
|
||||||
|
}).toList();
|
||||||
|
|
||||||
|
// Choose the best candidate
|
||||||
|
int bestCandidateIdx = 0;
|
||||||
|
for (int i = 1; i < candidatesSumOfDistances.length; i++) {
|
||||||
|
if (candidatesSumOfDistances[i] <
|
||||||
|
candidatesSumOfDistances[bestCandidateIdx]) {
|
||||||
|
bestCandidateIdx = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
final bestCandidate = candidateIds[bestCandidateIdx];
|
||||||
|
|
||||||
|
// Update closest distance and potential
|
||||||
|
for (int i = 0; i < closestDistSq.length; i++) {
|
||||||
|
closestDistSq[i] = math.min(
|
||||||
|
closestDistSq[i],
|
||||||
|
distancesToCandidates[bestCandidateIdx][i],
|
||||||
|
);
|
||||||
|
}
|
||||||
|
sumOfDistances = candidatesSumOfDistances[bestCandidateIdx];
|
||||||
|
|
||||||
|
// Pick best candidate
|
||||||
|
centers[c] = List<double>.from(X[bestCandidate]);
|
||||||
|
}
|
||||||
|
|
||||||
|
return centers;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Helper function to find closest centroid for a given point
|
||||||
|
int getClosestCentroid(
|
||||||
|
List<double> point,
|
||||||
|
List<List<double>> centroids, {
|
||||||
|
int excludeIndex = -1,
|
||||||
|
}) {
|
||||||
|
double minDistance = double.infinity;
|
||||||
|
int closestIndex = -1;
|
||||||
|
|
||||||
|
for (int i = 0; i < centroids.length; i++) {
|
||||||
|
final distance = euclideanDistance(point, centroids[i]);
|
||||||
|
if (distance < minDistance && i != excludeIndex) {
|
||||||
|
minDistance = distance;
|
||||||
|
closestIndex = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return closestIndex;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Helper function to compute Euclidean distance between two points
|
||||||
|
double euclideanDistance(
|
||||||
|
List<double> point1,
|
||||||
|
List<double> point2, {
|
||||||
|
bool squareResult = false,
|
||||||
|
}) {
|
||||||
|
double sum = 0;
|
||||||
|
for (int i = 0; i < point1.length; i++) {
|
||||||
|
sum += math.pow(point1[i] - point2[i], 2);
|
||||||
|
}
|
||||||
|
return squareResult ? sum : math.sqrt(sum);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Normalize a vector
|
||||||
|
List<double> vectorNormalize(List<double> vector) {
|
||||||
|
final magnitude = math.sqrt(vector.fold(0.0, (sum, c) => sum + c * c));
|
||||||
|
if (magnitude == 0) {
|
||||||
|
return List<double>.filled(vector.length, 0.0);
|
||||||
|
}
|
||||||
|
return vector.map((c) => c / magnitude).toList();
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Find average of two vectors
|
||||||
|
List<double> vectorMean(List<List<double>> vectors) {
|
||||||
|
if (vectors.isEmpty) {
|
||||||
|
return [];
|
||||||
|
}
|
||||||
|
final dims = vectors[0].length;
|
||||||
|
final sum = List<double>.filled(dims, 0.0);
|
||||||
|
|
||||||
|
for (final vector in vectors) {
|
||||||
|
for (int i = 0; i < dims; i++) {
|
||||||
|
sum[i] += vector[i];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return sum.map((a) => a / vectors.length).toList();
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Find distances from a single point to a list of points
|
||||||
|
List<double> euclideanDistancesSquared(
|
||||||
|
List<double> point,
|
||||||
|
List<List<double>> X,
|
||||||
|
) {
|
||||||
|
return X.map((row) {
|
||||||
|
double distSq = 0;
|
||||||
|
for (int i = 0; i < row.length; i++) {
|
||||||
|
distSq += math.pow(row[i] - point[i], 2);
|
||||||
|
}
|
||||||
|
return distSq;
|
||||||
|
}).toList();
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Cumulative sum for an array
|
||||||
|
List<double> stableCumsum(List<double> arr) {
|
||||||
|
double sum = 0;
|
||||||
|
return arr.map((value) => sum += value).toList();
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Binary search
|
||||||
|
int searchSorted(List<double> arr, double val) {
|
||||||
|
int low = 0;
|
||||||
|
int high = arr.length;
|
||||||
|
|
||||||
|
while (low < high) {
|
||||||
|
final mid = (low + high) ~/ 2;
|
||||||
|
if (arr[mid] < val) {
|
||||||
|
low = mid + 1;
|
||||||
|
} else {
|
||||||
|
high = mid;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return low;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Compute centroid of a cluster by reference
|
||||||
|
List<double> _computeCentroid(
|
||||||
|
List<int> cluster,
|
||||||
|
List<List<double>> data,
|
||||||
|
int dimensions,
|
||||||
|
) {
|
||||||
|
if (cluster.isEmpty) {
|
||||||
|
return List<double>.filled(dimensions, 0.0);
|
||||||
|
}
|
||||||
|
|
||||||
|
final centroid = List<double>.filled(dimensions, 0.0);
|
||||||
|
|
||||||
|
for (final index in cluster) {
|
||||||
|
final point = data[index];
|
||||||
|
for (int i = 0; i < dimensions; i++) {
|
||||||
|
centroid[i] += point[i];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for (int p = 0; p < dimensions; p++) {
|
||||||
|
centroid[p] /= cluster.length;
|
||||||
|
}
|
||||||
|
|
||||||
|
return centroid;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Returns true if both points have equal values
|
||||||
|
bool _arePointsEqual(List<double> point1, List<double> point2) {
|
||||||
|
if (point1.length != point2.length) return false;
|
||||||
|
for (int i = 0; i < point1.length; i++) {
|
||||||
|
if (point1[i] != point2[i]) return false;
|
||||||
|
}
|
||||||
|
return true;
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user