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 }