diff options
Diffstat (limited to 'research/knobloch/share.go')
| -rw-r--r-- | research/knobloch/share.go | 223 |
1 files changed, 223 insertions, 0 deletions
diff --git a/research/knobloch/share.go b/research/knobloch/share.go new file mode 100644 index 0000000..bc715f1 --- /dev/null +++ b/research/knobloch/share.go @@ -0,0 +1,223 @@ +package knobloch + +// Self-contained plot over the accumulated frequency_stats.csv; nothing else in +// the package depends on it. + +import ( + "encoding/csv" + "fmt" + "image/color" + "os" + "regexp" + "slices" + "strconv" + + "gonum.org/v1/plot" + "gonum.org/v1/plot/plotter" + "gonum.org/v1/plot/vg" + "gonum.org/v1/plot/vg/draw" +) + +// matches the alpha encoded in names like gpt2_50256_mi050_minipile +var modelPattern = regexp.MustCompile(`_(mi?)(\d{3})_`) + +type shareRow struct { + variant string + alpha float64 + types float64 + tokens float64 +} + +func PlotMorphemeShare(name, out string) error { + rows, err := readShareRows(name) + + if err != nil { + return err + } + + if len(rows) == 0 { + return fmt.Errorf("no usable rows in %s", name) + } + + p := plot.New() + + p.X.Label.Text = "Morpheme weight α (%)" + p.Y.Label.Text = "Share of vocabulary / token mass (%)" + + p.Add(plotter.NewGrid()) + + variants := make([]string, 0, 2) + + for _, r := range rows { + if !slices.Contains(variants, r.variant) { + variants = append(variants, r.variant) + } + } + + slices.Sort(variants) + + blue := color.NRGBA{R: 108, G: 126, B: 179, A: 255} + salmon := color.NRGBA{R: 214, G: 96, B: 77, A: 255} + + lo, hi := 100.0, 0.0 + + for _, variant := range variants { + var types, tokens plotter.XYs + + for _, r := range rows { + if r.variant != variant { + continue + } + + types = append(types, plotter.XY{X: r.alpha, Y: r.types * 100}) + tokens = append(tokens, plotter.XY{X: r.alpha, Y: r.tokens * 100}) + + lo = min(lo, r.types*100, r.tokens*100) + hi = max(hi, r.types*100, r.tokens*100) + } + + // a dashed line separates the mi variant from m at the same alpha + var dashes []vg.Length + + if variant != "m" { + dashes = []vg.Length{vg.Points(4), vg.Points(3)} + } + + if err := addShareSeries(p, types, blue, dashes, variant+" types"); err != nil { + return err + } + + if err := addShareSeries(p, tokens, salmon, dashes, variant+" tokens"); err != nil { + return err + } + } + + // keep the series off the frame and away from the legend + pad := max((hi-lo)*0.15, 2) + + p.Y.Min = lo - pad + p.Y.Max = hi + pad + + p.X.Min = -5 + p.X.Max = 105 + + p.Legend.Padding = vg.Points(4) + + return p.Save(7*vg.Inch, 5*vg.Inch, out) +} + +func addShareSeries(p *plot.Plot, pts plotter.XYs, c color.NRGBA, dashes []vg.Length, label string) error { + slices.SortFunc(pts, func(a, b plotter.XY) int { + switch { + case a.X < b.X: + return -1 + case a.X > b.X: + return 1 + default: + return 0 + } + }) + + line, points, err := plotter.NewLinePoints(pts) + + if err != nil { + return err + } + + line.Color = c + line.Width = vg.Points(1.5) + line.Dashes = dashes + + points.Color = c + points.Radius = vg.Points(3) + points.Shape = draw.CircleGlyph{} + + p.Add(line, points) + p.Legend.Add(label, line, points) + + return nil +} + +func readShareRows(name string) ([]shareRow, error) { + file, err := os.Open(name) + + if err != nil { + return nil, err + } + + defer file.Close() + + records, err := csv.NewReader(file).ReadAll() + + if err != nil { + return nil, err + } + + if len(records) == 0 { + return nil, fmt.Errorf("%s is empty", name) + } + + index := make(map[string]int) + + for i, h := range records[0] { + index[h] = i + } + + for _, h := range []string{"model", "type_share", "token_share"} { + if _, ok := index[h]; !ok { + return nil, fmt.Errorf("missing column %q in %s", h, name) + } + } + + // the stats file is appended to, so a rerun can repeat a model + seen := make(map[string]shareRow) + + var order []string + + for _, record := range records[1:] { + model := record[index["model"]] + + m := modelPattern.FindStringSubmatch(model) + + if m == nil { + continue + } + + alpha, err := strconv.ParseFloat(m[2], 64) + + if err != nil { + return nil, err + } + + types, err := strconv.ParseFloat(record[index["type_share"]], 64) + + if err != nil { + return nil, err + } + + tokens, err := strconv.ParseFloat(record[index["token_share"]], 64) + + if err != nil { + return nil, err + } + + if _, ok := seen[model]; !ok { + order = append(order, model) + } + + seen[model] = shareRow{ + variant: m[1], + alpha: alpha, + types: types, + tokens: tokens, + } + } + + rows := make([]shareRow, 0, len(order)) + + for _, model := range order { + rows = append(rows, seen[model]) + } + + return rows, nil +} |
