diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-13 23:04:57 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-17 22:41:09 +0100 |
| commit | 7073b124c5eb31169442ef18d1627b7ed7401280 (patch) | |
| tree | f472b8a937fb5e2e3f235723cd5bafe422b3b647 | |
| parent | 019a25a6082a8cde372304e34dcd0ae3d5e875ed (diff) | |
Scaffold byte-pair correction
| -rw-r--r-- | bpc/bpc.go | 28 | ||||
| -rw-r--r-- | bpc/cmd/bpc/main.go | 45 | ||||
| -rw-r--r-- | bpc/go.mod | 23 | ||||
| -rw-r--r-- | bpc/go.sum | 69 | ||||
| -rw-r--r-- | go.work | 6 | ||||
| -rw-r--r-- | go.work.sum | 47 | ||||
| -rw-r--r-- | gpt2/.gitignore | 1 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 18 | ||||
| -rw-r--r-- | gpt2/model.go (renamed from gpt2/main.go) | 46 | ||||
| -rw-r--r-- | gpt2/scripts/conv.py | 2 | ||||
| -rw-r--r-- | llm/causal.go | 5 | ||||
| -rw-r--r-- | llm/go.mod | 3 | ||||
| -rw-r--r-- | llm/tokenizer.go | 5 |
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= @@ -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 +} |
