Skip to content

Commit a5ed69d

Browse files
committed
feat: add multi-vector config to HNSW
1 parent a2d8bfb commit a5ed69d

12 files changed

Lines changed: 430 additions & 33 deletions

File tree

src/it/java/io/weaviate/integration/SearchITest.java

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
import io.weaviate.ConcurrentTest;
2020
import io.weaviate.client6.v1.api.WeaviateApiException;
2121
import io.weaviate.client6.v1.api.WeaviateClient;
22+
import io.weaviate.client6.v1.api.collections.ObjectMetadata;
2223
import io.weaviate.client6.v1.api.collections.Property;
2324
import io.weaviate.client6.v1.api.collections.ReferenceProperty;
2425
import io.weaviate.client6.v1.api.collections.VectorConfig;
@@ -31,7 +32,10 @@
3132
import io.weaviate.client6.v1.api.collections.query.QueryMetadata;
3233
import io.weaviate.client6.v1.api.collections.query.QueryResponseGroup;
3334
import io.weaviate.client6.v1.api.collections.query.SortBy;
35+
import io.weaviate.client6.v1.api.collections.query.Target;
3436
import io.weaviate.client6.v1.api.collections.query.Where;
37+
import io.weaviate.client6.v1.api.collections.vectorindex.Hnsw;
38+
import io.weaviate.client6.v1.api.collections.vectorindex.MultiVector;
3539
import io.weaviate.containers.Container;
3640
import io.weaviate.containers.Container.ContainerGroup;
3741
import io.weaviate.containers.Contextionary;
@@ -499,4 +503,50 @@ public void testMetadataAll() throws IOException {
499503
Assertions.assertThat(metadataNearText.distance()).as("distance").isNotNull();
500504
Assertions.assertThat(metadataNearText.certainty()).as("certainty").isNotNull();
501505
}
506+
507+
@Test
508+
public void testNearVector_targetVectors() throws IOException {
509+
// Arrange
510+
var nsThings = ns("Things");
511+
512+
client.collections.create(nsThings,
513+
c -> c.vectorConfig(
514+
VectorConfig.selfProvided("v1d"),
515+
VectorConfig.selfProvided("v2d",
516+
none -> none
517+
.vectorIndex(Hnsw.of(
518+
hnsw -> hnsw.multiVector(MultiVector.of()))))));
519+
520+
var things = client.collections.use(nsThings);
521+
522+
var thing123 = things.data.insert(Map.of(), thing -> thing.vectors(
523+
Vectors.of("v1d", new float[] { 1, 2, 3 }),
524+
Vectors.of("v2d", new float[][] { { 1, 2, 3 }, { 1, 2, 3 } })));
525+
526+
var thing456 = things.data.insertMany(List.of(
527+
WeaviateObject.of(thing -> thing
528+
.metadata(ObjectMetadata.of(
529+
meta -> meta
530+
.vectors(
531+
Vectors.of("v1d", new float[] { 4, 5, 6 }),
532+
Vectors.of("v2d", new float[][] { { 4, 5, 6 }, { 4, 5, 6 } })))))));
533+
Assertions.assertThat(thing456.errors()).as("insert many").isEmpty();
534+
535+
// Act
536+
var got123 = things.query.nearVector(
537+
Target.vector("v1d", new float[] { 1, 2, 3 }),
538+
q -> q.limit(1));
539+
Assertions.assertThat(got123.objects())
540+
.as("search v1d")
541+
.hasSize(1).extracting(WeaviateObject::uuid)
542+
.containsExactly(thing123.uuid());
543+
544+
var got456 = things.query.nearVector(
545+
Target.vector("v2d", new float[][] { { 4, 5, 6 }, { 4, 5, 6 } }),
546+
q -> q.limit(1));
547+
Assertions.assertThat(got456.objects())
548+
.as("search v2d")
549+
.hasSize(1).extracting(WeaviateObject::uuid)
550+
.containsExactly(thing456.uuids().get(0));
551+
}
502552
}
Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,108 @@
1+
package io.weaviate.client6.v1.api.collections;
2+
3+
import java.io.IOException;
4+
import java.util.EnumMap;
5+
import java.util.Map;
6+
import java.util.function.Function;
7+
8+
import com.google.gson.Gson;
9+
import com.google.gson.JsonParser;
10+
import com.google.gson.TypeAdapter;
11+
import com.google.gson.TypeAdapterFactory;
12+
import com.google.gson.reflect.TypeToken;
13+
import com.google.gson.stream.JsonReader;
14+
import com.google.gson.stream.JsonWriter;
15+
16+
import io.weaviate.client6.v1.api.collections.encoding.MuveraEncoding;
17+
import io.weaviate.client6.v1.internal.ObjectBuilder;
18+
import io.weaviate.client6.v1.internal.json.JsonEnum;
19+
20+
public interface Encoding {
21+
22+
enum Kind implements JsonEnum<Kind> {
23+
MUVERA("muvera");
24+
25+
private static final Map<String, Kind> jsonValueMap = JsonEnum.collectNames(Kind.values());
26+
private final String jsonValue;
27+
28+
private Kind(String jsonValue) {
29+
this.jsonValue = jsonValue;
30+
}
31+
32+
@Override
33+
public String jsonValue() {
34+
return this.jsonValue;
35+
}
36+
37+
public static Kind valueOfJson(String jsonValue) {
38+
return JsonEnum.valueOfJson(jsonValue, jsonValueMap, Kind.class);
39+
}
40+
}
41+
42+
Kind _kind();
43+
44+
Object _self();
45+
46+
public static Encoding muvera() {
47+
return MuveraEncoding.of();
48+
}
49+
50+
public static Encoding muvera(Function<MuveraEncoding.Builder, ObjectBuilder<MuveraEncoding>> fn) {
51+
return MuveraEncoding.of(fn);
52+
}
53+
54+
public enum CustomTypeAdapterFactory implements TypeAdapterFactory {
55+
INSTANCE;
56+
57+
private static final EnumMap<Encoding.Kind, TypeAdapter<? extends Encoding>> delegateAdapters = new EnumMap<>(
58+
Encoding.Kind.class);
59+
60+
private final void addAdapter(Gson gson, Encoding.Kind kind, Class<? extends Encoding> cls) {
61+
delegateAdapters.put(kind,
62+
(TypeAdapter<? extends Encoding>) gson.getDelegateAdapter(this, TypeToken.get(cls)));
63+
}
64+
65+
private final void init(Gson gson) {
66+
addAdapter(gson, Encoding.Kind.MUVERA, MuveraEncoding.class);
67+
}
68+
69+
@SuppressWarnings("unchecked")
70+
@Override
71+
public <T> TypeAdapter<T> create(Gson gson, TypeToken<T> type) {
72+
final var rawType = type.getRawType();
73+
if (!Encoding.class.isAssignableFrom(rawType)) {
74+
return null;
75+
}
76+
77+
if (delegateAdapters.isEmpty()) {
78+
init(gson);
79+
}
80+
81+
return (TypeAdapter<T>) new TypeAdapter<Encoding>() {
82+
83+
@Override
84+
public void write(JsonWriter out, Encoding value) throws IOException {
85+
TypeAdapter<T> adapter = (TypeAdapter<T>) delegateAdapters.get(value._kind());
86+
adapter.write(out, (T) value._self());
87+
}
88+
89+
@Override
90+
public Encoding read(JsonReader in) throws IOException {
91+
var encodingObject = JsonParser.parseReader(in).getAsJsonObject();
92+
var encodingName = encodingObject.keySet().iterator().next();
93+
94+
Encoding.Kind kind;
95+
try {
96+
kind = Encoding.Kind.valueOfJson(encodingName);
97+
} catch (IllegalArgumentException e) {
98+
return null;
99+
}
100+
101+
var adapter = delegateAdapters.get(kind);
102+
var concreteEncoding = encodingObject.get(encodingName).getAsJsonObject();
103+
return adapter.fromJsonTree(concreteEncoding);
104+
}
105+
}.nullSafe();
106+
}
107+
}
108+
}

