Files
AIRouter/internal/context/token_estimator.go
T
freedakgmail 93a469061d
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled
初始提交:边缘AI算力机统一AI通讯层
2026-08-03 07:44:05 +08:00

83 lines
2.5 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package context
import (
"strings"
"unicode"
)
// TokenEstimator estimates token counts for text using a simple heuristic.
// For production use, replace with a proper tokenizer (tiktoken, etc.).
type TokenEstimator struct {
charsPerToken float64
}
// NewTokenEstimator creates a new estimator with the default ratio.
// English text averages ~4 chars/token, Chinese ~1.5 chars/token.
func NewTokenEstimator() *TokenEstimator {
return &TokenEstimator{charsPerToken: 3.0}
}
// EstimateText estimates token count for a given text.
func (e *TokenEstimator) EstimateText(text string) int {
if text == "" {
return 0
}
// Count CJK characters as individual tokens
cjkCount := 0
nonCJKChars := 0
for _, r := range text {
if unicode.Is(unicode.Han, r) || unicode.Is(unicode.Hiragana, r) || unicode.Is(unicode.Katakana, r) || unicode.Is(unicode.Hangul, r) {
cjkCount++
} else {
nonCJKChars++
}
}
// Non-CJK: estimate by chars/token ratio
nonCJKTokens := int(float64(nonCJKChars) / e.charsPerToken)
if nonCJKChars > 0 && nonCJKTokens == 0 {
nonCJKTokens = 1
}
return cjkCount + nonCJKTokens
}
// EstimateMessage estimates token count for a single message (including role overhead).
func (e *TokenEstimator) EstimateMessage(msg interface{ GetRole() string; GetContent() string }) int {
role := msg.GetRole()
content := msg.GetContent()
// Role tokens: ~1-2 tokens for role name
roleTokens := len(strings.Fields(role)) + 1
return roleTokens + e.EstimateText(content)
}
// EstimateMessages estimates total token count for a list of messages.
func (e *TokenEstimator) EstimateMessages(messages []Message) int {
total := 0
for _, m := range messages {
total += e.EstimateText(m.Role) + e.EstimateText(m.Content) + 4 // role + content + formatting overhead
}
return total
}
// Message is a simplified message structure for estimation.
type Message struct {
Role string
Content string
}
func (m Message) GetRole() string { return m.Role }
func (m Message) GetContent() string { return m.Content }
// EstimateKVCache estimates the KV cache memory usage in bytes.
// Formula: input_tokens × layers × 2 (K+V) × hidden_dim × bytes_per_element
func EstimateKVCache(inputTokens, layers, hiddenDim, bytesPerElement int) int64 {
return int64(inputTokens) * int64(layers) * 2 * int64(hiddenDim) * int64(bytesPerElement)
}
// EstimateKVCachePerToken estimates KV cache per token in bytes.
func EstimateKVCachePerToken(layers, hiddenDim, bytesPerElement int) int64 {
return int64(layers) * 2 * int64(hiddenDim) * int64(bytesPerElement)
}