mirror of
https://github.com/vxcontrol/langchaingo.git
synced 2026-07-21 00:45:22 -04:00
feature-adding-max-token-size-memory | cp
This commit is contained in:
+2
-2
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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'`)
|
||||
|
||||
}
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user