diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-11 19:37:15 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-02-11 19:58:06 +0100 |
| commit | 966213896cae3faf8d6dbb0e07b98d9f370cb423 (patch) | |
| tree | cce4a2a904eab1145a317e68cb9f9e269afab596 /gpt2 | |
| parent | eb17c4545f525ac94700f61146cc341323a35a50 (diff) | |
Add top-K and logit RMSE inference tests
Diffstat (limited to 'gpt2')
| -rw-r--r-- | gpt2/model_test.go | 138 |
1 files changed, 138 insertions, 0 deletions
diff --git a/gpt2/model_test.go b/gpt2/model_test.go new file mode 100644 index 0000000..dcd00e6 --- /dev/null +++ b/gpt2/model_test.go @@ -0,0 +1,138 @@ +package gpt2 + +import ( + "encoding/binary" + "log" + "math" + "os" + "slices" + "testing" +) + +func fromModel() []float32 { + prompt := []int64{464, 2068, 7586, 21831, 18045, 625, 262, 16931, 3290} + + m := model() + + defer m.Destroy() + + logits := make([][]float32, 0) + + if _, err := m.Generate(prompt, 0, &logits); err != nil { + log.Fatal(err) + } + + return flatten(logits) +} + +func model() *Model { + m := NewModel("models/base/model.onnx", "0") // TODO check if CUDA is available + + if err := m.Init(); err != nil { + log.Fatal(err) + } + + return m +} + +func fromGold() []float32 { + shape := []int{1, 9, 50257} + + s, err := f32("test/logits.f32", shape[0]*shape[1]*shape[2]) + + if err != nil { + log.Fatal(err) + } + + return s +} + +func f32(name string, n int) ([]float32, error) { + var file *os.File + + if f, err := os.Open(name); err != nil { + return nil, err + } else { + file = f + } + + data := make([]float32, n) + + if err := binary.Read(file, binary.LittleEndian, data); err != nil { + return nil, err + } + + return data, nil +} + +func TestModel_GenerateTopK(t *testing.T) { + a := fromGold() + b := fromModel() + + shape := []int{1, 9, 50257} + + for i := range shape[0] { + for j := range shape[1] { + start := (i * shape[1] * shape[2]) + (j * shape[2]) + stop := start + shape[2] + + topA, _ := topK(a[start:stop], 64) + topB, _ := topK(b[start:stop], 64) + + if !slices.Equal(topA, topB) { + t.Errorf("token mismatch at token %d batch %d", j, i) + } + } + } +} + +func TestModel_GenerateRMSE(t *testing.T) { + a := fromGold() + b := fromModel() + + e := rmse(a, b) + + if e > 0 { + t.Fatal("RMSE exceeds threshold") + } +} + +func flatten(s [][]float32) []float32 { + n := 0 + + for _, row := range s { + n += len(row) + } + + r := make([]float32, n) + + i := 0 + + for _, row := range s { + for _, v := range row { + r[i] = v + i++ + } + } + + return r +} + +func mse(a, b []float32) float32 { + if len(a) != len(b) { + panic("length mismatch") + } + + s := float32(0) + + for i := range a { + d := a[i] - b[i] + s += d * d + } + + return s / float32(len(a)) +} + +func rmse(a, b []float32) float32 { + return float32(math.Sqrt(float64(mse(a, b)))) +} |