src/main/java/io/weaviate/client6/v1/api/collections/ObjectMetadata.java

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,12 +17,16 @@ public ObjectMetadata(Builder builder) {
1717
this(builder.uuid, builder.vectors, null, null);
1818
}
1919

20+
public static ObjectMetadata of() {
21+
return of(ObjectBuilder.identity());
22+
}
23+
2024
public static ObjectMetadata of(Function<Builder, ObjectBuilder<ObjectMetadata>> fn) {
2125
return fn.apply(new Builder()).build();
2226
}
2327

2428
public static class Builder implements ObjectBuilder<ObjectMetadata> {
25-
private String uuid;
29+
private String uuid = UUID.randomUUID().toString();
2630
private Vectors vectors;
2731

2832
/** Assign a custom UUID for the object. */

src/main/java/io/weaviate/client6/v1/api/collections/Quantization.java

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -148,7 +148,6 @@ public <T> TypeAdapter<T> create(Gson gson, TypeToken<T> type) {
148148
@Override
149149
public void write(JsonWriter out, Quantization value) throws IOException {
150150
if (value._kind() == Quantization.Kind.UNCOMPRESSED) {
151-
// out.name(value._kind().jsonValue());
152151
out.value(true);
153152
return;
154153
}

src/main/java/io/weaviate/client6/v1/api/collections/data/InsertManyRequest.java

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -31,9 +31,7 @@ public InsertManyRequest(WeaviateObject<T, Reference, ObjectMetadata>... objects
3131
public static final <T> InsertManyRequest<T> of(T... properties) {
3232
var objects = Arrays.stream(properties)
3333
.map(p -> WeaviateObject.<T, Reference, ObjectMetadata>of(
34-
obj -> obj
35-
.properties(p)
36-
.metadata(ObjectMetadata.of(m -> m.uuid(UUID.randomUUID())))))
34+
obj -> obj.properties(p).metadata(ObjectMetadata.of())))
3735
.toList();
3836
return new InsertManyRequest<T>(objects);
3937
}
@@ -101,9 +99,7 @@ public static <T> void buildObject(WeaviateProtoBatch.BatchObject.Builder object
10199

102100
var metadata = insert.metadata();
103101
if (metadata != null) {
104-
if (metadata.uuid() != null) {
105-
object.setUuid(metadata.uuid());
106-
}
102+
object.setUuid(metadata.uuid());
107103

108104
if (metadata.vectors() != null) {
109105
var vectors = metadata.vectors().asMap()
@@ -156,11 +152,15 @@ public static <T> void buildObject(WeaviateProtoBatch.BatchObject.Builder object
156152
}
157153
});
158154

159-
var nonRef = marshalStruct(collection.propertiesReader(insert.properties()).readProperties());
160-
object.setProperties(WeaviateProtoBatch.BatchObject.Properties.newBuilder()
161-
.setNonRefProperties(nonRef)
155+
var properties = WeaviateProtoBatch.BatchObject.Properties.newBuilder()
162156
.addAllSingleTargetRefProps(singleRef)
163-
.addAllMultiTargetRefProps(multiRef));
157+
.addAllMultiTargetRefProps(multiRef);
158+
159+
if (insert.properties() != null) {
160+
var nonRef = marshalStruct(collection.propertiesReader(insert.properties()).readProperties());
161+
properties.setNonRefProperties(nonRef);
162+
}
163+
object.setProperties(properties);
164164
}
165165

166166
@SuppressWarnings("unchecked")
Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,73 @@
1+
package io.weaviate.client6.v1.api.collections.encoding;
2+
3+
import java.util.function.Function;
4+
5+
import com.google.gson.annotations.SerializedName;
6+
7+
import io.weaviate.client6.v1.api.collections.Encoding;
8+
import io.weaviate.client6.v1.internal.ObjectBuilder;
9+
10+
public record MuveraEncoding(
11+
@SerializedName("enabled") boolean enabled,
12+
@SerializedName("ksim") Integer ksim,
13+
@SerializedName("dprojections") Integer dprojections,
14+
@SerializedName("repetitions") Integer repetitions) implements Encoding {
15+
16+
public static MuveraEncoding of() {
17+
return of(ObjectBuilder.identity());
18+
}
19+
20+
public static MuveraEncoding of(Function<Builder, ObjectBuilder<MuveraEncoding>> fn) {
21+
return fn.apply(new Builder()).build();
22+
}
23+
24+
public MuveraEncoding(Builder builder) {
25+
this(
26+
builder.enabled,
27+
builder.ksim,
28+
builder.dprojections,
29+
builder.repetitions);
30+
}
31+
32+
public static class Builder implements ObjectBuilder<MuveraEncoding> {
33+
private boolean enabled = true;
34+
private Integer ksim;
35+
private Integer dprojections;
36+
private Integer repetitions;
37+
38+
public Builder enabled(boolean enabled) {
39+
this.enabled = enabled;
40+
return this;
41+
}
42+
43+
public Builder ksim(int ksim) {
44+
this.ksim = ksim;
45+
return this;
46+
}
47+
48+
public Builder dprojections(int dprojections) {
49+
this.dprojections = dprojections;
50+
return this;
51+
}
52+
53+
public Builder repetitions(int repetitions) {
54+
this.repetitions = repetitions;
55+
return this;
56+
}
57+
58+
@Override
59+
public MuveraEncoding build() {
60+
return new MuveraEncoding(this);
61+
}
62+
}
63+
64+
@Override
65+
public Encoding.Kind _kind() {
66+
return Encoding.Kind.MUVERA;
67+
}
68+
69+
@Override
70+
public Object _self() {
71+
return this;
72+
}
73+
}

src/main/java/io/weaviate/client6/v1/api/collections/query/Target.java

Lines changed: 6 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -36,12 +36,11 @@ public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
3636
}
3737
req.addTargetVectors(vectorName);
3838

39-
var weightsForTarget = WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
40-
.setTarget(vectorName);
4139
if (weight != null) {
42-
weightsForTarget.setWeight(weight);
40+
req.addWeightsForTargets(WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
41+
.setTarget(vectorName)
42+
.setWeight(weight));
4343
}
44-
req.addWeightsForTargets(weightsForTarget);
4544
return true;
4645
}
4746

