From 8fd3f8aeec85f47d43ef36a81456d48640ca33f9 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Wed, 9 Sep 2026 02:40:41 +0200 Subject: Add dict module --- dict/dict.go | 85 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ dict/dict_test.go | 78 ++++++++++++++++++++++++++++++++++++++++++++++++++ dict/go.mod | 3 ++ go.work | 1 + 4 files changed, 167 insertions(+) create mode 100644 dict/dict.go create mode 100644 dict/dict_test.go create mode 100644 dict/go.mod diff --git a/dict/dict.go b/dict/dict.go new file mode 100644 index 0000000..bc0f8bd --- /dev/null +++ b/dict/dict.go @@ -0,0 +1,85 @@ +package dict + +import ( + "iter" + "slices" +) + +type Entry[V any] struct { + Key string + Val V +} +type Dict[V any] struct { + s []Entry[V] + m map[string]int +} + +func NewDict[V any]() *Dict[V] { + return &Dict[V]{ + s: make([]Entry[V], 0), + m: make(map[string]int), + } +} + +func (d *Dict[V]) Set(key string, val V) { + if i, ok := d.m[key]; ok { + d.s[i].Val = val + + return + } + + d.s = append(d.s, Entry[V]{ + Key: key, + Val: val, + }) + + d.m[key] = len(d.s) - 1 +} + +func (d *Dict[V]) Get(key string) (V, bool) { + i, ok := d.m[key] + + if !ok { + var zero V + + return zero, false + } + + return d.s[i].Val, true +} + +func (d *Dict[V]) GetBytes(key []byte) (V, bool) { + if i, ok := d.m[string(key)]; ok { + return d.s[i].Val, true + } + + var zero V + + return zero, false +} + +func (d *Dict[V]) Values() iter.Seq2[int, V] { + return func(yield func(int, V) bool) { + for i, entry := range d.s { + if !yield(i, entry.Val) { + return + } + } + } +} + +func (d *Dict[V]) Sort(cmp func(a, b Entry[V]) int) { + slices.SortFunc(d.s, cmp) + + for i, entry := range d.s { + d.m[entry.Key] = i + } +} + +// func CmpKey[V any](a, b Entry[V]) int { +// return strings.Compare(a.Key, b.Key) +// } + +// func CmpVal[V cmp.Ordered](a, b Entry[V]) int { +// return cmp.Compare(a.Val, b.Val) +// } diff --git a/dict/dict_test.go b/dict/dict_test.go new file mode 100644 index 0000000..c5fb717 --- /dev/null +++ b/dict/dict_test.go @@ -0,0 +1,78 @@ +package dict + +import ( + "fmt" + "slices" + "strings" + "testing" +) + +func TestDict_Sort(t *testing.T) { + d := NewDict[string]() + + d.Set("foo", "foo") + d.Set("bar", "bar") + d.Set("baz", "baz") + + d.Sort(func(a, b Entry[string]) int { + return strings.Compare(a.Key, b.Key) + }) + + expected := []Entry[string]{ + {"bar", "bar"}, + {"baz", "baz"}, + {"foo", "foo"}, + } + + verifyMap(t, d) + + if !slices.Equal(d.s, expected) { + t.Fatalf("expected %v but got %v", expected, d.s) + } +} + +func verifyMap(t *testing.T, dict *Dict[string]) { + if len(dict.m) != len(dict.s) { + t.Fatalf("length missmatch") + } + + for i, s := range dict.s { + v, ok := dict.m[s.Key] + + if !ok { + t.Fatalf("unknown key %s", s) + } + + if v != i { + t.Errorf("expected %d but got %d\n", i, v) + } + } +} + +func BenchmarkDict_Get(b *testing.B) { + d := NewDict[int]() + + for i := 0; i < 1000; i++ { + d.Set(fmt.Sprintf("key_%d", i), i) + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = d.Get("key_500") + } +} + +func BenchmarkDict_GetBytes(b *testing.B) { + d := NewDict[int]() + + for i := 0; i < 1000; i++ { + d.Set(fmt.Sprintf("key_%d", i), i) + } + + keyBytes := []byte("key_500") + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = d.GetBytes(keyBytes) + } +} diff --git a/dict/go.mod b/dict/go.mod new file mode 100644 index 0000000..23295f8 --- /dev/null +++ b/dict/go.mod @@ -0,0 +1,3 @@ +module go.jknobloc.com/x/dict + +go 1.25 diff --git a/go.work b/go.work index 9cdd338..b4654b7 100644 --- a/go.work +++ b/go.work @@ -2,6 +2,7 @@ go 1.25.0 use ( ./dataset + ./dict ./gpt2 ./llm ./llmc -- cgit v1.2.3