From d7b2028f013ebe1eba9fbd8a3c70e098b14c5f75 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Thu, 18 Apr 2024 15:48:13 -0400 Subject: [PATCH] Add Evaluator --- .../chat/agent/OpenAiDefaultChatAgentIT.java | 23 +++--- .../ai/evaluation/EvaluationRequest.java | 58 +++++++++++++ .../ai/evaluation/EvaluationResponse.java | 60 ++++++++++++++ .../ai/evaluation/Evaluator.java | 8 ++ .../ai/evaluation/RelevancyEvaluator.java | 82 +++++++++++++++++++ 5 files changed, 218 insertions(+), 13 deletions(-) create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationResponse.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/evaluation/Evaluator.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java index dcddf2520..16d3c4905 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java @@ -5,14 +5,14 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.agent.DefaultChatAgent; import org.springframework.ai.chat.agent.PromptContext; - import org.springframework.ai.chat.agent.transformer.QAPromptContextTransformer; - import org.springframework.ai.chat.agent.transformer.VectorStorePromptContextTransformer; - import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.evaluation.EvaluationRequest; +import org.springframework.ai.evaluation.EvaluationResponse; +import org.springframework.ai.evaluation.RelevancyEvaluator; import org.springframework.ai.openai.OpenAiChatClient; import org.springframework.ai.openai.OpenAiEmbeddingClient; import org.springframework.ai.openai.api.OpenAiApi; @@ -61,16 +61,13 @@ public class OpenAiDefaultChatAgentIT { var agentResponse = chatAgent.call(promptContext); System.out.println(agentResponse.getChatResponse().getResult().getOutput().getContent()); - // RelevancyEvaluator relevancyEvaluator = new - // RelevancyEvaluator(this.chatClient); - // EvaluationRequest evaluationRequest = new EvaluationRequest( - // agentResponse.getPromptContext().getOriginalPrompt(), - // agentResponse.getPromptContext().getDataList(), - // agentResponse.getChatResponse()); - // - // EvaluationResponse evaluationResponse = - // relevancyEvaluator.evaluate(evaluationRequest); - // System.out.println(evaluationResponse); + RelevancyEvaluator relevancyEvaluator = new RelevancyEvaluator(this.chatClient); + EvaluationRequest evaluationRequest = new EvaluationRequest( + agentResponse.getPromptContext().getPromptHistory().get(0), agentResponse.getPromptContext().getNodes(), + agentResponse.getChatResponse()); + + EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(evaluationRequest); + System.out.println(evaluationResponse); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java new file mode 100644 index 000000000..1081f37c6 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java @@ -0,0 +1,58 @@ +package org.springframework.ai.evaluation; + +import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.node.Node; + +import java.util.ArrayList; +import java.util.List; +import java.util.Objects; + +public class EvaluationRequest { + + private Prompt prompt; + + private List> dataList; + + private ChatResponse chatResponse; + + public EvaluationRequest(Prompt prompt, List> dataList, ChatResponse chatResponse) { + this.prompt = prompt; + this.dataList = dataList; + this.chatResponse = chatResponse; + } + + public Prompt getPrompt() { + return prompt; + } + + public List> getDataList() { + return dataList; + } + + public ChatResponse getChatResponse() { + return chatResponse; + } + + @Override + public String toString() { + return "EvaluationRequest{" + "prompt=" + prompt + ", dataList=" + dataList + ", chatResponse=" + chatResponse + + '}'; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof EvaluationRequest that)) + return false; + return Objects.equals(prompt, that.prompt) && Objects.equals(dataList, that.dataList) + && Objects.equals(chatResponse, that.chatResponse); + } + + @Override + public int hashCode() { + return Objects.hash(prompt, dataList, chatResponse); + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationResponse.java new file mode 100644 index 000000000..a22d738c6 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationResponse.java @@ -0,0 +1,60 @@ +package org.springframework.ai.evaluation; + +import java.util.Map; +import java.util.Objects; + +public class EvaluationResponse { + + private boolean pass; + + private float score; + + private String feedback; + + Map metadata; + + public EvaluationResponse(boolean pass, float score, String feedback, Map metadata) { + this.pass = pass; + this.score = score; + this.feedback = feedback; + this.metadata = metadata; + } + + public boolean isPass() { + return pass; + } + + public float getScore() { + return score; + } + + public String getFeedback() { + return feedback; + } + + public Map getMetadata() { + return metadata; + } + + @Override + public String toString() { + return "EvaluationResponse{" + "pass=" + pass + ", score=" + score + ", feedback='" + feedback + '\'' + + ", metadata=" + metadata + '}'; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof EvaluationResponse that)) + return false; + return pass == that.pass && Float.compare(score, that.score) == 0 && Objects.equals(feedback, that.feedback) + && Objects.equals(metadata, that.metadata); + } + + @Override + public int hashCode() { + return Objects.hash(pass, score, feedback, metadata); + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/Evaluator.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/Evaluator.java new file mode 100644 index 000000000..7cfdbbf67 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/Evaluator.java @@ -0,0 +1,8 @@ +package org.springframework.ai.evaluation; + +@FunctionalInterface +public interface Evaluator { + + EvaluationResponse evaluate(EvaluationRequest evaluationRequest); + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java new file mode 100644 index 000000000..38e08d068 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java @@ -0,0 +1,82 @@ +package org.springframework.ai.evaluation; + +import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.ai.node.Node; + +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; + +public class RelevancyEvaluator implements Evaluator { + + private static final String DEFAULT_EVALUATION_PROMPT_TEXT = """ + Your task is to evaluate if the response for the query + is in line with the context information provided.\\n + You have two options to answer. Either YES/ NO.\\n + Answer - YES, if the response for the query + is in line with context information otherwise NO.\\n + Query: \\n {query}\\n + Response: \\n {response}\\n + Context: \\n {context}\\n + Answer: " + """; + + private ChatClient chatClient; + + public RelevancyEvaluator(ChatClient chatClient) { + this.chatClient = chatClient; + } + + @Override + public EvaluationResponse evaluate(EvaluationRequest evaluationRequest) { + var query = doGetUserQuestion(evaluationRequest); + var response = doGetResponse(evaluationRequest); + var context = doGetSupportingData(evaluationRequest); + + var promptTemplate = new PromptTemplate(DEFAULT_EVALUATION_PROMPT_TEXT); + Message message = promptTemplate + .createMessage(Map.of("query", query, "response", response, "context", context)); + + ChatResponse chatResponse = this.chatClient.call(new Prompt(message)); + + var evaluationResponse = chatResponse.getResult().getOutput().getContent(); + boolean passing = false; + float score = 0; + if (evaluationResponse.toLowerCase().contains("yes")) { + passing = true; + score = 1; + } + + return new EvaluationResponse(passing, score, "", Collections.emptyMap()); + } + + protected String doGetResponse(EvaluationRequest evaluationRequest) { + return evaluationRequest.getChatResponse().getResult().getOutput().getContent(); + } + + protected String doGetSupportingData(EvaluationRequest evaluationRequest) { + List> data = evaluationRequest.getDataList(); + String supportingData = data.stream() + .filter(node -> node != null && node.getContent() instanceof String) + .map(node -> (Node) node) + .map(Node::getContent) + .collect(Collectors.joining("\n")); + return supportingData; + } + + protected String doGetUserQuestion(EvaluationRequest evaluationRequest) { + List instructions = evaluationRequest.getPrompt().getInstructions(); + String userMessage = instructions.stream() + .filter(m -> m.getMessageType() == MessageType.USER) + .map(m -> m.getContent()) + .collect(Collectors.joining(System.lineSeparator())); + return userMessage; + } + +}