Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 14 additions & 2 deletions benchmarker/cmd/ann_benchmark.go
Original file line number Diff line number Diff line change
Expand Up @@ -217,6 +217,13 @@ func createClient(cfg *Config) *weaviate.Client {
return client
}

func applyRQCentering(rqConfig map[string]interface{}, centering bool, trainingLimit int) {
if centering {
rqConfig["centering"] = true
rqConfig["trainingLimit"] = trainingLimit
}
}

// Re/create Weaviate schema
func createSchema(cfg *Config, client *weaviate.Client) {
err := client.Schema().ClassDeleter().WithClassName(cfg.ClassName).Do(context.Background())
Expand Down Expand Up @@ -302,6 +309,7 @@ func createSchema(cfg *Config, client *weaviate.Client) {
if cfg.RescoreLimit > -1 {
rqConfig["rescoreLimit"] = cfg.RescoreLimit
}
applyRQCentering(rqConfig, cfg.RQCentering, cfg.TrainingLimit)
vectorIndexConfig = map[string]interface{}{
"distance": cfg.DistanceMetric,
"efConstruction": float64(cfg.EfConstruction),
Expand Down Expand Up @@ -462,6 +470,7 @@ func createSchema(cfg *Config, client *weaviate.Client) {
if cfg.RescoreLimit > -1 {
rqConfig["rescoreLimit"] = cfg.RescoreLimit
}
applyRQCentering(rqConfig, cfg.RQCentering, cfg.TrainingLimit)

vectorIndexConfig = map[string]interface{}{
"distance": cfg.DistanceMetric,
Expand Down Expand Up @@ -702,6 +711,7 @@ func enableCompression(cfg *Config, client *weaviate.Client, dimensions uint, co
if cfg.RescoreLimit > -1 {
rqConfig["rescoreLimit"] = cfg.RescoreLimit
}
applyRQCentering(rqConfig, cfg.RQCentering, cfg.TrainingLimit)
vectorIndexConfig["rq"] = rqConfig
}

Expand Down Expand Up @@ -1119,7 +1129,9 @@ func initAnnBenchmark() {
annBenchmarkCommand.PersistentFlags().StringVar(&globalConfig.RQ,
"rq", "disabled", "Set RQ (disabled, auto, or enabled) (default disabled)")
annBenchmarkCommand.PersistentFlags().UintVar(&globalConfig.RQBits,
"rqBits", 8, "Set RQ bits (default 8)")
"rqBits", 8, "Set RQ bits: 1, 4 or 8; 4 requires indexType hnsw (default 8)")
annBenchmarkCommand.PersistentFlags().BoolVar(&globalConfig.RQCentering,
"rqCentering", false, "Enable RQ centering (requires rqBits=4)")
annBenchmarkCommand.PersistentFlags().IntVarP(&globalConfig.MultiVectorDimensions,
"multiVector", "m", 0, "Enable multi-dimensional vectors with the specified number of dimensions")
annBenchmarkCommand.PersistentFlags().BoolVar(&globalConfig.MuveraEnabled,
Expand All @@ -1139,7 +1151,7 @@ func initAnnBenchmark() {
annBenchmarkCommand.PersistentFlags().BoolVar(&globalConfig.SkipTombstonesEmpty,
"skipTombstonesEmpty", false, "Skip waiting for tombstone to be empty after update (default false)")
annBenchmarkCommand.PersistentFlags().IntVar(&globalConfig.TrainingLimit,
"trainingLimit", 100000, "Set PQ trainingLimit (default 100000)")
"trainingLimit", 0, "Set compression trainingLimit (default 10000 for 4-bit RQ, 100000 otherwise)")
annBenchmarkCommand.PersistentFlags().IntVar(&globalConfig.EfConstruction,
"efConstruction", 256, "Set Weaviate efConstruction parameter (default 256)")
annBenchmarkCommand.PersistentFlags().StringVar(&globalConfig.EfArray,
Expand Down
25 changes: 24 additions & 1 deletion benchmarker/cmd/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ type Config struct {
SQ string
RQ string
RQBits uint
RQCentering bool
SkipQuery bool
SkipAsyncReady bool
SkipTombstonesEmpty bool
Expand Down Expand Up @@ -173,7 +174,7 @@ func (c *Config) parseLabels() {
c.LabelMap = result
}

func (c Config) validateANN() error {
func (c *Config) validateANN() error {
if c.BenchmarkFile == "" && c.DatasetRepo == "" {
return errors.Errorf("a vector benchmark file or a dataset repository and dataset must be provided")
}
Expand All @@ -190,5 +191,27 @@ func (c Config) validateANN() error {
return errors.Errorf("distance metric must be set")
}

if c.RQ != "disabled" {
switch c.RQBits {
case 1, 4, 8:
default:
return errors.Errorf("rqBits must be 1, 4 or 8, got %d", c.RQBits)
}
if c.RQBits == 4 && c.IndexType != "hnsw" {
return errors.Errorf("rqBits=4 is only supported with indexType hnsw, got %q", c.IndexType)
}
if c.RQCentering && c.RQBits != 4 {
return errors.Errorf("rqCentering requires rqBits=4, got %d", c.RQBits)
}
}

if c.TrainingLimit == 0 {
if c.RQ != "disabled" && c.RQBits == 4 {
c.TrainingLimit = 10000
} else {
c.TrainingLimit = 100000
}
}

return nil
}
Loading