memoryweave/rust/src/cluster.rs

180 lines
5.3 KiB
Rust
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// DBSCAN 聚类 — 纯 Rust 手写实现
// 深度整合步骤 1发现新主题和重复模式
use serde::{Deserialize, Serialize};
/// 聚类结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClusteringResult {
pub num_clusters: usize,
pub num_noise: usize,
pub cluster_sizes: Vec<usize>,
pub labels: Vec<i32>,
}
/// DBSCAN 聚类器
pub struct Clusterer {
epsilon: f64,
min_points: usize,
}
impl Clusterer {
pub fn new(epsilon: f64, min_points: usize) -> Self {
Self { epsilon, min_points }
}
pub fn cluster(
&self,
vectors: &[Vec<f32>],
_ids: &[String],
) -> Result<ClusteringResult, Box<dyn std::error::Error>> {
if vectors.is_empty() {
return Ok(ClusteringResult {
num_clusters: 0, num_noise: 0,
cluster_sizes: vec![],
labels: vec![],
});
}
let n = vectors.len();
let mut labels = vec![-1_i32; n];
let mut cluster_id = 0;
// 朴<> DBSCAN
for i in 0..n {
if labels[i] != -1 {
continue;
}
let neighbors = self.region_query(vectors, i);
if neighbors.len() < self.min_points {
labels[i] = -1; // noise for now
continue;
}
// 扩展簇
labels[i] = cluster_id;
let mut seed = neighbors;
let mut idx = 0;
while idx < seed.len() {
let j = seed[idx];
idx += 1;
if labels[j] == -1 {
labels[j] = cluster_id;
let nbrs = self.region_query(vectors, j);
if nbrs.len() >= self.min_points {
for k in nbrs {
if labels[k] == -1 || labels[k] == -1 {
if labels[k] == -1 {
seed.push(k);
}
labels[k] = cluster_id;
}
}
}
}
}
cluster_id += 1;
}
let num_clusters = cluster_id as usize;
let mut noise_count = 0_usize;
let mut cluster_sizes = vec![0_usize; num_clusters];
for &label in &labels {
if label == -1 {
noise_count += 1;
} else {
cluster_sizes[label as usize] += 1;
}
}
eprintln!(
"[cluster] DBSCAN eps={} min_pts={}{} clusters + {} noise from {} items",
self.epsilon, self.min_points, num_clusters, noise_count, n
);
Ok(ClusteringResult {
num_clusters,
num_noise: noise_count,
cluster_sizes,
labels,
})
}
fn region_query(&self, vectors: &[Vec<f32>], idx: usize) -> Vec<usize> {
let mut neighbors = Vec::new();
let target = &vectors[idx];
for (i, v) in vectors.iter().enumerate() {
if i != idx && euclidean_sq(target, v) <= self.epsilon * self.epsilon {
neighbors.push(i);
}
}
neighbors
}
pub fn extract_centroids(
vectors: &[Vec<f32>],
labels: &[i32],
num_clusters: usize,
) -> Vec<Vec<f32>> {
let dim = vectors.first().map(|v| v.len()).unwrap_or(1024);
let mut centroids = vec![vec![0.0_f32; dim]; num_clusters];
let mut counts = vec![0_usize; num_clusters];
for (vec, &label) in vectors.iter().zip(labels.iter()) {
if label >= 0 {
let ci = label as usize;
if ci < num_clusters {
for (j, &val) in vec.iter().enumerate() {
centroids[ci][j] += val;
}
counts[ci] += 1;
}
}
}
for ci in 0..num_clusters {
if counts[ci] > 0 {
let inv = 1.0 / counts[ci] as f32;
for val in centroids[ci].iter_mut() {
*val *= inv;
}
}
}
centroids
}
pub fn find_duplicates(centroids: &[Vec<f32>], cluster_sizes: &[usize]) -> Vec<(usize, usize)> {
let mut pairs = Vec::new();
for i in 0..centroids.len() {
if cluster_sizes[i] < 2 { continue; }
for j in (i + 1)..centroids.len() {
if cluster_sizes[j] < 2 { continue; }
let sim = cosine_similarity(&centroids[i], &centroids[j]);
if sim > 0.95 {
pairs.push((i, j));
}
}
}
eprintln!("[cluster] found {} duplicate pairs (sim > 0.95)", pairs.len());
pairs
}
}
fn euclidean_sq(a: &[f32], b: &[f32]) -> f64 {
let n = a.len().min(b.len());
let mut sum = 0.0_f64;
for i in 0..n {
let diff = a[i] as f64 - b[i] as f64;
sum += diff * diff;
}
sum
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let n = a.len().min(b.len());
let (mut dot, mut na, mut nb) = (0.0_f32, 0.0_f32, 0.0_f32);
for i in 0..n {
dot += a[i] * b[i];
na += a[i] * a[i];
nb += b[i] * b[i];
}
(dot / (na.sqrt() * nb.sqrt().max(1e-10))).max(0.0)
}