feature-adding-max-token-size-memory | lint

This commit is contained in:
zivkovicn
2023-07-24 17:52:47 +02:00
parent 0b144d4f90
commit bf100cedd5
5 changed files with 15 additions and 46 deletions
-1
View File
@@ -8,7 +8,6 @@ import (
"time"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms"
"github.com/tmc/langchaingo/prompts"
"github.com/tmc/langchaingo/schema"
+10 -5
View File
@@ -7,7 +7,6 @@ import (
"testing"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms/openai"
"github.com/tmc/langchaingo/memory"
)
@@ -18,10 +17,10 @@ func TestConversation(t *testing.T) {
if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" {
t.Skip("OPENAI_API_KEY not set")
}
model, err := openai.New()
llm, err := openai.New()
require.NoError(t, err)
c := NewConversation(model, memory.NewBuffer())
c := NewConversation(llm, memory.NewBuffer())
_, err = Run(context.Background(), c, "Hi! I'm Jim")
require.NoError(t, err)
@@ -31,6 +30,8 @@ func TestConversation(t *testing.T) {
}
func TestConversationMemoryPrune(t *testing.T) {
t.Parallel()
if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" {
t.Skip("OPENAI_API_KEY not set")
}
@@ -38,12 +39,16 @@ func TestConversationMemoryPrune(t *testing.T) {
llm, err := openai.New()
require.NoError(t, err)
c := NewConversation(llm, memory.NewTokenBuffer(llm, 100, memory.WithReturnMessages(true)))
c := NewConversation(llm, memory.NewTokenBuffer(llm, 50))
_, err = Run(context.Background(), c, "Hi! I'm Jim")
require.NoError(t, err)
res, err := Run(context.Background(), c, "What is my name?")
require.NoError(t, err)
require.True(t, strings.Contains(res, "Jim"), `result does not contain the keyword 'Jim'`)
require.True(t, strings.Contains(res, "Jim"), `result does contain the keyword 'Jim'`)
// this message will hit the maxTokenLimit and will initiate the prune of the messages to fit the context
res, err = Run(context.Background(), c, "Are you sure that my name is Jim?")
require.NoError(t, err)
require.True(t, strings.Contains(res, "Jim"), `result does contain the keyword 'Jim'`)
}
+1 -2
View File
@@ -5,7 +5,6 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/schema"
)
@@ -87,7 +86,7 @@ func (t testChatMessageHistory) AddMessage(_ schema.ChatMessage) {
func (t testChatMessageHistory) Clear() {
}
func (t testChatMessageHistory) SetMessages(messages []schema.ChatMessage) {
func (t testChatMessageHistory) SetMessages(_ []schema.ChatMessage) {
}
func (t testChatMessageHistory) Messages() []schema.ChatMessage {
+4 -3
View File
@@ -36,7 +36,7 @@ func (tb *TokenBuffer) LoadMemoryVariables(inputs map[string]any) (map[string]an
return tb.Buffer.LoadMemoryVariables(inputs)
}
// SaveContext uses Buffer method for saving context and prunes memory buffer.
// SaveContext uses Buffer method for saving context and prunes memory buffer if needed.
func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[string]any) error {
err := tb.Buffer.SaveContext(inputValues, outputValues)
if err != nil {
@@ -49,8 +49,9 @@ func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[
if currBufferLength > tb.MaxTokenLimit {
// while currBufferLength is greater than MaxTokenLimit we keep removing messages from the memory
// from the oldest
for currBufferLength > tb.MaxTokenLimit {
tb.chatHistory.SetMessages(append(tb.chatHistory.Messages()[:0], tb.chatHistory.Messages()[1:]...))
tb.ChatHistory.SetMessages(append(tb.ChatHistory.Messages()[:0], tb.ChatHistory.Messages()[1:]...))
currBufferLength, err = tb.getNumTokensFromMessages()
if err != nil {
return err
@@ -68,7 +69,7 @@ func (tb *TokenBuffer) Clear() error {
func (tb *TokenBuffer) getNumTokensFromMessages() (int, error) {
sum := 0
for _, message := range tb.chatHistory.Messages() {
for _, message := range tb.ChatHistory.Messages() {
bufferString, err := schema.GetBufferString([]schema.ChatMessage{message}, tb.Buffer.HumanPrefix, tb.Buffer.AIPrefix)
if err != nil {
return 0, err
-35
View File
@@ -5,7 +5,6 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms/openai"
"github.com/tmc/langchaingo/schema"
)
@@ -79,37 +78,3 @@ func TestTokenBufferMemoryWithPreLoadedHistory(t *testing.T) {
expected := map[string]any{"history": "Human: bar\nAI: foo"}
assert.Equal(t, expected, result)
}
func TestTokenBufferMemoryPrune(t *testing.T) {
t.Parallel()
llm, err := openai.New()
require.NoError(t, err)
m := NewTokenBuffer(llm, 20, WithChatHistory(NewChatMessageHistory(
WithPreviousMessages([]schema.ChatMessage{
schema.HumanChatMessage{Text: "human message test for max token"},
schema.AIChatMessage{Text: "ai message test for max token"},
}),
)))
buffStringMsg1, err := schema.GetBufferString([]schema.ChatMessage{
schema.HumanChatMessage{Text: "human message test for max token"},
}, "Human", "AI")
require.NoError(t, err)
tokenNumMsg1 := m.LLM.GetNumTokens(buffStringMsg1)
assert.Equal(t, 9, tokenNumMsg1)
buffStringMsg2, err := schema.GetBufferString([]schema.ChatMessage{
schema.AIChatMessage{Text: "ai message test for max token"},
}, "Human", "AI")
require.NoError(t, err)
tokenNumMsg2 := m.LLM.GetNumTokens(buffStringMsg2)
assert.Equal(t, 8, tokenNumMsg2)
assert.Equal(t, tokenNumMsg1+tokenNumMsg2, 17)
_, err = llm.Call()
require.NoError(t, err)
}