summaryrefslogtreecommitdiff
path: root/llm/tokenbuffer.go
blob: 4df0af35ab8160799063720de92db279aaf4915a (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
package llm

import "iter"

type TokenBuffer struct {
	tokenizer   Tokenizer
	window      int
	stride      int
	buffer      []int64
	document    int
	includeTail bool
}

func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer {
	if stride > window {
		panic("stride exceeds window")
	}

	return &TokenBuffer{
		tokenizer:   tokenizer,
		window:      window,
		stride:      stride,
		buffer:      make([]int64, 0, 2*window),
		document:    -1,
		includeTail: true,
	}
}

func (tb *TokenBuffer) IncludeTail() bool {
	return tb.includeTail
}

func (tb *TokenBuffer) SetIncludeTail(includeTail bool) {
	tb.includeTail = includeTail
}

func (tb *TokenBuffer) Push(document int, text string) iter.Seq[[]int64] {
	return func(yield func([]int64) bool) {
		if tb.document != -1 && document != tb.document {
			if tb.includeTail && len(tb.buffer) > 0 {
				if !yield(tb.buffer) {
					return
				}
			}

			tb.buffer = nil
		}

		tb.document = document

		if text == "" {
			return
		}

		ids := toInt64(tb.tokenizer.Tokenize(text))

		tb.buffer = append(tb.buffer, ids...)

		for len(tb.buffer) >= tb.window {
			w := make([]int64, tb.window)

			copy(w, tb.buffer[:tb.window])

			if !yield(w) {
				return
			}

			tb.buffer = append(tb.buffer[:0], tb.buffer[tb.stride:]...)
		}
	}
}

func (tb *TokenBuffer) Tail() []int64 {
	if !tb.includeTail || len(tb.buffer) == 0 {
		return nil
	}

	w := make([]int64, len(tb.buffer))

	copy(w, tb.buffer)

	tb.buffer = nil
	tb.document = -1

	return w
}