@@ -226,12 +225,11 @@ public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
226225
}
227226
req.addTargetVectors(vectorName);
228227

229-
var weightsForTarget = WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
230-
.setTarget(vectorName);
231228
if (weight != null) {
232-
weightsForTarget.setWeight(weight);
229+
req.addWeightsForTargets(WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
230+
.setTarget(vectorName)
231+
.setWeight(weight));
233232
}
234-
req.addWeightsForTargets(weightsForTarget);
235233
return true;
236234
}
237235
}

src/main/java/io/weaviate/client6/v1/api/collections/vectorindex/Hnsw.java

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@ public record Hnsw(
1515
@SerializedName("vectorCacheMaxObjects") Long vectorCacheMaxObjects,
1616
@SerializedName("cleanupIntervalSeconds") Integer cleanupIntervalSeconds,
1717
@SerializedName("filterStrategy") FilterStrategy filterStrategy,
18+
@SerializedName("multivector") MultiVector multiVector,
1819

1920
@SerializedName("dynamicEfMin") Integer dynamicEfMin,
2021
@SerializedName("dynamicEfMax") Integer dynamicEfMax,
@@ -49,6 +50,7 @@ public Hnsw(Builder builder) {
4950
builder.vectorCacheMaxObjects,
5051
builder.cleanupIntervalSeconds,
5152
builder.filterStrategy,
53+
builder.multiVector,
5254
builder.dynamicEfMin,
5355
builder.dynamicEfMax,
5456
builder.dynamicEfFactor,
@@ -64,6 +66,7 @@ public static class Builder implements ObjectBuilder<Hnsw> {
6466
private Long vectorCacheMaxObjects;
6567
private Integer cleanupIntervalSeconds;
6668
private FilterStrategy filterStrategy;
69+
private MultiVector multiVector;
6770

6871
private Integer dynamicEfMin;
6972
private Integer dynamicEfMax;
@@ -106,6 +109,11 @@ public final Builder filterStrategy(FilterStrategy filterStrategy) {
106109
return this;
107110
}
108111

112+
public final Builder multiVector(MultiVector multiVector) {
113+
this.multiVector = multiVector;
114+
return this;
115+
}
116+
109117
public final Builder dynamicEfMin(int dynamicEfMin) {
110118
this.dynamicEfMin = dynamicEfMin;
111119
return this;

0 commit comments

Comments
 (0)