From 13f1509e2a8059a700e2625f7bd82b07b2cee924 Mon Sep 17 00:00:00 2001 From: Ashwin Krishna Kumar Date: Tue, 28 Jul 2026 12:20:31 +0530 Subject: [PATCH] Improve PQ training error message --- .../jvector/quantization/ProductQuantization.java | 3 +++ .../quantization/TestProductQuantization.java | 13 +++++++++++++ 2 files changed, 16 insertions(+) diff --git a/jvector-base/src/main/java/io/github/jbellis/jvector/quantization/ProductQuantization.java b/jvector-base/src/main/java/io/github/jbellis/jvector/quantization/ProductQuantization.java index c84b7b955..b54d0684b 100644 --- a/jvector-base/src/main/java/io/github/jbellis/jvector/quantization/ProductQuantization.java +++ b/jvector-base/src/main/java/io/github/jbellis/jvector/quantization/ProductQuantization.java @@ -116,6 +116,9 @@ public static ProductQuantization compute(RandomAccessVectorValues ravv, { checkClusterCount(clusterCount); + if (ravv.size() < clusterCount) { + throw new IllegalArgumentException("Cannot train PQ with %d clusters on %d points, supply more training vectors or lower cluster count."); + } var subvectorSizesAndOffsets = getSubvectorSizesAndOffsets(ravv.dimension(), M); var vectors = extractTrainingVectors(ravv, parallelExecutor); diff --git a/jvector-tests/src/test/java/io/github/jbellis/jvector/quantization/TestProductQuantization.java b/jvector-tests/src/test/java/io/github/jbellis/jvector/quantization/TestProductQuantization.java index 992604b7e..16e3036bc 100644 --- a/jvector-tests/src/test/java/io/github/jbellis/jvector/quantization/TestProductQuantization.java +++ b/jvector-tests/src/test/java/io/github/jbellis/jvector/quantization/TestProductQuantization.java @@ -433,4 +433,17 @@ public void testPQCodebookSums() { } } } + + @Test + public void testErrorOnInsufficientVectors() { + int dim = 384; + var pqM = 48; + + var ravv100 = new ListRandomAccessVectorValues(createRandomVectors(100, dim), dim); + var ravv200 = new ListRandomAccessVectorValues(createRandomVectors(200, dim), dim); + ProductQuantization.compute(ravv100, pqM, 50, false); // should not throw + assertThrows(IllegalArgumentException.class, () -> ProductQuantization.compute(ravv100, pqM, 150, false)); + ProductQuantization.compute(ravv200, pqM, 150, false); // should not throw + assertThrows(IllegalArgumentException.class, () -> ProductQuantization.compute(ravv200, pqM, 256, false)); + } }