初始提交:边缘AI算力机统一AI通讯层
This commit is contained in:
@@ -0,0 +1,82 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user