Skip to content
34 changes: 32 additions & 2 deletions diskann-benchmark/src/disk_index/benchmarks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,17 +22,21 @@ use half::f16;

use crate::{
disk_index::{
build::{build_disk_index, DiskBuildStats},
build::{
build_disk_index, build_pq_kmeans_router, DiskBuildStats, PqKmeansRouterBuildStats,
},
search::{search_disk_index, DiskSearchStats},
},
inputs::disk::{DiskIndexLoad, DiskIndexOperation, DiskIndexSource},
inputs::disk::{DiskIndexLoad, DiskIndexOperation, DiskIndexSource, PqKmeansRouterBuild},
};

/// Disk Index
struct DiskIndex<T> {
_vector_type: std::marker::PhantomData<T>,
}

struct PqKmeansRouterBuildJob;

#[derive(Debug, Serialize, Deserialize)]
pub(super) struct DiskIndexStats {
pub(super) build: Option<DiskBuildStats>,
Expand Down Expand Up @@ -108,11 +112,37 @@ where
}
}

impl Benchmark for PqKmeansRouterBuildJob {
type Input = PqKmeansRouterBuild;
type Output = PqKmeansRouterBuildStats;

fn try_match(&self, _input: &PqKmeansRouterBuild, context: &MatchContext) -> Score {
context.success(0)
}

fn description(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "PQ-kmeans start-point router build")
}

fn run(
&self,
input: &PqKmeansRouterBuild,
_checkpoint: Checkpoint<'_>,
mut output: &mut dyn Output,
) -> anyhow::Result<PqKmeansRouterBuildStats> {
writeln!(output, "{}", input)?;
let stats = build_pq_kmeans_router(input)?;
writeln!(output, "{}", stats)?;
Ok(stats)
}
}

////////////////////////////
// Benchmark Registration //
////////////////////////////

pub(super) fn register_benchmarks(registry: &mut Registry) -> anyhow::Result<()> {
registry.register("pq-kmeans-router-build", PqKmeansRouterBuildJob)?;
registry.register_regression("disk-index-f32", DiskIndex::<f32>::new())?;
registry.register_regression("disk-index-f16", DiskIndex::<f16>::new())?;
registry.register_regression("disk-index-u8", DiskIndex::<u8>::new())?;
Expand Down
84 changes: 80 additions & 4 deletions diskann-benchmark/src/disk_index/build.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,21 +13,35 @@ use diskann::{
use diskann_benchmark_runner::utils::MicroSeconds;
use diskann_disk::{
build::builder::build::DiskIndexBuilder,
data_model::AdHoc,
data_model::{AdHoc, CachingStrategy},
disk_index_build_parameter::{
DiskIndexBuildParameters, MemoryBudget, NumPQChunks, DISK_SECTOR_LEN,
},
storage::DiskIndexWriter,
search::{
pq_kmeans_router::{PqKmeansRouterBuildParams, PqKmeansRouterData},
provider::{
aligned_file_reader::AlignedFileReaderFactory,
disk_vertex_provider_factory::DiskVertexProviderFactory,
},
traits::VertexProviderFactory,
},
storage::{disk_index_reader::DiskIndexReader, DiskIndexWriter},
};
use diskann_providers::storage::{
get_compressed_pq_file, get_disk_index_file, get_pq_pivot_file, FileStorageProvider,
StorageReadProvider, StorageWriteProvider,
};
use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
use diskann_providers::{model::IndexConfiguration, utils::load_metadata_from_file};
use diskann_vector::distance::Metric;
use opentelemetry::global;
use opentelemetry::trace::Tracer;
use opentelemetry_sdk::trace::SdkTracerProvider;
use scopeguard::defer;

use crate::{disk_index::json_spancollector::JsonSpanCollector, inputs::disk::DiskIndexBuild};
use crate::{
disk_index::json_spancollector::JsonSpanCollector,
inputs::disk::{DiskIndexBuild, PqKmeansRouterBuild},
};

#[derive(Serialize, Deserialize, Debug)]
pub(super) struct DiskBuildStats {
Expand Down Expand Up @@ -55,6 +69,68 @@ impl fmt::Display for DiskBuildStats {
}
}

#[derive(Serialize, Deserialize, Debug)]
pub(super) struct PqKmeansRouterBuildStats {
build_time: MicroSeconds,
num_points: usize,
num_pq_chunks: usize,
num_representatives: usize,
artifact_bytes: u64,
}

impl fmt::Display for PqKmeansRouterBuildStats {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(
f,
"PQ-kmeans router build time: {:.3}s",
self.build_time.as_seconds()
)?;
writeln!(f, "Points: {}", self.num_points)?;
writeln!(f, "PQ chunks: {}", self.num_pq_chunks)?;
writeln!(f, "Representatives: {}", self.num_representatives)?;
writeln!(f, "Artifact bytes: {}", self.artifact_bytes)
}
}

pub(super) fn build_pq_kmeans_router(
params: &PqKmeansRouterBuild,
) -> anyhow::Result<PqKmeansRouterBuildStats> {
let start = std::time::Instant::now();
let index_reader = DiskIndexReader::new(
get_pq_pivot_file(&params.load_path),
get_compressed_pq_file(&params.load_path),
&FileStorageProvider,
)?;
let vertex_provider_factory =
DiskVertexProviderFactory::<AdHoc<f32>, AlignedFileReaderFactory>::from_disk_index_path(
get_disk_index_file(&params.load_path),
CachingStrategy::None,
)?;
let graph_header = vertex_provider_factory.get_header()?;
let pq_data = index_reader.get_pq_data();
let router_data = PqKmeansRouterData::build_from_pq_data(
pq_data.as_ref(),
PqKmeansRouterBuildParams {
metric: params.distance.into(),
num_representatives: params.num_representatives,
training_sample_size: params.training_sample_size,
max_iterations: params.max_iterations,
},
Some(graph_header.metadata().medoid as u32),
)?;
router_data.save_to_path(&params.artifact)?;
let artifact_bytes = std::fs::metadata(&params.artifact)?.len();
let build_time = start.elapsed().into();

Ok(PqKmeansRouterBuildStats {
build_time,
num_points: router_data.num_points,
num_pq_chunks: router_data.num_pq_chunks,
num_representatives: router_data.representative_ids.len(),
artifact_bytes,
})
}

pub(super) fn build_disk_index<T, StorageProviderType>(
storage_provider: &StorageProviderType,
params: &DiskIndexBuild,
Expand Down
Loading