From 47852e8d15190e8befd664fbec5895061b640fbd Mon Sep 17 00:00:00 2001 From: Andrew Kane Date: Sun, 3 Dec 2023 15:12:50 -0800 Subject: [PATCH] Started half indexing --- sql/vector--0.5.1--0.6.0.sql | 5 ++ sql/vector.sql | 7 +++ src/hnswbuild.c | 4 +- test/t/019_hnsw_half.pl | 93 ++++++++++++++++++++++++++++++++++++ 4 files changed, 107 insertions(+), 2 deletions(-) create mode 100644 test/t/019_hnsw_half.pl diff --git a/sql/vector--0.5.1--0.6.0.sql b/sql/vector--0.5.1--0.6.0.sql index 61348fa..ff903de 100644 --- a/sql/vector--0.5.1--0.6.0.sql +++ b/sql/vector--0.5.1--0.6.0.sql @@ -75,3 +75,8 @@ CREATE OPERATOR <=> ( LEFTARG = half[], RIGHTARG = half[], PROCEDURE = cosine_distance, COMMUTATOR = '<=>' ); + +CREATE OPERATOR CLASS half_l2_ops + FOR TYPE half[] USING hnsw AS + OPERATOR 1 <-> (half[], half[]) FOR ORDER BY float_ops, + FUNCTION 1 half_l2_squared_distance(half[], half[]); diff --git a/sql/vector.sql b/sql/vector.sql index c86d6d9..dcc1feb 100644 --- a/sql/vector.sql +++ b/sql/vector.sql @@ -377,3 +377,10 @@ CREATE OPERATOR <=> ( LEFTARG = half[], RIGHTARG = half[], PROCEDURE = cosine_distance, COMMUTATOR = '<=>' ); + +-- half opclasses + +CREATE OPERATOR CLASS half_l2_ops + FOR TYPE half[] USING hnsw AS + OPERATOR 1 <-> (half[], half[]) FOR ORDER BY float_ops, + FUNCTION 1 half_l2_squared_distance(half[], half[]); diff --git a/src/hnswbuild.c b/src/hnswbuild.c index ea7b889..5a861e2 100644 --- a/src/hnswbuild.c +++ b/src/hnswbuild.c @@ -462,8 +462,8 @@ InitBuildState(HnswBuildState * buildstate, Relation heap, Relation index, Index buildstate->dimensions = TupleDescAttr(index->rd_att, 0)->atttypmod; /* Require column to have dimensions to be indexed */ - if (buildstate->dimensions < 0) - elog(ERROR, "column does not have dimensions"); + // if (buildstate->dimensions < 0) + // elog(ERROR, "column does not have dimensions"); if (buildstate->dimensions > HNSW_MAX_DIM) elog(ERROR, "column cannot have more than %d dimensions for hnsw index", HNSW_MAX_DIM); diff --git a/test/t/019_hnsw_half.pl b/test/t/019_hnsw_half.pl new file mode 100644 index 0000000..3e7ed93 --- /dev/null +++ b/test/t/019_hnsw_half.pl @@ -0,0 +1,93 @@ +use strict; +use warnings; +use PostgresNode; +use TestLib; +use Test::More; + +my $node; +my @queries = (); +my @expected; +my $limit = 20; + +sub test_recall +{ + my ($min, $operator) = @_; + my $correct = 0; + my $total = 0; + + my $explain = $node->safe_psql("postgres", qq( + SET enable_seqscan = off; + EXPLAIN ANALYZE SELECT i FROM tst ORDER BY v $operator '$queries[0]' LIMIT $limit; + )); + like($explain, qr/Index Scan/); + + for my $i (0 .. $#queries) + { + my $actual = $node->safe_psql("postgres", qq( + SET enable_seqscan = off; + SELECT i FROM tst ORDER BY v $operator '$queries[$i]' LIMIT $limit; + )); + my @actual_ids = split("\n", $actual); + my %actual_set = map { $_ => 1 } @actual_ids; + + my @expected_ids = split("\n", $expected[$i]); + + foreach (@expected_ids) + { + if (exists($actual_set{$_})) + { + $correct++; + } + $total++; + } + } + + cmp_ok($correct / $total, ">=", $min, $operator); +} + +# Initialize node +$node = get_new_node('node'); +$node->init; +$node->start; + +# Create table +$node->safe_psql("postgres", "CREATE EXTENSION vector;"); +$node->safe_psql("postgres", "CREATE TABLE tst (i int4, v half[3]);"); +$node->safe_psql("postgres", + "INSERT INTO tst SELECT i, ARRAY[random(), random(), random()]::numeric[]::half[] FROM generate_series(1, 10000) i;" +); + +# Generate queries +for (1 .. 20) +{ + my $r1 = rand(); + my $r2 = rand(); + my $r3 = rand(); + push(@queries, "{$r1,$r2,$r3}"); +} + +# Check each index type +my @operators = ("<->"); +my @opclasses = ("half_l2_ops"); + +for my $i (0 .. $#operators) +{ + my $operator = $operators[$i]; + my $opclass = $opclasses[$i]; + + # Get exact results + @expected = (); + foreach (@queries) + { + my $res = $node->safe_psql("postgres", "SELECT i FROM tst ORDER BY v $operator '$_' LIMIT $limit;"); + push(@expected, $res); + } + + # Add index + $node->safe_psql("postgres", "CREATE INDEX ON tst USING hnsw (v $opclass);"); + + my $min = $operator eq "<#>" ? 0.80 : 0.99; + test_recall($min, $operator); +} + +done_testing();