forked from marcboeker/go-duckdb
-
Notifications
You must be signed in to change notification settings - Fork 41
Expand file tree
/
Copy pathvector_udf_benchmark_test.go
More file actions
143 lines (126 loc) · 3.68 KB
/
Copy pathvector_udf_benchmark_test.go
File metadata and controls
143 lines (126 loc) · 3.68 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
package duckdb
import (
"context"
"database/sql/driver"
"fmt"
"strings"
"testing"
"unicode"
"github.com/stretchr/testify/require"
)
type varcharTransformBenchmarkUDF struct {
info TypeInfo
useVector bool
}
func (udf *varcharTransformBenchmarkUDF) Config() ScalarFuncConfig {
return ScalarFuncConfig{
InputTypeInfos: []TypeInfo{udf.info},
ResultTypeInfo: udf.info,
}
}
func (udf *varcharTransformBenchmarkUDF) Executor() ScalarFuncExecutor {
if !udf.useVector {
return ScalarFuncExecutor{
RowExecutor: func(values []driver.Value) (any, error) {
return transformBenchmarkVarchar(values[0].(string)), nil
},
}
}
return ScalarFuncExecutor{
ChunkContextExecutor: func(_ context.Context, state *ChunkIteratorState) error {
inputVector, err := state.GetInputChunk().GetVector(0)
if err != nil {
return err
}
input, err := GetVectorView[string](inputVector)
if err != nil {
return err
}
output, err := GetVectorWriter[string](state.GetResultVector())
if err != nil {
return err
}
for row := range input.Len() {
value, valid, err := input.GetValueBorrowed(row)
if err != nil {
return err
}
if !valid {
// Default NULL handling sets this result to NULL after the callback.
continue
}
if err = output.Set(row, transformBenchmarkVarchar(value)); err != nil {
return err
}
}
return nil
},
}
}
func transformBenchmarkVarchar(value string) string {
value = strings.TrimSpace(value)
return strings.Map(func(r rune) rune {
if r == '-' {
return '_'
}
return unicode.ToUpper(r)
}, value)
}
var varcharTransformBenchmarkSink int64
const varcharTransformBenchmarkRowCount = 2_000_000
func BenchmarkVarcharTransformUDF(b *testing.B) {
db := openDbWrapper(b, ``)
defer closeDbWrapper(b, db)
conn := openConnWrapper(b, db, context.Background())
defer closeConnWrapper(b, conn)
_, err := conn.ExecContext(context.Background(), `SET threads = 1`)
require.NoError(b, err)
info := mustTypeInfo(b, TYPE_VARCHAR)
require.NoError(b, RegisterScalarUDF(conn, "varchar_rows_transform", &varcharTransformBenchmarkUDF{
info: info,
}))
require.NoError(b, RegisterScalarUDF(conn, "varchar_vector_transform", &varcharTransformBenchmarkUDF{
info: info,
useVector: true,
}))
_, err = conn.ExecContext(context.Background(), fmt.Sprintf(`
CREATE TABLE varchar_transform_benchmark AS
SELECT CASE i %% 6
WHEN 0 THEN NULL::VARCHAR
WHEN 1 THEN ''
WHEN 2 THEN ' customer-' || i::VARCHAR || '-alpha '
WHEN 3 THEN repeat('duckdb-go-', 4) || i::VARCHAR
WHEN 4 THEN ' München-' || i::VARCHAR || '-straße '
ELSE repeat('long-varchar-value-', 8) || i::VARCHAR
END AS value
FROM range(%d) values(i)
`, varcharTransformBenchmarkRowCount))
require.NoError(b, err)
benchmarks := []struct {
name string
function string
}{
{name: "RowExecutor", function: "varchar_rows_transform"},
{name: "ChunkVector", function: "varchar_vector_transform"},
}
for _, benchmark := range benchmarks {
b.Run(benchmark.name, func(b *testing.B) {
query := `SELECT coalesce(sum(length(` + benchmark.function +
`(value))), 0) FROM varchar_transform_benchmark`
stmt, prepareErr := conn.PrepareContext(context.Background(), query)
require.NoError(b, prepareErr)
defer func() {
require.NoError(b, stmt.Close())
}()
require.NoError(b, stmt.QueryRowContext(context.Background()).Scan(&varcharTransformBenchmarkSink))
b.ReportAllocs()
b.ResetTimer()
for b.Loop() {
if err := stmt.QueryRowContext(context.Background()).Scan(&varcharTransformBenchmarkSink); err != nil {
b.Fatal(err)
}
}
b.ReportMetric(varcharTransformBenchmarkRowCount, "rows/op")
})
}
}