diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:39:02 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:39:02 +0200 |
| commit | 9e1b8c4bde9b0263a4c4d2278e3c283ca0eb07f1 (patch) | |
| tree | bf6b604a150a196490b92f89ca40a335aca83d31 /research/knobloch/shared.go | |
| parent | 75e581b3bc19a73d0c1f48b32f2e58121c9741c1 (diff) | |
Diffstat (limited to 'research/knobloch/shared.go')
| -rw-r--r-- | research/knobloch/shared.go | 152 |
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 +} |
