diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-30 14:43:28 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-30 14:43:28 +0200 |
| commit | f9332fe85bf12433286fc77456ae2705c494a4f9 (patch) | |
| tree | 5f254e600b3916478302580383fd2b5611af122c | |
| parent | cc7aa68fa63cb69cb3a565364ee0bee1f482f6f5 (diff) | |
Add knobloch module
| -rw-r--r-- | go.work | 1 | ||||
| m--------- | mbpe | 0 | ||||
| -rw-r--r-- | research/knobloch/cmd/train/main.go | 124 | ||||
| -rw-r--r-- | research/knobloch/cmd/train/serialize.go | 108 | ||||
| -rw-r--r-- | research/knobloch/go.mod | 3 |
5 files changed, 236 insertions, 0 deletions
@@ -7,6 +7,7 @@ use ( ./llmc ./mbpe ./onnx + ./research/knobloch ./research/lesci ./research/sander ./shelf diff --git a/mbpe b/mbpe -Subproject d6c99eeabc16e7c58dcc179cf4cbf6e72d556db +Subproject f958f46bd12abd1d740b9c733d9fcb9f7cf7816 diff --git a/research/knobloch/cmd/train/main.go b/research/knobloch/cmd/train/main.go new file mode 100644 index 0000000..a946ce4 --- /dev/null +++ b/research/knobloch/cmd/train/main.go @@ -0,0 +1,124 @@ +package main + +import ( + "errors" + "fmt" + "log" + "os" + "path/filepath" + + "github.com/jonasknobloch/mbpe" + + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tokenizer/bpe" + "go.jknobloc.com/x/tokenizer/bpe/split" +) + +func main() { + // train() + // serialize() +} + +func train() { + out := shelf.Abs("results/knobloch/minipile") + + if err := os.MkdirAll(out, os.ModePerm); err != nil { + log.Fatal(err) + } + + morfessor := func(alpha float64) mbpe.Segmenter { + m := mbpe.NewMorfessor(alpha) + + if err := m.LoadModel(shelf.Abs("morfessor/semisup_model.proto")); err != nil { + log.Fatal(err) + } + + return m + } + + // mbpe.InvertWeightFunction = true + + m000 := morfessor(0.0) + m010 := morfessor(0.1) + m020 := morfessor(0.2) + m030 := morfessor(0.3) + m040 := morfessor(0.4) + m050 := morfessor(0.5) + m060 := morfessor(0.6) + m070 := morfessor(0.7) + m080 := morfessor(0.8) + m090 := morfessor(0.9) + m100 := morfessor(1.0) + + newTrainer := func(segmenter mbpe.Segmenter) *mbpe.MBPETrainer { + alphabet := make(map[string]struct{}) + + for _, r := range bpe.InitialAlphabet() { + alphabet[string(r)] = struct{}{} + } + + b := mbpe.NewByteLevel(true) + + b.SetMatcher(split.NewFSA()) + + return mbpe.NewMBPETrainer(b, segmenter, mbpe.NewMBPE(), 1<<17, alphabet) + } + + trainers := []struct { + *mbpe.MBPETrainer + string + }{ + {newTrainer(m000), "m000_minipile"}, + {newTrainer(m010), "m010_minipile"}, + {newTrainer(m020), "m020_minipile"}, + {newTrainer(m030), "m030_minipile"}, + {newTrainer(m040), "m040_minipile"}, + {newTrainer(m050), "m050_minipile"}, + {newTrainer(m060), "m060_minipile"}, + {newTrainer(m070), "m070_minipile"}, + {newTrainer(m080), "m080_minipile"}, + {newTrainer(m090), "m090_minipile"}, + {newTrainer(m100), "m100_minipile"}, + } + + for i, t := range trainers { + dict := filepath.Join(out, "dict.txt") + + if dictErr := t.LoadDict(dict); dictErr != nil { + var reader dataset.Reader + + if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/train")); err != nil { + log.Fatal(err) + } else { + reader = r + } + + if err := t.InitDict(reader); err != nil { + log.Fatal(err) + } + + if err := t.Dict().Save(dict); err != nil { + log.Fatal(err) + } + } + + if i > 0 { + fmt.Println() + } + + fmt.Printf("%s\n\n", t.string) + + t.Train() + + dir := filepath.Join(out, t.string) + + if err := os.Mkdir(dir, 0755); err != nil && !errors.Is(err, os.ErrExist) { + log.Fatal(err) + } + + if err := t.Model().Save(filepath.Join(dir, "vocab.json"), filepath.Join(dir, "merges.txt")); err != nil { + log.Fatal(err) + } + } +} diff --git a/research/knobloch/cmd/train/serialize.go b/research/knobloch/cmd/train/serialize.go new file mode 100644 index 0000000..9fbf968 --- /dev/null +++ b/research/knobloch/cmd/train/serialize.go @@ -0,0 +1,108 @@ +package main + +import ( + "fmt" + "log" + "os" + "path/filepath" + "strings" + + "github.com/jonasknobloch/mbpe" + + "go.jknobloc.com/x/shelf" +) + +func serialize() { + base := shelf.Abs("results/knobloch/minipile") + + var paths []string + + if ps, err := subDirs(base); err != nil { + log.Fatal(err) + } else { + paths = ps + } + + steps := []int{100512, 50256, 32768, 16384, 8192} + + outRoot := shelf.Abs("tokenizers") + + if err := os.MkdirAll(outRoot, os.ModePerm); err != nil { + log.Fatal(err) + } + + for _, step := range steps { + for _, path := range paths { + model := mbpe.NewMBPE() + + if err := model.Load(filepath.Join(path, "vocab.json"), filepath.Join(path, "merges.txt")); err != nil { + log.Fatal(err) + } + + model.Trim(step) + + dir := fmt.Sprintf("tokenizer_gpt2_%d_%s", step, filepath.Base(path)) + + out := filepath.Join(outRoot, dir) + + if err := os.Mkdir(out, os.ModePerm); err != nil { + log.Fatal(err) + } + + if err := model.Save(filepath.Join(out, "vocab.json"), filepath.Join(out, "merges.txt")); err != nil { + log.Fatal(err) + } + + var config string + var special string + + if bs, err := os.ReadFile(shelf.Abs("models/gpt2/tokenizer_config.json")); err != nil { + log.Fatal(bs) + } else { + config = string(bs) + } + + if bs, err := os.ReadFile(shelf.Abs("models/gpt2/special_tokens_map.json")); err != nil { + log.Fatal(err) + } else { + special = string(bs) + } + + config = strings.Replace(config, "50256", fmt.Sprintf("%d", step), -1) + + if err := os.WriteFile(filepath.Join(out, "tokenizer_config.json"), []byte(config), os.ModePerm); err != nil { + log.Fatal(err) + } + + if err := os.WriteFile(filepath.Join(out, "special_tokens_map.json"), []byte(special), os.ModePerm); err != nil { + log.Fatal(err) + } + } + } +} + +func subDirs(base string) ([]string, error) { + paths := make([]string, 0) + + err := filepath.WalkDir(base, func(path string, d os.DirEntry, err error) error { + if err != nil { + return err + } + + rel, err := filepath.Rel(base, path) + + if err != nil { + return err + } + + depth := strings.Count(rel, string(os.PathSeparator)) + + if d.IsDir() && rel != "." && depth == 0 { + paths = append(paths, path) + } + + return nil + }) + + return paths, err +} diff --git a/research/knobloch/go.mod b/research/knobloch/go.mod new file mode 100644 index 0000000..e1a110b --- /dev/null +++ b/research/knobloch/go.mod @@ -0,0 +1,3 @@ +module go.jknobloc.com/x/research/knobloch + +go 1.25 |
