use std::io::Write;
use std::path::Path;
use std::path::PathBuf;
use std::time::Instant;
use kdam::{tqdm, BarExt};
use ndarray::{Array1, Array2};
use ndarray_npy::write_npy;
use num_format::{Locale, ToFormattedString};
use serde::{Deserialize, Serialize};
use abd_clam::cluster::PartitionCriteria;
use abd_clam::dataset::{Dataset, VecVec};
use abd_clam::search::cakes::CAKES;
pub mod utils;
use utils::distances;
use utils::search_readers;
fn main() {
let reports_root = get_reports_root();
for &(data_name, metric_name) in search_readers::SEARCH_DATASETS {
if ["deep-image", "nytimes", "lastfm"].contains(&data_name) {
continue;
}
if metric_name == "jaccard" {
continue;
}
let data_dir = {
let mut path = reports_root.clone();
path.push(data_name);
if path.exists() {
std::fs::remove_dir_all(&path).unwrap();
}
std::fs::create_dir(&path).unwrap();
path
};
for &(metric_name, metric) in distances::METRICS {
let out_dir = {
let mut path = data_dir.clone();
path.push(metric_name);
if path.exists() {
std::fs::remove_dir_all(&path).unwrap();
}
std::fs::create_dir(&path).unwrap();
path
};
println!();
println!("Making reports on {data_name} with {metric_name} ...");
let (data, queries) = search_readers::read_search_data(data_name).unwrap();
let data = VecVec::new(data, metric, data_name.to_string(), false);
let car = data.cardinality().to_formatted_string(&Locale::en);
let dim = data.dimensionality().to_formatted_string(&Locale::en);
println!("Got data with shape ({car} x {dim}) ...");
let start = Instant::now();
let criteria = PartitionCriteria::new(true).with_min_cardinality(1);
let cakes = CAKES::new(data, Some(42)).build(&criteria);
let build_time = start.elapsed().as_secs_f32();
println!("Built CAKES on {data_name} with {metric_name} in {build_time:.3} seconds ...");
let data = cakes.data();
let batch_size = 500_000;
let linear_time = report_linear(data, &queries, &out_dir, batch_size);
let time = CakesTime {
data_name,
metric_name,
cardinality: data.cardinality(),
dimensionality: data.dimensionality(),
build_time,
num_queries: queries.len(),
linear_time,
batch_size,
};
let time = serde_json::to_string_pretty(&time).unwrap();
let time_path = {
let mut path = out_dir.clone();
path.push("time-taken.json");
path
};
let mut time_file = std::fs::File::create(&time_path).unwrap();
time_file.write_all(time.as_bytes()).unwrap();
println!("Wrote timings file {time_path:?} ...");
}
}
}
fn get_reports_root() -> PathBuf {
let mut path = std::env::current_dir().unwrap();
path.push("reports");
assert!(
path.exists(),
"Please create a `reports` directory in the root of the clam repo."
);
assert!(
path.is_dir(),
"Please create a `reports` directory in the root of the clam repo."
);
path
}
fn report_linear(data: &VecVec<f32, f32>, queries: &[Vec<f32>], out_dir: &Path, batch_size: usize) -> f32 {
let indices = data.indices();
let num_batches = {
let num_batches = indices.len() / batch_size;
if indices.len() % batch_size == 0 {
num_batches
} else {
num_batches + 1
}
};
let mut time = 0.;
for (i, batch) in indices.chunks(batch_size).enumerate() {
let n = i + 1;
let mut pb = tqdm!(total = queries.len(), desc = format!("Linear Batch {n}/{num_batches}"));
let mut array = Array2::<f32>::default((0, batch.len()));
for query in queries.iter() {
let start = Instant::now();
let distances = data.query_to_many(query, batch);
time += start.elapsed().as_secs_f32();
array.push_row(Array1::from_vec(distances).view()).unwrap();
pb.update(1);
}
let out_path = {
let mut path = out_dir.to_path_buf();
path.push(format!("query-distances-batch-{n}-{num_batches}.npy"));
path
};
write_npy(&out_path, &array).unwrap();
}
time
}
#[derive(Debug, Serialize, Deserialize)]
struct CakesTime<'a> {
data_name: &'a str,
metric_name: &'a str,
cardinality: usize,
dimensionality: usize,
build_time: f32,
num_queries: usize,
linear_time: f32,
batch_size: usize,
}