pgvector 0.3.5__tar.gz → 0.3.6__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (58) hide show
  1. {pgvector-0.3.5 → pgvector-0.3.6}/PKG-INFO +3 -3
  2. {pgvector-0.3.5 → pgvector-0.3.6}/README.md +2 -2
  3. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg2/halfvec.py +7 -2
  4. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg2/register.py +6 -5
  5. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg2/sparsevec.py +7 -2
  6. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg2/vector.py +7 -2
  7. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector.egg-info/PKG-INFO +3 -3
  8. {pgvector-0.3.5 → pgvector-0.3.6}/pyproject.toml +1 -1
  9. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_asyncpg.py +18 -0
  10. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_django.py +41 -1
  11. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_peewee.py +22 -0
  12. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_psycopg2.py +30 -3
  13. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_sqlalchemy.py +18 -2
  14. {pgvector-0.3.5 → pgvector-0.3.6}/LICENSE.txt +0 -0
  15. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/asyncpg/__init__.py +0 -0
  16. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/asyncpg/register.py +0 -0
  17. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/__init__.py +0 -0
  18. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/bit.py +0 -0
  19. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/extensions.py +0 -0
  20. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/functions.py +0 -0
  21. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/halfvec.py +0 -0
  22. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/indexes.py +0 -0
  23. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/sparsevec.py +0 -0
  24. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/django/vector.py +0 -0
  25. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/peewee/__init__.py +0 -0
  26. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/peewee/bit.py +0 -0
  27. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/peewee/halfvec.py +0 -0
  28. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/peewee/sparsevec.py +0 -0
  29. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/peewee/vector.py +0 -0
  30. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/__init__.py +0 -0
  31. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/bit.py +0 -0
  32. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/halfvec.py +0 -0
  33. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/register.py +0 -0
  34. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/sparsevec.py +0 -0
  35. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg/vector.py +0 -0
  36. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/psycopg2/__init__.py +0 -0
  37. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/__init__.py +0 -0
  38. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/bit.py +0 -0
  39. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/functions.py +0 -0
  40. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/halfvec.py +0 -0
  41. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/sparsevec.py +0 -0
  42. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/sqlalchemy/vector.py +0 -0
  43. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/utils/__init__.py +0 -0
  44. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/utils/bit.py +0 -0
  45. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/utils/halfvec.py +0 -0
  46. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/utils/sparsevec.py +0 -0
  47. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector/utils/vector.py +0 -0
  48. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector.egg-info/SOURCES.txt +0 -0
  49. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector.egg-info/dependency_links.txt +0 -0
  50. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector.egg-info/requires.txt +0 -0
  51. {pgvector-0.3.5 → pgvector-0.3.6}/pgvector.egg-info/top_level.txt +0 -0
  52. {pgvector-0.3.5 → pgvector-0.3.6}/setup.cfg +0 -0
  53. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_bit.py +0 -0
  54. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_half_vector.py +0 -0
  55. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_psycopg.py +0 -0
  56. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_sparse_vector.py +0 -0
  57. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_sqlmodel.py +0 -0
  58. {pgvector-0.3.5 → pgvector-0.3.6}/tests/test_vector.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: pgvector
3
- Version: 0.3.5
3
+ Version: 0.3.6
4
4
  Summary: pgvector support for Python
5
5
  Author-email: Andrew Kane <andrew@ankane.org>
6
6
  License: MIT
@@ -197,7 +197,7 @@ Average vectors
197
197
  ```python
198
198
  from pgvector.sqlalchemy import avg
199
199
 
200
- session.scalars(select(func.avg(Item.embedding))).first()
200
+ session.scalars(select(avg(Item.embedding))).first()
201
201
  ```
202
202
 
203
203
  Also supports `sum`
@@ -279,7 +279,7 @@ Average vectors
279
279
  ```python
280
280
  from pgvector.sqlalchemy import avg
281
281
 
282
- session.exec(select(func.avg(Item.embedding))).first()
282
+ session.exec(select(avg(Item.embedding))).first()
283
283
  ```
