summaryrefslogtreecommitdiff
path: root/research/sander/plot.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-17 16:07:14 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-17 16:07:14 +0200
commit3fe37224ccc2500e7e767e5fc00163ef61182828 (patch)
tree7bd14e080c080b24026660beda42817a3c9dc77f /research/sander/plot.go
parent4d8f647bfb54899e3052931790e288da7af97050 (diff)
Add sander module
Diffstat (limited to 'research/sander/plot.go')
-rw-r--r--research/sander/plot.go100
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
+}