use super::*;
#[test]
fn test_should_use_gpu_threshold() {
assert!(!should_use_gpu(100, 16, 8));
assert!(should_use_gpu(10000, 256, 8));
assert!(!should_use_gpu(10_000_000 / (256 * 8), 256, 8));
}
#[test]
fn test_gpu_context_new_does_not_panic() {
let _ctx = PqGpuContext::new();
}
#[test]
fn test_gpu_kmeans_assign_matches_cpu() {
let sub_vectors = vec![
vec![1.0, 0.0, 0.0, 0.0],
vec![0.9, 0.1, 0.0, 0.0],
vec![0.0, 1.0, 0.0, 0.0],
vec![0.1, 0.9, 0.0, 0.0],
vec![0.0, 0.0, 1.0, 0.0],
vec![0.0, 0.0, 0.9, 0.1],
vec![1.0, 1.0, 0.0, 0.0],
vec![0.0, 0.0, 0.0, 1.0],
vec![0.5, 0.5, 0.0, 0.0],
vec![0.0, 0.0, 0.5, 0.5],
];
let centroids = vec![
vec![1.0, 0.0, 0.0, 0.0],
vec![0.0, 1.0, 0.0, 0.0],
vec![0.0, 0.0, 1.0, 0.0],
];
let cpu_assignments: Vec<usize> = sub_vectors
.iter()
.map(|v| {
centroids
.iter()
.enumerate()
.map(|(idx, c)| {
let dist: f32 = v.iter().zip(c.iter()).map(|(a, b)| (a - b) * (a - b)).sum();
(idx, dist)
})
.min_by(|a, b| a.1.total_cmp(&b.1))
.map(|(idx, _)| idx)
.unwrap()
})
.collect();
if let Some(ctx) = PqGpuContext::new() {
if let Some(gpu_assignments) = gpu_kmeans_assign(&ctx, &sub_vectors, ¢roids, 4) {
assert_eq!(
gpu_assignments.len(),
sub_vectors.len(),
"GPU must return one assignment per vector"
);
assert_eq!(
gpu_assignments, cpu_assignments,
"GPU assignments must match CPU"
);
}
}
}
#[test]
fn test_gpu_kmeans_assign_empty_input() {
if let Some(ctx) = PqGpuContext::new() {
assert!(gpu_kmeans_assign(&ctx, &[], &[vec![1.0]], 1).is_none());
assert!(gpu_kmeans_assign(&ctx, &[vec![1.0]], &[], 1).is_none());
assert!(gpu_kmeans_assign(&ctx, &[vec![1.0]], &[vec![1.0]], 0).is_none());
}
}
#[test]
fn test_gpu_kmeans_assign_dimension_mismatch_returns_none() {
if let Some(ctx) = PqGpuContext::new() {
let sub_vectors = vec![vec![1.0, 0.0, 0.0]]; let centroids = vec![vec![1.0, 0.0, 0.0, 0.0]]; assert!(
gpu_kmeans_assign(&ctx, &sub_vectors, ¢roids, 4).is_none(),
"mismatched sub_vector dim must return None"
);
}
}
#[test]
fn test_pq_context_shares_global_device() {
let gpu_available = GpuAccelerator::is_available();
let pq_ctx = PqGpuContext::new();
assert_eq!(
pq_ctx.is_some(),
gpu_available,
"PqGpuContext availability must match GpuAccelerator::is_available()"
);
}