summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-30 14:43:28 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-30 14:43:28 +0200
commitf9332fe85bf12433286fc77456ae2705c494a4f9 (patch)
tree5f254e600b3916478302580383fd2b5611af122c
parentcc7aa68fa63cb69cb3a565364ee0bee1f482f6f5 (diff)
Add knobloch module
-rw-r--r--go.work1
m---------mbpe0
-rw-r--r--research/knobloch/cmd/train/main.go124
-rw-r--r--research/knobloch/cmd/train/serialize.go108
-rw-r--r--research/knobloch/go.mod3
5 files changed, 236 insertions, 0 deletions
diff --git a/go.work b/go.work
index b99b737..fd6b5b5 100644
--- a/go.work
+++ b/go.work
@@ -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