summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2025-11-13 23:04:57 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2025-11-17 22:41:09 +0100
commit7073b124c5eb31169442ef18d1627b7ed7401280 (patch)
treef472b8a937fb5e2e3f235723cd5bafe422b3b647
parent019a25a6082a8cde372304e34dcd0ae3d5e875ed (diff)
Scaffold byte-pair correction
-rw-r--r--bpc/bpc.go28
-rw-r--r--bpc/cmd/bpc/main.go45
-rw-r--r--bpc/go.mod23
-rw-r--r--bpc/go.sum69
-rw-r--r--go.work6
-rw-r--r--go.work.sum47
-rw-r--r--gpt2/.gitignore1
-rw-r--r--gpt2/cmd/gpt2/main.go18
-rw-r--r--gpt2/model.go (renamed from gpt2/main.go)46
-rw-r--r--gpt2/scripts/conv.py2
-rw-r--r--llm/causal.go5
-rw-r--r--llm/go.mod3
-rw-r--r--llm/tokenizer.go5
13 files changed, 280 insertions, 18 deletions
diff --git a/bpc/bpc.go b/bpc/bpc.go
new file mode 100644
index 0000000..67804d3
--- /dev/null
+++ b/bpc/bpc.go
@@ -0,0 +1,28 @@
+package bpc
+
+import (
+ "fmt"
+ "llm"
+)
+
+func Run(model llm.Causal, tokenizer llm.Tokenizer) {
+ t := tokenizer.Tokenize("The quick brown")
+
+ logits := make([][]float32, 0)
+
+ if _, err := model.Generate(toInt64(t), 0, &logits); err != nil {
+ // TODO handle
+ }
+
+ fmt.Println(logits)
+}
+
+func toInt64(s []int) []int64 {
+ r := make([]int64, len(s))
+
+ for i, v := range s {
+ r[i] = int64(v)
+ }
+
+ return r
+}
diff --git a/bpc/cmd/bpc/main.go b/bpc/cmd/bpc/main.go
new file mode 100644
index 0000000..b6ebef1
--- /dev/null
+++ b/bpc/cmd/bpc/main.go
@@ -0,0 +1,45 @@
+package main
+
+import (
+ "bpc"
+ "gpt2"
+ "log"
+ mbpe "mbpe-dyn"
+)
+
+func main() {
+ m := model()
+
+ if err := m.Init(); err != nil {
+ log.Fatal(err)
+ }
+
+ t := tokenizer()
+
+ bpc.Run(m, t)
+
+ if err := m.Destroy(); err != nil {
+ log.Fatal(err)
+ }
+}
+
+func model() *gpt2.Model {
+ return gpt2.NewModel("../gpt2/models/base/model.onnx")
+}
+
+func tokenizer() *mbpe.Tokenizer {
+ m := mbpe.NewMBPE()
+
+ if err := m.Load("../gpt2/models/base/vocab.json", "../gpt2/models/base/merges.txt"); err != nil {
+ panic(err)
+ }
+
+ t := mbpe.NewTokenizer(m)
+
+ byteLevel := mbpe.NewByteLevel(false)
+
+ t.SetPreTokenizer(byteLevel)
+ t.SetDecoder(byteLevel)
+
+ return t
+}
diff --git a/bpc/go.mod b/bpc/go.mod
new file mode 100644
index 0000000..42c95a3
--- /dev/null
+++ b/bpc/go.mod
@@ -0,0 +1,23 @@
+module bpc
+
+go 1.24.7
+
+replace mbpe-dyn => github.com/jonasknobloch/mbpe-dyn v0.0.0-20251113214706-ba5a18b759bd
+
+require mbpe-dyn v0.0.0-00010101000000-000000000000
+
+require (
+ git.sr.ht/~sbinet/gg v0.6.0 // indirect
+ github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b // indirect
+ github.com/campoy/embedmd v1.0.0 // indirect
+ github.com/go-fonts/liberation v0.3.3 // indirect
+ github.com/go-latex/latex v0.0.0-20240709081214-31cef3c7570e // indirect
+ github.com/go-pdf/fpdf v0.9.0 // indirect
+ github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 // indirect
+ github.com/pmezard/go-difflib v1.0.0 // indirect
+ golang.org/x/image v0.24.0 // indirect
+ golang.org/x/text v0.22.0 // indirect
+ gonum.org/v1/gonum v0.15.1 // indirect
+ gonum.org/v1/plot v0.15.0 // indirect
+ google.golang.org/protobuf v1.36.5 // indirect
+)
diff --git a/bpc/go.sum b/bpc/go.sum
new file mode 100644
index 0000000..a7e8e62
--- /dev/null
+++ b/bpc/go.sum
@@ -0,0 +1,69 @@
+git.sr.ht/~sbinet/cmpimg v0.1.0 h1:E0zPRk2muWuCqSKSVZIWsgtU9pjsw3eKHi8VmQeScxo=
+git.sr.ht/~sbinet/cmpimg v0.1.0/go.mod h1:FU12psLbF4TfNXkKH2ZZQ29crIqoiqTZmeQ7dkp/pxE=
+git.sr.ht/~sbinet/gg v0.6.0 h1:RIzgkizAk+9r7uPzf/VfbJHBMKUr0F5hRFxTUGMnt38=
+git.sr.ht/~sbinet/gg v0.6.0/go.mod h1:uucygbfC9wVPQIfrmwM2et0imr8L7KQWywX0xpFMm94=
+github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
+github.com/ajstarks/deck v0.0.0-20200831202436-30c9fc6549a9/go.mod h1:JynElWSGnm/4RlzPXRlREEwqTHAN3T56Bv2ITsFT3gY=
+github.com/ajstarks/deck/generate v0.0.0-20210309230005-c3f852c02e19/go.mod h1:T13YZdzov6OU0A1+RfKZiZN9ca6VeKdBdyDV+BY97Tk=
+github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b h1:slYM766cy2nI3BwyRiyQj/Ud48djTMtMebDqepE95rw=
+github.com/ajstarks/svgo v0.0.0-20211024235047-1546f124cd8b/go.mod h1:1KcenG0jGWcpt8ov532z81sp/kMMUG485J2InIOyADM=
+github.com/campoy/embedmd v1.0.0 h1:V4kI2qTJJLf4J29RzI/MAt2c3Bl4dQSYPuflzwFH2hY=
+github.com/campoy/embedmd v1.0.0/go.mod h1:oxyr9RCiSXg0M3VJ3ks0UGfp98BpSSGr0kpiX3MzVl8=
+github.com/dlclark/regexp2 v1.11.5 h1:Q/sSnsKerHeCkc/jSTNq1oCm7KiVgUMZRDUoRu0JQZQ=
+github.com/dlclark/regexp2 v1.11.5/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
+github.com/go-fonts/dejavu v0.3.4 h1:Qqyx9IOs5CQFxyWTdvddeWzrX0VNwUAvbmAzL0fpjbc=
+github.com/go-fonts/dejavu v0.3.4/go.mod h1:D1z0DglIz+lmpeNYMYlxW4r22IhcdOYnt+R3PShU/Kg=
+github.com/go-fonts/latin-modern v0.3.3 h1:g2xNgI8yzdNzIVm+qvbMryB6yGPe0pSMss8QT3QwlJ0=
+github.com/go-fonts/latin-modern v0.3.3/go.mod h1:tHaiWDGze4EPB0Go4cLT5M3QzRY3peya09Z/8KSCrpY=
+github.com/go-fonts/liberation v0.3.3 h1:tM/T2vEOhjia6v5krQu8SDDegfH1SfXVRUNNKpq0Usk=
+github.com/go-fonts/liberation v0.3.3/go.mod h1:eUAzNRuJnpSnd1sm2EyloQfSOT79pdw7X7++Ri+3MCU=
+github.com/go-latex/latex v0.0.0-20240709081214-31cef3c7570e h1:xcdj0LWnMSIU1j8+jIeJyfvk6SjgJedFQssSqFthJ2E=
+github.com/go-latex/latex v0.0.0-20240709081214-31cef3c7570e/go.mod h1:J4SAGzkcl+28QWi7yz72tyC/4aGnppOvya+AEv4TaAQ=
+github.com/go-pdf/fpdf v0.9.0 h1:PPvSaUuo1iMi9KkaAn90NuKi+P4gwMedWPHhj8YlJQw=
+github.com/go-pdf/fpdf v0.9.0/go.mod h1:oO8N111TkmKb9D7VvWGLvLJlaZUQVPM+6V42pp3iV4Y=
+github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0 h1:DACJavvAHhabrF08vX0COfcOBJRhZ8lUbR+ZWIs0Y5g=
+github.com/golang/freetype v0.0.0-20170609003504-e2365dfdc4a0/go.mod h1:E/TSTwGwJL78qG/PmXZO1EjYhfJinVAhrmmHX6Z8B9k=
+github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
+github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
+github.com/jonasknobloch/mbpe-dyn v0.0.0-20251113214706-ba5a18b759bd h1:tH5hi+l4vT7moNvVm0UflhIvNErZdXH8221jTLFafak=
+github.com/jonasknobloch/mbpe-dyn v0.0.0-20251113214706-ba5a18b759bd/go.mod h1:9ygffEKXqjku7v4qOBdw9FXRBCzn8SmOPG4nLJnXUSY=
+github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
+github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
+github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
+github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74=
+golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
+golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
+golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
+golang.org/x/exp v0.0.0-20241009180824-f66d83c29e7c h1:7dEasQXItcW1xKJ2+gg5VOiBnqWrJc+rq0DPKyvvdbY=
+golang.org/x/exp v0.0.0-20241009180824-f66d83c29e7c/go.mod h1:NQtJDoLvd6faHhE7m4T/1IY708gDefGGjR/iUW8yQQ8=
+golang.org/x/image v0.24.0 h1:AN7zRgVsbvmTfNyqIbbOraYL8mSwcKncEj8ofjgzcMQ=
+golang.org/x/image v0.24.0/go.mod h1:4b/ITuLfqYq1hqZcjofwctIhi7sZh2WaCjvsBNjjya8=
+golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
+golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
+golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
+golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
+golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
+golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
+golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
+golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
+golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
+golang.org/x/sys v0.0.0-20210119212857-b64e53b001e4/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
+golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
+golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
+golang.org/x/text v0.22.0 h1:bofq7m3/HAFvbF51jz3Q9wLg3jkvSPuiZu/pD1XwgtM=
+golang.org/x/text v0.22.0/go.mod h1:YRoo4H8PVmsu+E3Ou7cqLVH8oXWIHVoX0jqUWALQhfY=
+golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
+golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
+golang.org/x/tools v0.1.0/go.mod h1:xkSsbof2nBLbhDlRMhhhyNLN/zl3eTqcnHD5viDpcZ0=
+golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
+golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
+golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
+gonum.org/v1/gonum v0.15.1 h1:FNy7N6OUZVUaWG9pTiD+jlhdQ3lMP+/LcTpJ6+a8sQ0=
+gonum.org/v1/gonum v0.15.1/go.mod h1:eZTZuRFrzu5pcyjN5wJhcIhnUdNijYxX1T2IcrOGY0o=
+gonum.org/v1/plot v0.15.0 h1:SIFtFNdZNWLRDRVjD6CYxdawcpJDWySZehJGpv1ukkw=
+gonum.org/v1/plot v0.15.0/go.mod h1:3Nx4m77J4T/ayr/b8dQ8uGRmZF6H3eTqliUExDrQHnM=
+google.golang.org/protobuf v1.36.5 h1:tPhr+woSbjfYvY6/GPufUoYizxw1cF/yFoxJ2fmpwlM=
+google.golang.org/protobuf v1.36.5/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
+honnef.co/go/tools v0.1.3/go.mod h1:NgwopIslSNH47DimFoV78dnkksY2EFtX0ajyb3K/las=
+rsc.io/pdf v0.1.1 h1:k1MczvYDUvJBe93bYd7wrZLLUEcLZAuF824/I4e5Xr4=
+rsc.io/pdf v0.1.1/go.mod h1:n8OzWcQ6Sp37PL01nO98y4iUCRdTGarVfzxY20ICaU4=
diff --git a/go.work b/go.work
index bddb430..0891173 100644
--- a/go.work
+++ b/go.work
@@ -1,3 +1,7 @@
go 1.24.7
-use ./gpt2
+use (
+ ./bpc
+ ./gpt2
+ ./llm
+)
diff --git a/go.work.sum b/go.work.sum
new file mode 100644
index 0000000..faef89e
--- /dev/null
+++ b/go.work.sum
@@ -0,0 +1,47 @@
+gioui.org v0.2.0 h1:RbzDn1h/pCVf/q44ImQSa/J3MIFpY3OWphzT/Tyei+w=
+gioui.org v0.2.0/go.mod h1:1H72sKEk/fNFV+l0JNeM2Dt3co3Y4uaQcD+I+/GQ0e4=
+gioui.org/cpu v0.0.0-20220412190645-f1e9e8c3b1f7 h1:tNJdnP5CgM39PRc+KWmBRRYX/zJ+rd5XaYxY5d5veqA=
+gioui.org/cpu v0.0.0-20220412190645-f1e9e8c3b1f7/go.mod h1:A8M0Cn5o+vY5LTMlnRoK3O5kG+rH0kWfJjeKd9QpBmQ=
+gioui.org/shader v1.0.6 h1:cvZmU+eODFR2545X+/8XucgZdTtEjR3QWW6W65b0q5Y=
+gioui.org/shader v1.0.6/go.mod h1:mWdiME581d/kV7/iEhLmUgUK5iZ09XR5XpduXzbePVM=
+gioui.org/x v0.2.0 h1:/MbdjKH19F16auv19UiQxli2n6BYPw7eyh9XBOTgmEw=
+gioui.org/x v0.2.0/go.mod h1:rCGN2nZ8ZHqrtseJoQxCMZpt2xrZUrdZ2WuMRLBJmYs=
+github.com/BurntSushi/toml v0.3.1 h1:WXkYYl6Yr3qBf1K79EBnL4mak0OimBfB0XUf9Vl28OQ=
+github.com/ajstarks/deck v0.0.0-20200831202436-30c9fc6549a9 h1:7kQgkwGRoLzC9K0oyXdJo7nve/bynv/KwUsxbiTlzAM=
+github.com/ajstarks/deck/generate v0.0.0-20210309230005-c3f852c02e19 h1:iXUgAaqDcIUGbRoy2TdeofRG/j1zpGRSEmNK05T+bi8=
+github.com/andybalholm/stroke v0.0.0-20221221101821-bd29b49d73f0 h1:uF5Q/hWnDU1XZeT6CsrRSxHLroUSEYYO3kgES+yd+So=
+github.com/andybalholm/stroke v0.0.0-20221221101821-bd29b49d73f0/go.mod h1:ccdDYaY5+gO+cbnQdFxEXqfy0RkoV25H3jLXUDNM3wg=
+github.com/boombuler/barcode v1.0.1 h1:NDBbPmhS+EqABEs5Kg3n/5ZNjy73Pz7SIV+KCeqyXcs=
+github.com/boombuler/barcode v1.0.1/go.mod h1:paBWMcWSl3LHKBqUq+rly7CNSldXjb2rDl3JlRe0mD8=
+github.com/go-fonts/stix v0.2.2 h1:v9krocr13J1llaOHLEol1eaHsv8S43UuFX/1bFgEJJ4=
+github.com/go-fonts/stix v0.2.2/go.mod h1:SUxggC9dxd/Q+rb5PkJuvfvTbOPtNc2Qaua00fIp9iU=
+github.com/go-text/typesetting v0.0.0-20230803102845-24e03d8b5372 h1:FQivqchis6bE2/9uF70M2gmmLpe82esEm2QadL0TEJo=
+github.com/go-text/typesetting v0.0.0-20230803102845-24e03d8b5372/go.mod h1:evDBbvNR/KaVFZ2ZlDSOWWXIUKq0wCOEtzLxRM8SG3k=
+github.com/goccmack/gocc v0.0.0-20230228185258-2292f9e40198 h1:FSii2UQeSLngl3jFoR4tUKZLprO7qUlh/TKKticc0BM=
+github.com/goccmack/gocc v0.0.0-20230228185258-2292f9e40198/go.mod h1:DTh/Y2+NbnOVVoypCCQrovMPDKUGp4yZpSbWg5D0XIM=
+github.com/golang/protobuf v1.5.0 h1:LUVKkCeviFUMKqHa4tXIIij/lbhnMbP7Fn5wKdKkRh4=
+github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk=
+github.com/jonasknobloch/mbpe-dyn v0.0.0-20250825202016-24984e94f5fc h1:PowV6bLABtPgYl6NTa9e+AJYuL5bebLqkurp8dqO1LE=
+github.com/jonasknobloch/mbpe-dyn v0.0.0-20250825202016-24984e94f5fc/go.mod h1:9ygffEKXqjku7v4qOBdw9FXRBCzn8SmOPG4nLJnXUSY=
+github.com/kisielk/gotool v1.0.0 h1:AV2c/EiW3KqPNT9ZKl07ehoAGi4C5/01Cfbblndcapg=
+github.com/phpdave11/gofpdi v1.0.13 h1:o61duiW8M9sMlkVXWlvP92sZJtGKENvW3VExs6dZukQ=
+github.com/phpdave11/gofpdi v1.0.13/go.mod h1:vBmVV0Do6hSBHC8uKUQ71JGW+ZGQq74llk/7bXwjDoI=
+github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
+github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
+github.com/ruudk/golang-pdf417 v0.0.0-20201230142125-a7e3863a1245 h1:K1Xf3bKttbF+koVGaX5xngRIZ5bVjbmPnaxE/dR08uY=
+github.com/ruudk/golang-pdf417 v0.0.0-20201230142125-a7e3863a1245/go.mod h1:pQAZKsJ8yyVxGRWYNEm9oFB8ieLgKFnamEyDmSA0BRk=
+github.com/yuin/goldmark v1.2.1 h1:ruQGxdhGHe7FWOJPT0mKs5+pD2Xs1Bm/kdGlHO04FmM=
+golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9 h1:psW17arqaxU48Z5kZ0CQnkZWQJsqcURM6tKiBApRjXI=
+golang.org/x/exp/shiny v0.0.0-20241009180824-f66d83c29e7c h1:jTMrjjZRcSH/BDxWhXCP6OWsfVgmnwI7J+F4/nyVXaU=
+golang.org/x/exp/shiny v0.0.0-20241009180824-f66d83c29e7c/go.mod h1:3F+MieQB7dRYLTmnncoFbb1crS5lfQoTfDgQy6K4N0o=
+golang.org/x/mod v0.17.0 h1:zY54UmvipHiNd+pm+m0x9KhZ9hl1/7QNMyxXbc6ICqA=
+golang.org/x/mod v0.17.0/go.mod h1:hTbmBsO62+eylJbnUtE2MGJUyE7QWk4xUqPFrRgJ+7c=
+golang.org/x/net v0.0.0-20201021035429-f5854403a974 h1:IX6qOQeG5uLjB/hjjwjedwfjND0hgjPMMyO1RoIXQNI=
+golang.org/x/sync v0.11.0 h1:GGz8+XQP4FvTTrjZPzNKTMFtSXH80RAzG+5ghFPgK9w=
+golang.org/x/sync v0.11.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
+golang.org/x/sys v0.26.0 h1:KHjCJyddX0LoSTb3J+vWpupP9p0oznkqVk/IfjymZbo=
+golang.org/x/sys v0.26.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d h1:vU5i/LfpvrRCpgM/VPfJLg5KjxD3E+hfT1SH+d9zLwg=
+golang.org/x/tools v0.21.1-0.20240508182429-e35e4ccd0d2d/go.mod h1:aiJjzUbINMkxbQROHiO6hDPo2LHcIPhhQsa9DLh0yGk=
+golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 h1:go1bK/D/BFZV2I8cIQd1NKEZ+0owSTG1fDTci4IqFcE=
+honnef.co/go/tools v0.1.3 h1:qTakTkI6ni6LFD5sBwwsdSO+AQqbSIxOauHTTQKZ/7o=
diff --git a/gpt2/.gitignore b/gpt2/.gitignore
new file mode 100644
index 0000000..8c6790b
--- /dev/null
+++ b/gpt2/.gitignore
@@ -0,0 +1 @@
+/models \ No newline at end of file
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index ff84f42..96dabeb 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -4,24 +4,24 @@ import (
"fmt"
"gpt2"
"log"
-
- ort "github.com/yalue/onnxruntime_go"
)
func main() {
- ort.SetSharedLibraryPath("lib/onnxruntime-osx-arm64-1.22.0/lib/libonnxruntime.1.22.0.dylib")
+ prompt := []int64{464, 2068, 7586, 21831}
+
+ m := gpt2.NewModel("models/base/model.onnx")
- if err := ort.InitializeEnvironment(); err != nil {
+ if err := m.Init(); err != nil {
log.Fatal(err)
}
- defer ort.DestroyEnvironment()
-
- prompt := []int64{464, 2068, 7586, 21831}
-
- if out, err := gpt2.Generate("scripts/onnx-gpt2/model.onnx", prompt, 5, nil); err != nil {
+ if out, err := m.Generate(prompt, 5, nil); err != nil {
log.Fatal(err)
} else {
fmt.Printf("\n%v\n", out)
}
+
+ if err := m.Destroy(); err != nil {
+ log.Fatal(err)
+ }
}
diff --git a/gpt2/main.go b/gpt2/model.go
index ee9e19f..1b7a9d1 100644
--- a/gpt2/main.go
+++ b/gpt2/model.go
@@ -3,8 +3,10 @@ package gpt2
import (
"errors"
"fmt"
+ _ "llm"
"log"
"math"
+ "os"
"sort"
ort "github.com/yalue/onnxruntime_go"
@@ -17,7 +19,37 @@ const (
headDim = 64
)
-func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
+type Model struct {
+ name string
+}
+
+func NewModel(name string) *Model {
+ return &Model{
+ name: name,
+ }
+}
+
+func (m *Model) SharedLibraryPath() string {
+ p, ok := os.LookupEnv("ONNXRUNTIME_SHARED_LIBRARY_PATH")
+
+ if !ok {
+ // TODO embed runtime binaries
+ }
+
+ return p
+}
+
+func (m *Model) Init() error {
+ ort.SetSharedLibraryPath(m.SharedLibraryPath())
+
+ return ort.InitializeEnvironment()
+}
+
+func (m *Model) Destroy() error {
+ return ort.DestroyEnvironment()
+}
+
+func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
if len(prompt) == 0 {
return nil, errors.New("empty prompt")
}
@@ -31,7 +63,7 @@ func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([
out := make([]int64, 0, steps+1)
for step := range context + steps {
- _, _, outputs, err := forward(model, token, step, cacheNames, cacheValues)
+ _, _, outputs, err := forward(m.name, token, step, cacheNames, cacheValues)
if err != nil {
return nil, err
@@ -43,13 +75,13 @@ func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([
*logits = append(*logits, l)
}
- idx, p := topK(softmax(l), 5)
+ idx, _ := topK(softmax(l), 5)
- fmt.Printf("\n%d\n\n", token)
+ // fmt.Printf("\n%d\n\n", token)
- for i, t := range idx {
- fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
- }
+ // for i, t := range idx {
+ // fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t)
+ // }
if step < context-1 {
token = prompt[step+1]
diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py
index cf0cbb8..e3773e2 100644
--- a/gpt2/scripts/conv.py
+++ b/gpt2/scripts/conv.py
@@ -12,4 +12,4 @@ from optimum.onnxruntime import ORTModelForCausalLM
model_id = "gpt2"
model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True)
-model.save_pretrained("onnx-gpt2") \ No newline at end of file
+model.save_pretrained("../models/base") \ No newline at end of file
diff --git a/llm/causal.go b/llm/causal.go
new file mode 100644
index 0000000..296a744
--- /dev/null
+++ b/llm/causal.go
@@ -0,0 +1,5 @@
+package llm
+
+type Causal interface {
+ Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error)
+}
diff --git a/llm/go.mod b/llm/go.mod
new file mode 100644
index 0000000..9a09bef
--- /dev/null
+++ b/llm/go.mod
@@ -0,0 +1,3 @@
+module llm
+
+go 1.24.7
diff --git a/llm/tokenizer.go b/llm/tokenizer.go
new file mode 100644
index 0000000..c6ea3c1
--- /dev/null
+++ b/llm/tokenizer.go
@@ -0,0 +1,5 @@
+package llm
+
+type Tokenizer interface {
+ Tokenize(s string) []int
+}