summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-09-09 02:40:41 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-09-09 02:40:41 +0200
commit8fd3f8aeec85f47d43ef36a81456d48640ca33f9 (patch)
tree3c42c767dba661dfa28275379de35145f71611d8
parent08400210cfdfd5769ce2d48db5234fbeae2e0c6e (diff)
Add dict module
-rw-r--r--dict/dict.go85
-rw-r--r--dict/dict_test.go78
-rw-r--r--dict/go.mod3
-rw-r--r--go.work1
4 files changed, 167 insertions, 0 deletions
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