summaryrefslogtreecommitdiff
path: root/research/sander/extract.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/extract.go
parent4d8f647bfb54899e3052931790e288da7af97050 (diff)
Add sander module
Diffstat (limited to 'research/sander/extract.go')
-rw-r--r--research/sander/extract.go52
1 files changed, 52 insertions, 0 deletions
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
+ })
+}