summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--research/knobloch/cmd/train/main.go48
1 files changed, 19 insertions, 29 deletions
diff --git a/research/knobloch/cmd/train/main.go b/research/knobloch/cmd/train/main.go
index a2dbdc2..5436784 100644
--- a/research/knobloch/cmd/train/main.go
+++ b/research/knobloch/cmd/train/main.go
@@ -42,23 +42,11 @@ func train() {
}
return m
- }
+ }()
// mbpe.InvertWeightFunction = true
// mbpe.UseSimpleClashes = true
- m000 := morfessor()
- m010 := morfessor()
- m020 := morfessor()
- m030 := morfessor()
- m040 := morfessor()
- m050 := morfessor()
- m060 := morfessor()
- m070 := morfessor()
- m080 := morfessor()
- m090 := morfessor()
- m100 := morfessor()
-
newTrainer := func(segmenter mbpe.Segmenter, alpha float64) *mbpe.MBPETrainer {
alphabet := make(map[string]struct{})
@@ -73,24 +61,26 @@ func train() {
return mbpe.NewMBPETrainer(b, segmenter, alpha, mbpe.NewMBPE(), 1<<17, alphabet)
}
- trainers := []struct {
- *mbpe.MBPETrainer
+ configs := []struct {
+ float64
string
}{
- {newTrainer(m000, 0.0), "m000_minipile_v2"},
- {newTrainer(m010, 0.1), "m010_minipile_v2"},
- {newTrainer(m020, 0.2), "m020_minipile_v2"},
- {newTrainer(m030, 0.3), "m030_minipile_v2"},
- {newTrainer(m040, 0.4), "m040_minipile_v2"},
- {newTrainer(m050, 0.5), "m050_minipile_v2"},
- {newTrainer(m060, 0.6), "m060_minipile_v2"},
- {newTrainer(m070, 0.7), "m070_minipile_v2"},
- {newTrainer(m080, 0.8), "m080_minipile_v2"},
- {newTrainer(m090, 0.9), "m090_minipile_v2"},
- {newTrainer(m100, 1.0), "m100_minipile_v2"},
+ {0.0, "m000_minipile_v2"},
+ {0.1, "m010_minipile_v2"},
+ {0.2, "m020_minipile_v2"},
+ {0.3, "m030_minipile_v2"},
+ {0.4, "m040_minipile_v2"},
+ {0.5, "m050_minipile_v2"},
+ {0.6, "m060_minipile_v2"},
+ {0.7, "m070_minipile_v2"},
+ {0.8, "m080_minipile_v2"},
+ {0.9, "m090_minipile_v2"},
+ {1.0, "m100_minipile_v2"},
}
- for i, t := range trainers {
+ for i, c := range configs {
+ t := newTrainer(morfessor, c.float64)
+
dict := filepath.Join(out, "dict.txt")
if dictErr := t.LoadDict(dict); dictErr != nil {
@@ -115,7 +105,7 @@ func train() {
fmt.Println()
}
- fmt.Printf("%s\n\n", t.string)
+ fmt.Printf("%s\n\n", c.string)
if fileOpen, errOpen := os.Open(filepath.Join(out, "segments.gob")); errOpen != nil {
if fileCreate, errCreate := os.Create(filepath.Join(out, "segments.gob")); errCreate != nil {
@@ -137,7 +127,7 @@ func train() {
t.Train()
- dir := filepath.Join(out, t.string)
+ dir := filepath.Join(out, c.string)
if err := os.Mkdir(dir, 0755); err != nil && !errors.Is(err, os.ErrExist) {
log.Fatal(err)