summaryrefslogtreecommitdiff
path: root/research/knobloch/share.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/knobloch/share.go')
-rw-r--r--research/knobloch/share.go223
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
+}