summaryrefslogtreecommitdiff
path: root/llm/tokenbuffer.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/tokenbuffer.go')
-rw-r--r--llm/tokenbuffer.go37
1 files changed, 32 insertions, 5 deletions
diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go
index e2d8a91..7bd427c 100644
--- a/llm/tokenbuffer.go
+++ b/llm/tokenbuffer.go
@@ -10,17 +10,23 @@ type TokenBuffer struct {
document int
position int
includeTail bool
+ config TokenBufferConfig
}
type TokenBufferConfig struct {
- Window int
- Stride int
+ Window int
+ Stride int
+ PadLeft bool
+ PadRight bool
+ PadTokenID int64
}
type TokenWindow struct {
- Document int
- Tokens []int64
- Seen int
+ Document int
+ Tokens []int64
+ Seen int
+ PaddingLeft int
+ PaddingRight int
}
func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
@@ -28,6 +34,10 @@ func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
panic("stride exceeds window")
}
+ if cfg.PadLeft && cfg.PadRight {
+ panic("either pad left or right")
+ }
+
return &TokenBuffer{
tokenizer: tokenizer,
window: cfg.Window,
@@ -36,6 +46,7 @@ func NewTokenBuffer(tokenizer Tokenizer, cfg TokenBufferConfig) *TokenBuffer {
document: -1,
position: 0,
includeTail: true,
+ config: cfg,
}
}
@@ -112,6 +123,22 @@ func (tb *TokenBuffer) Tail() (TokenWindow, bool) {
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