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

This commit is contained in:
zivkovicn
2023-07-24 17:25:57 +02:00
parent a698f9b493
commit bc1ac3e00a
9 changed files with 233 additions and 5 deletions
+2 -2
View File
@@ -19,9 +19,9 @@ type Chain interface {
Call(ctx context.Context, inputs map[string]any, options ...ChainCallOption) (map[string]any, error)
// GetMemory gets the memory of the chain.
GetMemory() schema.Memory
// InputKeys returns the input keys the chain expects.
// GetInputKeys returns the input keys the chain expects.
GetInputKeys() []string
// OutputKeys returns the output keys the chain returns.
// GetOutputKeys returns the output keys the chain returns.
GetOutputKeys() []string
}
+1
View File
@@ -8,6 +8,7 @@ import (
"time"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms"
"github.com/tmc/langchaingo/prompts"
"github.com/tmc/langchaingo/schema"
+20
View File
@@ -7,12 +7,14 @@ import (
"testing"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms/openai"
"github.com/tmc/langchaingo/memory"
)
func TestConversation(t *testing.T) {
t.Parallel()
if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" {
t.Skip("OPENAI_API_KEY not set")
}
@@ -27,3 +29,21 @@ func TestConversation(t *testing.T) {
require.NoError(t, err)
require.True(t, strings.Contains(res, "Jim"), `result does not contain the keyword 'Jim'`)
}
func TestConversationMemoryPrune(t *testing.T) {
if openaiKey := os.Getenv("OPENAI_API_KEY"); openaiKey == "" {
t.Skip("OPENAI_API_KEY not set")
}
llm, err := openai.New()
require.NoError(t, err)
c := NewConversation(llm, memory.NewTokenBuffer(llm, 100, memory.WithReturnMessages(true)))
_, 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'`)
}
+4
View File
@@ -5,6 +5,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/schema"
)
@@ -86,6 +87,9 @@ func (t testChatMessageHistory) AddMessage(_ schema.ChatMessage) {
func (t testChatMessageHistory) Clear() {
}
func (t testChatMessageHistory) SetMessages(messages []schema.ChatMessage) {
}
func (t testChatMessageHistory) Messages() []schema.ChatMessage {
return []schema.ChatMessage{
schema.HumanChatMessage{Text: "user message test"},
+4
View File
@@ -37,3 +37,7 @@ func (h *ChatMessageHistory) Clear() {
func (h *ChatMessageHistory) AddMessage(message schema.ChatMessage) {
h.messages = append(h.messages, message)
}
func (h *ChatMessageHistory) SetMessages(messages []schema.ChatMessage) {
h.messages = messages
}
+81
View File
@@ -0,0 +1,81 @@
package memory
import (
"github.com/tmc/langchaingo/llms"
"github.com/tmc/langchaingo/schema"
)
// TokenBuffer for storing conversation memory.
type TokenBuffer struct {
Buffer
LLM llms.LanguageModel
MaxTokenLimit int
}
// Statically assert that TokenBuffer implement the memory interface.
var _ schema.Memory = &TokenBuffer{}
// NewTokenBuffer is a function for crating a new token buffer memory.
func NewTokenBuffer(llm llms.LanguageModel, maxTokenLimit int, options ...BufferOption) *TokenBuffer {
tb := &TokenBuffer{
LLM: llm,
MaxTokenLimit: maxTokenLimit,
Buffer: *applyBufferOptions(options...),
}
return tb
}
// MemoryVariables uses Buffer method for memory variables.
func (tb *TokenBuffer) MemoryVariables() []string {
return tb.Buffer.MemoryVariables()
}
// LoadMemoryVariables uses Buffer method for loading memory variables.
func (tb *TokenBuffer) LoadMemoryVariables(inputs map[string]any) (map[string]any, error) {
return tb.Buffer.LoadMemoryVariables(inputs)
}
// SaveContext uses Buffer method for saving context and prunes memory buffer.
func (tb *TokenBuffer) SaveContext(inputValues map[string]any, outputValues map[string]any) error {
err := tb.Buffer.SaveContext(inputValues, outputValues)
if err != nil {
return err
}
currBufferLength, err := tb.getNumTokensFromMessages()
if err != nil {
return err
}
if currBufferLength > tb.MaxTokenLimit {
// while currBufferLength is greater than MaxTokenLimit we keep removing messages from the memory
for currBufferLength > tb.MaxTokenLimit {
tb.chatHistory.SetMessages(append(tb.chatHistory.Messages()[:0], tb.chatHistory.Messages()[1:]...))
currBufferLength, err = tb.getNumTokensFromMessages()
if err != nil {
return err
}
}
}
return nil
}
// Clear uses Buffer method for clearing buffer memory.
func (tb *TokenBuffer) Clear() error {
return tb.Buffer.Clear()
}
func (tb *TokenBuffer) getNumTokensFromMessages() (int, error) {
sum := 0
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
}
sum += tb.LLM.GetNumTokens(bufferString)
}
return sum, nil
}
+115
View File
@@ -0,0 +1,115 @@
package memory
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tmc/langchaingo/llms/openai"
"github.com/tmc/langchaingo/schema"
)
func TestTokenBufferMemory(t *testing.T) {
t.Parallel()
llm, err := openai.New()
require.NoError(t, err)
m := NewTokenBuffer(llm, 2000)
result1, err := m.LoadMemoryVariables(map[string]any{})
require.NoError(t, err)
expected1 := map[string]any{"history": ""}
assert.Equal(t, expected1, result1)
err = m.SaveContext(map[string]any{"foo": "bar"}, map[string]any{"bar": "foo"})
require.NoError(t, err)
result2, err := m.LoadMemoryVariables(map[string]any{})
require.NoError(t, err)
expected2 := map[string]any{"history": "Human: bar\nAI: foo"}
assert.Equal(t, expected2, result2)
}
func TestTokenBufferMemoryReturnMessage(t *testing.T) {
t.Parallel()
llm, err := openai.New()
require.NoError(t, err)
m := NewTokenBuffer(llm, 2000, WithReturnMessages(true))
expected1 := map[string]any{"history": []schema.ChatMessage{}}
result1, err := m.LoadMemoryVariables(map[string]any{})
require.NoError(t, err)
assert.Equal(t, expected1, result1)
err = m.SaveContext(map[string]any{"foo": "bar"}, map[string]any{"bar": "foo"})
require.NoError(t, err)
result2, err := m.LoadMemoryVariables(map[string]any{})
require.NoError(t, err)
expectedChatHistory := NewChatMessageHistory(
WithPreviousMessages([]schema.ChatMessage{
schema.HumanChatMessage{Text: "bar"},
schema.AIChatMessage{Text: "foo"},
}),
)
expected2 := map[string]any{"history": expectedChatHistory.Messages()}
assert.Equal(t, expected2, result2)
}
func TestTokenBufferMemoryWithPreLoadedHistory(t *testing.T) {
t.Parallel()
llm, err := openai.New()
require.NoError(t, err)
m := NewTokenBuffer(llm, 2000, WithChatHistory(NewChatMessageHistory(
WithPreviousMessages([]schema.ChatMessage{
schema.HumanChatMessage{Text: "bar"},
schema.AIChatMessage{Text: "foo"},
}),
)))
result, err := m.LoadMemoryVariables(map[string]any{})
require.NoError(t, err)
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)
}
+3
View File
@@ -16,4 +16,7 @@ type ChatMessageHistory interface {
// Messages get all messages from the store
Messages() []ChatMessage
// SetMessages replaces existing messages in the store
SetMessages(messages []ChatMessage)
}
+3 -3
View File
@@ -2,12 +2,12 @@ package schema
// Memory is the interface for memory in chains.
type Memory interface {
// Input keys this memory class will load dynamically.
// MemoryVariables Input keys this memory class will load dynamically.
MemoryVariables() []string
// Return key-value pairs given the text input to the chain.
// LoadMemoryVariables Return key-value pairs given the text input to the chain.
// If None, return all memories
LoadMemoryVariables(inputs map[string]any) (map[string]any, error)
// Save the context of this model run to memory.
// SaveContext Save the context of this model run to memory.
SaveContext(inputs map[string]any, outputs map[string]any) error
// Clear memory contents.
Clear() error