284
284
 
285
285
  Also supports `sum`
@@ -185,7 +185,7 @@ Average vectors
185
185
  ```python
186
186
  from pgvector.sqlalchemy import avg
187
187
 
188
- session.scalars(select(func.avg(Item.embedding))).first()
188
+ session.scalars(select(avg(Item.embedding))).first()
189
189
  ```
190
190
 
191
191
  Also supports `sum`
@@ -267,7 +267,7 @@ Average vectors
267
267
  ```python
268
268
  from pgvector.sqlalchemy import avg
269
269
 
270
- session.exec(select(func.avg(Item.embedding))).first()
270
+ session.exec(select(avg(Item.embedding))).first()
271
271
  ```
272
272
 
273
273
  Also supports `sum`
@@ -1,4 +1,4 @@
1
- from psycopg2.extensions import adapt, new_type, register_adapter, register_type
1
+ from psycopg2.extensions import adapt, new_array_type, new_type, register_adapter, register_type
2
2
  from ..utils import HalfVector
3
3
 
4
4
 
@@ -14,7 +14,12 @@ def cast_halfvec(value, cur):
14
14
  return HalfVector._from_db(value)
15
15
 
16
16
 
17
- def register_halfvec_info(oid, scope):
17
+ def register_halfvec_info(oid, array_oid, scope):
18
18
  halfvec = new_type((oid,), 'HALFVEC', cast_halfvec)
19
19
  register_type(halfvec, scope)
20
+
21
+ if array_oid is not None:
22
+ halfvecarray = new_array_type((array_oid,), 'HALFVECARRAY', halfvec)
23
+ register_type(halfvecarray, scope)
24
+
20
25
  register_adapter(HalfVector, HalfvecAdapter)
@@ -7,22 +7,23 @@ from .vector import register_vector_info
7
7
 
8
8
  # TODO make globally False by default in 0.4.0
9
9
  # note: register_adapter is always global
10
- def register_vector(conn_or_curs=None, globally=True):
10
+ # TODO make arrays True by defalt in 0.4.0
11
+ def register_vector(conn_or_curs=None, globally=True, arrays=False):
11
12
  conn = conn_or_curs if hasattr(conn_or_curs, 'cursor') else conn_or_curs.connection
12
13
  cur = conn.cursor(cursor_factory=cursor)
13
14
  scope = None if globally else conn_or_curs
14
15
 
15
16
  # use to_regtype to get first matching type in search path
16
- cur.execute("SELECT typname, oid FROM pg_type WHERE oid IN (to_regtype('vector'), to_regtype('halfvec'), to_regtype('sparsevec'))")
17
+ cur.execute("SELECT typname, oid FROM pg_type WHERE oid IN (to_regtype('vector'), to_regtype('_vector'), to_regtype('halfvec'), to_regtype('_halfvec'), to_regtype('sparsevec'), to_regtype('_sparsevec'))")
17
18
  type_info = dict(cur.fetchall())
18
19
 
19
20
  if 'vector' not in type_info:
20
21
  raise psycopg2.ProgrammingError('vector type not found in the database')
21
22
 
22
- register_vector_info(type_info['vector'], scope)
23
+ register_vector_info(type_info['vector'], type_info['_vector'] if arrays else None, scope)
23
24
 
24
25
  if 'halfvec' in type_info:
25
- register_halfvec_info(type_info['halfvec'], scope)
26
+ register_halfvec_info(type_info['halfvec'], type_info['_halfvec'] if arrays else None, scope)
26
27
 
27
28
  if 'sparsevec' in type_info:
28
- register_sparsevec_info(type_info['sparsevec'], scope)
29
+ register_sparsevec_info(type_info['sparsevec'], type_info['_sparsevec'] if arrays else None, scope)
@@ -1,4 +1,4 @@
1
- from psycopg2.extensions import adapt, new_type, register_adapter, register_type
1
+ from psycopg2.extensions import adapt, new_array_type, new_type, register_adapter, register_type
2
2
  from ..utils import SparseVector
3
3
 
4
4
 
