diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-17 16:07:14 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-17 16:07:14 +0200 |
| commit | 3fe37224ccc2500e7e767e5fc00163ef61182828 (patch) | |
| tree | 7bd14e080c080b24026660beda42817a3c9dc77f /research/sander/plot.go | |
| parent | 4d8f647bfb54899e3052931790e288da7af97050 (diff) | |
Add sander module
Diffstat (limited to 'research/sander/plot.go')
| -rw-r--r-- | research/sander/plot.go | 100 |
1 files changed, 100 insertions, 0 deletions
diff --git a/research/sander/plot.go b/research/sander/plot.go new file mode 100644 index 0000000..bc6c483 --- /dev/null +++ b/research/sander/plot.go @@ -0,0 +1,100 @@ +package sander + +import ( + "context" + "database/sql" + "image/color" + "path" + + "gonum.org/v1/plot" + "gonum.org/v1/plot/plotter" + "gonum.org/v1/plot/vg" + "gonum.org/v1/plot/vg/draw" +) + +func (e *Experiment) Plot(db *sql.DB) error { + var rows *sql.Rows + + if r, err := db.QueryContext(context.Background(), ` + SELECT s.token_id, s.distance, e.reference + FROM similarity_buffer s + JOIN embeddings e ON s.token_id = e.token_id + ORDER BY s.token_id ASC + `); err != nil { + return err + } else { + rows = r + } + + defer rows.Close() + + var pts plotter.XYZs + + for rows.Next() { + var id, distance float64 + var reference bool + + if err := rows.Scan(&id, &distance, &reference); err != nil { + return err + } + + z := 0.0 + + if reference { + z = 1.0 + } + + pts = append(pts, plotter.XYZ{X: id, Y: distance, Z: z}) + } + + if err := rows.Err(); err != nil { + return err + } + + return render(pts, path.Join(e.name, "similarity_scatter.png")) +} + +func render(pts plotter.XYZs, out string) error { + p := plot.New() + + p.Y.Label.Text = "Euclidean Distance" + + p.Y.Min = 0 + p.Y.Max = 5 + + var candidate, reference plotter.XYZs + + for _, pt := range pts { + if pt.Z == 1.0 { + reference = append(reference, pt) + } else { + candidate = append(candidate, pt) + } + } + + if err := addScatter(p, candidate, color.RGBA{A: 64}); err != nil { + return err + } + + if err := addScatter(p, reference, color.RGBA{R: 255, G: 69, B: 0, A: 255}); err != nil { + return err + } + + return p.Save(10*vg.Inch, 6*vg.Inch, out) +} + +func addScatter(p *plot.Plot, pts plotter.XYZs, c color.Color) error { + s, err := plotter.NewScatter(pts) + + if err != nil { + return err + } + + s.Color = c + s.Radius = vg.Points(2) + s.Shape = draw.CircleGlyph{} + + p.Add(s) + + return nil +} |
