summaryrefslogtreecommitdiff
path: root/research/knobloch/cmd/train/main.go
blob: a946ce49d5a28ae1cd4dfa1cb4bb0936a402e484 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
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)
		}
	}
}