@@ -14,7 +14,12 @@ def cast_sparsevec(value, cur):
14
14
  return SparseVector._from_db(value)
15
15
 
16
16
 
17
- def register_sparsevec_info(oid, scope):
17
+ def register_sparsevec_info(oid, array_oid, scope):
18
18
  sparsevec = new_type((oid,), 'SPARSEVEC', cast_sparsevec)
19
19
  register_type(sparsevec, scope)
20
+
21
+ if array_oid is not None:
22
+ sparsevecarray = new_array_type((array_oid,), 'SPARSEVECARRAY', sparsevec)
23
+ register_type(sparsevecarray, scope)
24
+
20
25
  register_adapter(SparseVector, SparsevecAdapter)
@@ -1,5 +1,5 @@
1
1
  import numpy as np
2
- from psycopg2.extensions import adapt, new_type, register_adapter, register_type
2
+ from psycopg2.extensions import adapt, new_array_type, new_type, register_adapter, register_type
3
3
  from ..utils import Vector
4
4
 
5
5
 
@@ -15,7 +15,12 @@ def cast_vector(value, cur):
15
15
  return Vector._from_db(value)
16
16
 
17
17
 
18
- def register_vector_info(oid, scope):
18
+ def register_vector_info(oid, array_oid, scope):
19
19
  vector = new_type((oid,), 'VECTOR', cast_vector)
20
20
  register_type(vector, scope)
21
+
22
+ if array_oid is not None:
23
+ vectorarray = new_array_type((array_oid,), 'VECTORARRAY', vector)
24
+ register_type(vectorarray, scope)
25
+
21
26
  register_adapter(np.ndarray, VectorAdapter)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: pgvector
3
- Version: 0.3.5
3
+ Version: 0.3.6
4
4
  Summary: pgvector support for Python
5
5
  Author-email: Andrew Kane <andrew@ankane.org>
6
6
  License: MIT
@@ -197,7 +197,7 @@ Average vectors
197
197
  ```python
198
198
  from pgvector.sqlalchemy import avg
199
199
 
200
- session.scalars(select(func.avg(Item.embedding))).first()
200
+ session.scalars(select(avg(Item.embedding))).first()
201
201
  ```
202
202
 
203
203
  Also supports `sum`
@@ -279,7 +279,7 @@ Average vectors
279
279
  ```python
280
280
  from pgvector.sqlalchemy import avg
281
281
 
282
- session.exec(select(func.avg(Item.embedding))).first()
282
+ session.exec(select(avg(Item.embedding))).first()
283
283
  ```
284
284
 
285
285
  Also supports `sum`
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "pgvector"
7
- version = "0.3.5"
7
+ version = "0.3.6"
8
8
  description = "pgvector support for Python"
9
9
  readme = "README.md"
