summaryrefslogtreecommitdiff
path: root/research/knobloch/shared.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-09-11 18:39:02 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-09-11 18:39:02 +0200
commit9e1b8c4bde9b0263a4c4d2278e3c283ca0eb07f1 (patch)
treebf6b604a150a196490b92f89ca40a335aca83d31 /research/knobloch/shared.go
parent75e581b3bc19a73d0c1f48b32f2e58121c9741c1 (diff)
Diffstat (limited to 'research/knobloch/shared.go')
-rw-r--r--research/knobloch/shared.go152
1 files changed, 152 insertions, 0 deletions
diff --git a/research/knobloch/shared.go b/research/knobloch/shared.go
new file mode 100644
index 0000000..5656668
--- /dev/null
+++ b/research/knobloch/shared.go
@@ -0,0 +1,152 @@
+package knobloch
+
+import (
+ "encoding/json"
+ "fmt"
+ "os"
+ "sync"
+
+ "go.jknobloc.com/x/shelf"
+ "go.jknobloc.com/x/tokenizer/bpe"
+)
+
+// SharedVocabs lists the vocabularies intersected for the *_shared plots, which
+// restrict every figure to tokens all of these models have in common. Set it to
+// the family being swept so a model is only compared against its own kind; an
+// empty list skips those plots.
+var SharedVocabs = []shelf.Item{
+ "models/mbpe/minipile/gpt2_50256_m000_minipile/vocab.json",
+ "models/mbpe/minipile/gpt2_50256_m030_minipile/vocab.json",
+ "models/mbpe/minipile/gpt2_50256_m050_minipile/vocab.json",
+ "models/mbpe/minipile/gpt2_50256_m100_minipile/vocab.json",
+}
+
+var sharedCache struct {
+ mu sync.Mutex
+ m map[string]map[string]struct{}
+}
+
+// sharedTokens caches the intersection per vocabulary list, so a sweep that
+// switches SharedVocabs between families gets each family's own intersection
+// while still reading the vocabularies only once per family.
+func sharedTokens() (map[string]struct{}, error) {
+ key := ""
+
+ for _, v := range SharedVocabs {
+ key += string(v) + "\n"
+ }
+
+ sharedCache.mu.Lock()
+
+ defer sharedCache.mu.Unlock()
+
+ if tokens, ok := sharedCache.m[key]; ok {
+ return tokens, nil
+ }
+
+ tokens, err := intersectVocabs(SharedVocabs)
+
+ if err != nil {
+ return nil, err
+ }
+
+ if sharedCache.m == nil {
+ sharedCache.m = make(map[string]map[string]struct{})
+ }
+
+ sharedCache.m[key] = tokens
+
+ return tokens, nil
+}
+
+func intersectVocabs(vocabs []shelf.Item) (map[string]struct{}, error) {
+ if len(vocabs) == 0 {
+ return nil, fmt.Errorf("no vocabularies to intersect")
+ }
+
+ var keep map[string]struct{}
+
+ for _, v := range vocabs {
+ tokens, err := loadVocab(shelf.Abs(v))
+
+ if err != nil {
+ return nil, err
+ }
+
+ if keep == nil {
+ keep = tokens
+
+ continue
+ }
+
+ for token := range keep {
+ if _, ok := tokens[token]; !ok {
+ delete(keep, token)
+ }
+ }
+ }
+
+ return keep, nil
+}
+
+func loadVocab(name string) (map[string]struct{}, error) {
+ file, err := os.Open(name)
+
+ if err != nil {
+ return nil, err
+ }
+
+ defer file.Close()
+
+ var m map[string]int64
+
+ if err := json.NewDecoder(file).Decode(&m); err != nil {
+ return nil, err
+ }
+
+ tokens := make(map[string]struct{}, len(m))
+
+ for token := range m {
+ tokens[token] = struct{}{}
+ }
+
+ return tokens, nil
+}
+
+// sharedMask selects the vocab ids on one side of the intersection: shared marks
+// the tokens every vocabulary in SharedVocabs has, and its complement isolates
+// what a listed model added on its own.
+func sharedMask(n int, t *bpe.Tokenizer, shared bool) ([]bool, error) {
+ keep, err := sharedTokens()
+
+ if err != nil {
+ return nil, err
+ }
+
+ itoa := bpe.Itoa(t)
+
+ mask := make([]bool, n)
+
+ for id := range mask {
+ _, ok := keep[itoa[int64(id)]]
+
+ mask[id] = ok == shared
+ }
+
+ return mask, nil
+}
+
+// maskCounts zeroes the counts outside the mask. Every plot skips zero counts
+// already, so the existing set of figures works unchanged while the axes stay
+// pinned to the full vocabulary.
+func maskCounts(r []int, keep []bool) []int {
+ masked := make([]int, len(r))
+
+ for id, v := range r {
+ if keep[id] {
+ masked[id] = v
+ }
+ }
+
+ return masked
+}