Add EmbeddingClient and OpenAI implementation

This commit is contained in:
Mark Pollack
2023-08-12 17:58:58 -04:00
parent 2608c95894
commit abb29137c3
10 changed files with 274 additions and 6 deletions

View File

@@ -27,20 +27,22 @@ public abstract class AbstractChain implements Chain {
@Override
public Map<String, Object> apply(Map<String, Object> inputMap) {
Map<String, Object> inputMapToUse = processBeforeApply(inputMap);
Map<String, Object> inputMapToUse = preProcess(inputMap);
Map<String, Object> outputMap = doApply(inputMapToUse);
Map<String, Object> outputMapToUse = processAfterApply(inputMapToUse, outputMap);
Map<String, Object> outputMapToUse = postProcess(inputMapToUse, outputMap);
return outputMapToUse;
}
protected Map<String, Object> processBeforeApply(Map<String, Object> inputMap) {
@Override
public Map<String, Object> preProcess(Map<String, Object> inputMap) {
validateInputs(inputMap);
return inputMap;
}
protected abstract Map<String, Object> doApply(Map<String, Object> inputMap);
private Map<String, Object> processAfterApply(Map<String, Object> inputMap, Map<String, Object> outputMap) {
@Override
public Map<String, Object> postProcess(Map<String, Object> inputMap, Map<String, Object> outputMap) {
validateOutputs(outputMap);
Map<String, Object> combindedMap = new HashMap<>();
combindedMap.putAll(inputMap);

View File

@@ -10,4 +10,8 @@ public interface Chain extends Function<Map<String, Object>, Map<String, Object>
List<String> getOutputKeys();
Map<String, Object> preProcess(Map<String, Object> inputMap);
Map<String, Object> postProcess(Map<String, Object> inputMap, Map<String, Object> outputMap);
}

View File

@@ -0,0 +1,46 @@
package org.springframework.ai.core.embedding;
import java.util.List;
import java.util.Objects;
public class Embedding {
private List<Double> embedding;
private Integer index;
public Embedding(List<Double> embedding, Integer index) {
this.embedding = embedding;
this.index = index;
}
public List<Double> getEmbedding() {
return embedding;
}
public Integer getIndex() {
return index;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (o == null || getClass() != o.getClass())
return false;
Embedding embedding1 = (Embedding) o;
return Objects.equals(embedding, embedding1.embedding) && Objects.equals(index, embedding1.index);
}
@Override
public int hashCode() {
return Objects.hash(embedding, index);
}
@Override
public String toString() {
String messsage = this.embedding.size() == 0 ? "<empty>" : "<has data>";
return "Embedding{" + "embedding=" + messsage + ", index=" + index + '}';
}
}

View File

@@ -0,0 +1,11 @@
package org.springframework.ai.core.embedding;
import java.util.List;
public interface EmbeddingClient {
List<Double> embed(String text);
EmbeddingResult embed(List<String> texts);
}

View File

@@ -0,0 +1,47 @@
package org.springframework.ai.core.embedding;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
public class EmbeddingResult {
private List<Embedding> data;
private Map<String, Object> metadata = new HashMap<>();
public EmbeddingResult(List<Embedding> data, Map<String, Object> metadata) {
this.data = data;
this.metadata = metadata;
}
public List<Embedding> getData() {
return data;
}
public Map<String, Object> getMetadata() {
return metadata;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (o == null || getClass() != o.getClass())
return false;
EmbeddingResult that = (EmbeddingResult) o;
return Objects.equals(data, that.data) && Objects.equals(metadata, that.metadata);
}
@Override
public int hashCode() {
return Objects.hash(data, metadata);
}
@Override
public String toString() {
return "EmbeddingResult{" + "data=" + data + ", metadata=" + metadata + '}';
}
}

View File

@@ -50,6 +50,7 @@
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -0,0 +1,75 @@
package org.springframework.ai.openai.embedding;
import com.theokanning.openai.Usage;
import com.theokanning.openai.embedding.EmbeddingRequest;
import com.theokanning.openai.service.OpenAiService;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.core.embedding.Embedding;
import org.springframework.ai.core.embedding.EmbeddingClient;
import org.springframework.ai.core.embedding.EmbeddingResult;
import org.springframework.util.Assert;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
public class OpenAiEmbeddingClient implements EmbeddingClient {
private static final Logger logger = LoggerFactory.getLogger(OpenAiEmbeddingClient.class);
private final OpenAiService openAiService;
private String model = "text-embedding-ada-002";
public OpenAiEmbeddingClient(OpenAiService openAiService) {
Assert.notNull(openAiService, "OpenAiService must not be null");
this.openAiService = openAiService;
}
@Override
public EmbeddingResult embed(List<String> texts) {
EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(texts).model(this.model).build();
com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult = this.openAiService
.createEmbeddings(embeddingRequest);
return generateEmbeddingResult(nativeEmbeddingResult);
}
@Override
public List<Double> embed(String text) {
EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(List.of(text)).model(this.model).build();
com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult = this.openAiService
.createEmbeddings(embeddingRequest);
return generateEmbeddingResult(nativeEmbeddingResult).getData().get(0).getEmbedding();
}
private EmbeddingResult generateEmbeddingResult(
com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult) {
List<Embedding> data = generateEmbeddingList(nativeEmbeddingResult.getData());
Map<String, Object> metadata = generateMetadata(nativeEmbeddingResult.getModel(),
nativeEmbeddingResult.getUsage());
return new EmbeddingResult(data, metadata);
}
private Map<String, Object> generateMetadata(String model, Usage usage) {
Map<String, Object> metadata = new HashMap<>();
metadata.put("model", model);
metadata.put("prompt-tokens", usage.getPromptTokens());
metadata.put("completion-tokens", usage.getCompletionTokens());
metadata.put("total-tokens", usage.getTotalTokens());
return metadata;
}
private List<Embedding> generateEmbeddingList(List<com.theokanning.openai.embedding.Embedding> nativeData) {
List<Embedding> data = new ArrayList<>();
for (com.theokanning.openai.embedding.Embedding nativeDatum : nativeData) {
List<Double> nativeDatumEmbedding = nativeDatum.getEmbedding();
int nativeIndex = nativeDatum.getIndex();
Embedding embedding = new Embedding(nativeDatumEmbedding, nativeIndex);
data.add(embedding);
}
return data;
}
}

View File

@@ -0,0 +1,39 @@
package org.springframework.ai.openai;
import com.theokanning.openai.service.OpenAiService;
import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient;
import org.springframework.ai.openai.llm.OpenAiClient;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.context.annotation.Bean;
import org.springframework.util.StringUtils;
import java.io.IOException;
import java.util.Properties;
@SpringBootConfiguration
public class OpenAiTestConfiguration {
@Bean
public OpenAiService theoOpenAiService() throws IOException {
// get api token in file ~/.openai
String apiKey = System.getenv("OPENAI_API_KEY");
if (!StringUtils.hasText(apiKey)) {
throw new IllegalArgumentException(
"You must provide an API key. Put it in an environment variable under the name OPENAI_API_KEY");
}
return new OpenAiService(apiKey);
}
@Bean
public OpenAiClient openAiClient(OpenAiService theoOpenAiService) {
OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService);
return openAiClient;
}
@Bean
public OpenAiEmbeddingClient openAiEmbeddingClient(OpenAiService theoOpenAiService) {
return new OpenAiEmbeddingClient(theoOpenAiService);
}
}

View File

@@ -0,0 +1,33 @@
package org.springframework.ai.openai.embedding;
import org.junit.jupiter.api.Test;
import org.springframework.ai.core.embedding.EmbeddingResult;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import java.util.List;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest
class EmbeddingIntegrationTest {
@Autowired
private OpenAiEmbeddingClient embeddingClient;
@Test
void simpleEmbedding() {
assertThat(embeddingClient).isNotNull();
EmbeddingResult embeddingResult = embeddingClient.embed(List.of("Hello World"));
System.out.println(embeddingResult);
assertThat(embeddingResult.getData()).hasSize(1);
assertThat(embeddingResult.getData().get(0).getEmbedding()).isNotEmpty();
assertThat(embeddingResult.getMetadata()).containsEntry("model", "text-embedding-ada-002-v2");
assertThat(embeddingResult.getMetadata()).containsEntry("completion-tokens", 0L);
assertThat(embeddingResult.getMetadata()).containsEntry("total-tokens", 2L);
assertThat(embeddingResult.getMetadata()).containsEntry("prompt-tokens", 2L);
}
}

View File

@@ -18,6 +18,7 @@ package org.springframework.ai.autoconfigure.openai;
import com.theokanning.openai.service.OpenAiService;
import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient;
import org.springframework.ai.openai.llm.OpenAiClient;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
@@ -39,16 +40,25 @@ public class OpenAiAutoConfiguration {
}
@Bean
public OpenAiClient openAiClient(OpenAiProperties openAiProperties) {
public OpenAiService theoOpenAiService(OpenAiProperties openAiProperties) {
if (!StringUtils.hasText(openAiProperties.getApiKey())) {
throw new IllegalArgumentException(
"You must provide an API key with the property name " + CONFIG_PREFIX + ".api-key");
}
OpenAiService theoOpenAiService = new OpenAiService(openAiProperties.getApiKey());
return new OpenAiService(openAiProperties.getApiKey());
}
@Bean
public OpenAiClient openAiClient(OpenAiProperties openAiProperties, OpenAiService theoOpenAiService) {
OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService);
openAiClient.setTemperature(openAiProperties.getTemperature());
openAiClient.setModel(openAiProperties.getModel());
return openAiClient;
}
@Bean
public OpenAiEmbeddingClient openAiEmbeddingClient(OpenAiService theoOpenAiService) {
return new OpenAiEmbeddingClient(theoOpenAiService);
}
}