10
10
  authors = [
@@ -94,6 +94,24 @@ class TestAsyncpg:
94
94
 
95
95
  await conn.close()
96
96
 
97
+ @pytest.mark.asyncio
98
+ async def test_vector_array(self):
99
+ conn = await asyncpg.connect(database='pgvector_python_test')
100
+ await conn.execute('CREATE EXTENSION IF NOT EXISTS vector')
101
+ await conn.execute('DROP TABLE IF EXISTS asyncpg_items')
102
+ await conn.execute('CREATE TABLE asyncpg_items (id bigserial PRIMARY KEY, embeddings vector[])')
103
+
104
+ await register_vector(conn)
105
+
106
+ embeddings = [np.array([1.5, 2, 3]), np.array([4.5, 5, 6])]
107
+ await conn.execute("INSERT INTO asyncpg_items (embeddings) VALUES (ARRAY[$1, $2]::vector[])", embeddings[0], embeddings[1])
108
+
109
+ res = await conn.fetch("SELECT * FROM asyncpg_items ORDER BY id")
110
+ assert np.array_equal(res[0]['embeddings'][0], embeddings[0])
111
+ assert np.array_equal(res[0]['embeddings'][1], embeddings[1])
112
+
113
+ await conn.close()
114
+
97
115
  @pytest.mark.asyncio
98
116
  async def test_pool(self):
99
117
  async def init(conn):
@@ -1,8 +1,10 @@
1
1
  import django
2
2
  from django.conf import settings
3
+ from django.contrib.postgres.fields import ArrayField
3
4
  from django.core import serializers
4
5
  from django.db import connection, migrations, models
5
- from django.db.models import Avg, Sum
6
+ from django.db.models import Avg, Sum, FloatField, DecimalField
7
+ from django.db.models.functions import Cast
6
8
  from django.db.migrations.loader import MigrationLoader
7
9
  from django.forms import ModelForm
8
10
  from math import sqrt
@@ -46,6 +48,9 @@ class Item(models.Model):
46
48
  half_embedding = HalfVectorField(dimensions=3, null=True, blank=True)
47
49
  binary_embedding = BitField(length=3, null=True, blank=True)
48
50
  sparse_embedding = SparseVectorField(dimensions=3, null=True, blank=True)
51
+ embeddings = ArrayField(VectorField(dimensions=3), null=True, blank=True)
52
+ double_embedding = ArrayField(FloatField(), null=True, blank=True)
53
+ numeric_embedding = ArrayField(DecimalField(max_digits=20, decimal_places=10), null=True, blank=True)
49
54
 
50
55
  class Meta:
51
56
  app_label = 'django_app'
@@ -82,6 +87,9 @@ class Migration(migrations.Migration):
82
87
  ('half_embedding', pgvector.django.HalfVectorField(dimensions=3, null=True, blank=True)),
83
88
  ('binary_embedding', pgvector.django.BitField(length=3, null=True, blank=True)),
84
89
  ('sparse_embedding', pgvector.django.SparseVectorField(dimensions=3, null=True, blank=True)),
90
+ ('embeddings', ArrayField(pgvector.django.VectorField(dimensions=3), null=True, blank=True)),
91
+ ('double_embedding', ArrayField(FloatField(), null=True, blank=True)),
92
+ ('numeric_embedding', ArrayField(DecimalField(max_digits=20, decimal_places=10), null=True, blank=True)),
85
93
  ],
86
94
  ),
