From 3fe37224ccc2500e7e767e5fc00163ef61182828 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Fri, 17 Apr 2026 16:07:14 +0200 Subject: Add sander module --- research/sander/extract.go | 52 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 52 insertions(+) create mode 100644 research/sander/extract.go (limited to 'research/sander/extract.go') diff --git a/research/sander/extract.go b/research/sander/extract.go new file mode 100644 index 0000000..de31b1f --- /dev/null +++ b/research/sander/extract.go @@ -0,0 +1,52 @@ +package sander + +import ( + "database/sql" + "database/sql/driver" + + "go.jknobloc.com/x/onnx" + "go.jknobloc.com/x/research/lesci" + "go.jknobloc.com/x/tensor" +) + +func (e *Experiment) Extract(db *sql.DB) error { + var data []float32 + var shape []int + + if d, s, err := onnx.ExtractInitializer(e.model, "transformer.wte.weight"); err != nil { + return err + } else { + data = d + shape = s + } + + if shape[1] != 768 { + panic("unimplemented") + } + + t := tensor.NewDense[float32](shape, data) + + if _, err := db.Exec(` + DROP TABLE IF EXISTS embeddings; + CREATE TABLE embeddings (token_id INTEGER, embedding FLOAT[768], reference BOOLEAN); + `); err != nil { + return err + } + + return lesci.AppendRows(db, "embeddings", func(appendFunc lesci.AppendFunc) error { + for i := range shape[0] { + embedding, ok := t.Select(0, i).Contiguous().Data() + + if !ok { + panic("") + } + + _, reference := e.reference[int64(i)] + + if err := appendFunc([]driver.Value{int32(i), embedding, reference}); err != nil { + return err + } + } + return nil + }) +} -- cgit v1.3.1