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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
|
package llm
import "iter"
type TokenBuffer struct {
tokenizer Tokenizer
window int
stride int
buffer []int64
document int
position int
includeTail bool
config TokenBufferConfig
}
type TokenBufferConfig struct {
Window int
Stride int
PadLeft bool
PadRight bool
PadTokenID int64
}
type TokenWindow struct {
Document int
Tokens []int64
Seen int
PaddingLeft int
PaddingRight int
}
func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
if cfg.Stride > cfg.Window {
panic("stride exceeds window")
}
if cfg.PadLeft && cfg.PadRight {
panic("either pad left or right")
}
return &TokenBuffer{
tokenizer: tokenizer,
window: cfg.Window,
stride: cfg.Stride,
buffer: make([]int64, 0, 2*cfg.Window),
document: -1,
position: 0,
includeTail: true,
config: cfg,
}
}
func (tb *TokenBuffer) IncludeTail() bool {
return tb.includeTail
}
func (tb *TokenBuffer) SetIncludeTail(includeTail bool) {
tb.includeTail = includeTail
}
func (tb *TokenBuffer) Position() int {
return tb.position
}
func (tb *TokenBuffer) Push(document int, text string) iter.Seq[TokenWindow] {
return func(yield func(window TokenWindow) bool) {
if tb.document != -1 && document != tb.document {
w, ok := tb.Tail()
if ok && !yield(w) {
return
}
}
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 := TokenWindow{
Document: tb.document,
Tokens: make([]int64, tb.window),
Seen: min(tb.Position(), tb.window-tb.stride),
}
copy(w.Tokens, tb.buffer[:tb.window])
if !yield(w) {
return
}
tb.position += tb.stride
tb.buffer = append(tb.buffer[:0], tb.buffer[tb.stride:]...)
}
}
}
func (tb *TokenBuffer) Tail() (TokenWindow, bool) {
seen := min(tb.Position(), min(len(tb.buffer), tb.window-tb.stride))
w := TokenWindow{
Document: tb.document,
Seen: seen,
}
tb.document = -1
tb.position = 0
if !tb.includeTail || len(tb.buffer) == 0 {
tb.buffer = tb.buffer[:0]
return w, false
}
w.Tokens = make([]int64, len(tb.buffer))
copy(w.Tokens, tb.buffer)
if tb.config.PadLeft || tb.config.PadRight {
padding := make([]int64, tb.window-len(w.Tokens))
for i := range padding {
padding[i] = tb.config.PadTokenID
}
if tb.config.PadLeft {
w.Tokens = append(padding, w.Tokens...)
w.PaddingLeft = len(padding)
} else {
w.Tokens = append(w.Tokens, padding...)
w.PaddingRight = len(padding)
}
}
tb.buffer = tb.buffer[:0]
return w, true
}
func (tb *TokenBuffer) Stream(docs iter.Seq2[int, string]) iter.Seq[TokenWindow] {
return func(yield func(window TokenWindow) bool) {
for n, d := range docs {
for w := range tb.Push(n, d) {
if !yield(w) {
return
}
}
if w, ok := tb.Tail(); ok && !yield(w) {
return
}
}
}
}
|