87
95
  migrations.AddIndex(
@@ -433,3 +441,35 @@ class TestDjango:
433
441
  assert Item.objects.first().half_embedding is None
434
442
  assert Item.objects.first().binary_embedding is None
435
443
  assert Item.objects.first().sparse_embedding is None
444
+
445
+ def test_vector_array(self):
446
+ Item(id=1, embeddings=[np.array([1, 2, 3]), np.array([4, 5, 6])]).save()
447
+
448
+ with connection.cursor() as cursor:
449
+ from pgvector.psycopg import register_vector
450
+ register_vector(cursor.connection)
451
+
452
+ # this fails if the driver does not cast arrays
453
+ item = Item.objects.get(pk=1)
454
+ assert item.embeddings[0].tolist() == [1, 2, 3]
455
+ assert item.embeddings[1].tolist() == [4, 5, 6]
456
+
457
+ def test_double_array(self):
458
+ Item(id=1, double_embedding=[1, 1, 1]).save()
459
+ Item(id=2, double_embedding=[2, 2, 2]).save()
460
+ Item(id=3, double_embedding=[1, 1, 2]).save()
461
+ distance = L2Distance(Cast('double_embedding', VectorField()), [1, 1, 1])
462
+ items = Item.objects.annotate(distance=distance).order_by(distance)
463
+ assert [v.id for v in items] == [1, 3, 2]
464
+ assert [v.distance for v in items] == [0, 1, sqrt(3)]
465
+ assert items[1].double_embedding == [1, 1, 2]
466
+
467
+ def test_numeric_array(self):
468
+ Item(id=1, numeric_embedding=[1, 1, 1]).save()
469
+ Item(id=2, numeric_embedding=[2, 2, 2]).save()
470
+ Item(id=3, numeric_embedding=[1, 1, 2]).save()
471
+ distance = L2Distance(Cast('numeric_embedding', VectorField()), [1, 1, 1])
472
+ items = Item.objects.annotate(distance=distance).order_by(distance)
473
+ assert [v.id for v in items] == [1, 3, 2]
474
+ assert [v.distance for v in items] == [0, 1, sqrt(3)]
475
+ assert items[1].numeric_embedding == [1, 1, 2]
@@ -199,3 +199,25 @@ class TestPeewee:
199
199
  Item.get_or_create(id=1, defaults={'embedding': [1, 2, 3]})
200
200
  Item.get_or_create(embedding=np.array([4, 5, 6]))
201
201
  Item.get_or_create(embedding=Item.embedding.to_value([7, 8, 9]))
202
+
203
+ def test_vector_array(self):
204
+ from playhouse.postgres_ext import PostgresqlExtDatabase, ArrayField
205
+
206
+ ext_db = PostgresqlExtDatabase('pgvector_python_test')
207
+
208
+ class ExtItem(BaseModel):
209
+ embeddings = ArrayField(VectorField, field_kwargs={'dimensions': 3}, index=False)
210
+
211
+ class Meta:
212
+ database = ext_db
213
+ table_name = 'peewee_ext_item'
214
+
215
+ ext_db.connect()
216
+ ext_db.drop_tables([ExtItem])
217
+ ext_db.create_tables([ExtItem])
218
+
219
+ # fails with column "embeddings" is of type vector[] but expression is of type text[]
220
+ # ExtItem.create(id=1, embeddings=[np.array([1, 2, 3]), np.array([4, 5, 6])])
221
+ # item = ExtItem.get_by_id(1)
222
+ # assert np.array_equal(item.embeddings[0], np.array([1, 2, 3]))
223
+ # assert np.array_equal(item.embeddings[1], np.array([4, 5, 6]))
@@ -1,5 +1,5 @@
1
1
  import numpy as np
2
- from pgvector.psycopg2 import register_vector, SparseVector
2
+ from pgvector.psycopg2 import register_vector, HalfVector, SparseVector
3
3
  import psycopg2
4
4
  from psycopg2.extras import DictCursor, RealDictCursor, NamedTupleCursor
5
5
 
@@ -9,9 +9,9 @@ conn.autocommit = True
9
9
  cur = conn.cursor()
10
10
  cur.execute('CREATE EXTENSION IF NOT EXISTS vector')
11
11
  cur.execute('DROP TABLE IF EXISTS psycopg2_items')
12
- cur.execute('CREATE TABLE psycopg2_items (id bigserial PRIMARY KEY, embedding vector(3), half_embedding halfvec(3), binary_embedding bit(3), sparse_embedding sparsevec(3))')
12
+ cur.execute('CREATE TABLE psycopg2_items (id bigserial PRIMARY KEY, embedding vector(3), half_embedding halfvec(3), binary_embedding bit(3), sparse_embedding sparsevec(3), embeddings vector[], half_embeddings halfvec[], sparse_embeddings sparsevec[])')
13
13
 
14
- register_vector(cur, globally=False)
14
+ register_vector(cur, globally=False, arrays=True)
15
15
 
16
16
 
17
17
  class TestPsycopg2:
@@ -55,6 +55,33 @@ class TestPsycopg2:
55
55
  assert res[0][0].to_list() == [1.5, 2, 3]
56
56
  assert res[1][0] is None
57
57
 
58
+ def test_vector_array(self):
59
+ embeddings = [np.array([1.5, 2, 3]), np.array([4.5, 5, 6])]
60
+ cur.execute('INSERT INTO psycopg2_items (embeddings) VALUES (%s::vector[])', (embeddings,))
61
+
62
+ cur.execute('SELECT embeddings FROM psycopg2_items ORDER BY id')
63
+ res = cur.fetchone()
64
+ assert np.array_equal(res[0][0], embeddings[0])
65
+ assert np.array_equal(res[0][1], embeddings[1])
66
+
67
+ def test_halfvec_array(self):
68
+ embeddings = [HalfVector([1.5, 2, 3]), HalfVector([4.5, 5, 6])]
69
+ cur.execute('INSERT INTO psycopg2_items (half_embeddings) VALUES (%s::halfvec[])', (embeddings,))
70
+
71
+ cur.execute('SELECT half_embeddings FROM psycopg2_items ORDER BY id')
72
+ res = cur.fetchone()
73
+ assert res[0][0].to_list() == [1.5, 2, 3]
74
+ assert res[0][1].to_list() == [4.5, 5, 6]
75
+
76
+ def test_sparsevec_array(self):
77
+ embeddings = [SparseVector([1.5, 2, 3]), SparseVector([4.5, 5, 6])]
78
+ cur.execute('INSERT INTO psycopg2_items (sparse_embeddings) VALUES (%s::sparsevec[])', (embeddings,))
79
+
80
+ cur.execute('SELECT sparse_embeddings FROM psycopg2_items ORDER BY id')
81
+ res = cur.fetchone()
82
+ assert res[0][0].to_list() == [1.5, 2, 3]
83
+ assert res[0][1].to_list() == [4.5, 5, 6]
84
+
58
85
  def test_cursor_factory(self):
59
86
  for cursor_factory in [DictCursor, RealDictCursor, NamedTupleCursor]:
60
87
  conn = psycopg2.connect(dbname='pgvector_python_test')
@@ -1,7 +1,7 @@
1
1
  import numpy as np
2
2
  from pgvector.sqlalchemy import VECTOR, HALFVEC, BIT, SPARSEVEC, SparseVector, avg, sum
3
3
  import pytest
4
- from sqlalchemy import create_engine, insert, inspect, select, text, MetaData, Table, Column, Index, Integer
4
+ from sqlalchemy import create_engine, insert, inspect, select, text, MetaData, Table, Column, Index, Integer, ARRAY
5
5
  from sqlalchemy.exc import StatementError
6
6
  from sqlalchemy.ext.automap import automap_base
7
7
  from sqlalchemy.orm import declarative_base, Session
@@ -31,6 +31,7 @@ class Item(Base):
31
31
  half_embedding = mapped_column(HALFVEC(3))
32
32
  binary_embedding = mapped_column(BIT(3))
33
33
  sparse_embedding = mapped_column(SPARSEVEC(3))
34
+ embeddings = mapped_column(ARRAY(VECTOR(3)))
34
35
 
35
36
 
36
37
  Base.metadata.drop_all(engine)
@@ -70,7 +71,8 @@ class TestSqlalchemy:
70
71
  Column('embedding', VECTOR(3)),
71
72
  Column('half_embedding', HALFVEC(3)),
72
73
  Column('binary_embedding', BIT(3)),
73
- Column('sparse_embedding', SPARSEVEC(3))
74
+ Column('sparse_embedding', SPARSEVEC(3)),
75
+ Column('embeddings', ARRAY(VECTOR(3)))
74
76
  )
75
77
 
76
78
  metadata.drop_all(engine)
@@ -422,6 +424,20 @@ class TestSqlalchemy:
422
424
  item = session.query(AutoItem).first()
423
425
  assert item.embedding.tolist() == [1, 2, 3]
424
426
 
427
+ def test_vector_array(self):
428
+ session = Session(engine)
429
+ session.add(Item(id=1, embeddings=[np.array([1, 2, 3]), np.array([4, 5, 6])]))
430
+ session.commit()
431
+
432
+ with engine.connect() as connection:
433
+ from pgvector.psycopg2 import register_vector
434
+ register_vector(connection.connection.dbapi_connection, globally=False, arrays=True)
435
+
436
+ # this fails if the driver does not cast arrays
437
+ item = Session(bind=connection).get(Item, 1)
438
+ assert item.embeddings[0].tolist() == [1, 2, 3]
439
+ assert item.embeddings[1].tolist() == [4, 5, 6]
440
+
425
441
  @pytest.mark.asyncio
426
442
  @pytest.mark.skipif(sqlalchemy_version == 1, reason='Requires SQLAlchemy 2+')
427
443
  async def test_async(self):
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes